Compare commits
3 Commits
28ef360d5f
...
d5e30f8c36
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d5e30f8c36 | ||
|
|
a803b80d59 | ||
|
|
8094aaab10 |
52
README.md
52
README.md
@@ -64,33 +64,45 @@ python main.py
|
|||||||
> `nginx`. Это также позволит вам использовать защищённое соединение, что
|
> `nginx`. Это также позволит вам использовать защищённое соединение, что
|
||||||
> исключит возможность применения атаки Man-in-the-Middle для перехвата токена.
|
> исключит возможность применения атаки Man-in-the-Middle для перехвата токена.
|
||||||
|
|
||||||
Все запросы к веб-серверу требуют авторизации, используя HTTP заголовок
|
Все запросы к веб-серверу являются GET-запросами и требуют авторизации.
|
||||||
`Authorization` и схему `Bearer`. Токен авторизации генерируется посредством
|
Авторизоваться можно двумя путями:
|
||||||
взаимодействия с ботом в Matrix.
|
1. Использовать HTTP-заголовок `Authorization` и схему `Bearer`. Например:
|
||||||
|
```plain
|
||||||
|
Authorization: Bearer 1234567890abcdef
|
||||||
|
```
|
||||||
|
2. Использовать URL параметр `token`. Например:
|
||||||
|
```plain
|
||||||
|
https://csonac.su/notify?token=f1829e94d...
|
||||||
|
```
|
||||||
|
Рекомендуется использовать первый способ (HTTP-заголовок), так как это позволяет
|
||||||
|
избежать раскрытия токена в логах сервера и других местах, где можно посмотреть
|
||||||
|
URL прошлых запросов.
|
||||||
|
|
||||||
> Для каждого сервиса, использующего бота, рекомендуется генерировать свой
|
> Для каждого сервиса, использующего бота, рекомендуется генерировать свой
|
||||||
> собственный токен. Это позволит отозвать токен только для одного серсива, если
|
> собственный токен. Это позволит отозвать токен только для одного серсива, если
|
||||||
> токен будет украден.
|
> токен будет украден.
|
||||||
|
|
||||||
Доступные эндпоинты:
|
Доступные эндпоинты:
|
||||||
- `POST /<channel>/notify`
|
- `GET /notify`
|
||||||
- **Описание.** Используется, чтобы отправить уведомление в указанный канал.
|
- **Описание.** Используется, чтобы отправить уведомление в канал.
|
||||||
Вместо `<channel>` указывается код канала, получаемый при помощи команды
|
- **Параметры запроса**
|
||||||
`!info`, выполненной в комнате Matrix.
|
- `channel` - код канала, получаемый при помощи команды `!info`,
|
||||||
- **Тело запроса.** Тело запроса представляет собой `json` объект:
|
выполненной в комнате Matrix
|
||||||
```json
|
- `service` - имя сервиса (учитывается, только если для токена не было
|
||||||
|
настроено имя сервиса через бота)
|
||||||
|
- `text` - текст уведомления (форматирование не поддерживается)
|
||||||
|
- **Ответ**
|
||||||
|
- В случае успеха сервер вернёт `200` и JSON следующего формата:
|
||||||
|
```json
|
||||||
{
|
{
|
||||||
"service": "<название сервиса; указывается, если токен это позволяет>",
|
"notification_id": "Notification ID will be here"
|
||||||
"text": "<текст уведомления>"
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
- **Тело ответа.** Тело ответа представляет собой `json` объект. В случае
|
> В текущей версии `notification_id` не имеет практической пользы и его
|
||||||
успеха в нём будут все поля, перечисляемые ниже. В случае провала - только
|
> формат будет меняться.
|
||||||
поле `error`, содержащее текстовое описание ошибки.
|
- В случае ошибки сервер вернёт JSON следующего формата:
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"error": null,
|
"detail": "Error description in English"
|
||||||
"notification_id": "<здесь будет Notification ID>"
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
- ``
|
|
||||||
@@ -11,7 +11,11 @@ import util
|
|||||||
DEFAULT_CONFIG = {
|
DEFAULT_CONFIG = {
|
||||||
"matrix_homeserver": "https://matrix.domain.net",
|
"matrix_homeserver": "https://matrix.domain.net",
|
||||||
"matrix_user": "short_username",
|
"matrix_user": "short_username",
|
||||||
"store_dir": "session_storage"
|
"store_dir": "session_storage",
|
||||||
|
"web_ip": "0.0.0.0",
|
||||||
|
"web_port": 4980,
|
||||||
|
"fails_to_ban": 50,
|
||||||
|
"ban_duration": 600
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -50,4 +54,5 @@ async def load_config(path: str = "config.json") -> AppConfig | None:
|
|||||||
cfg = AppConfig(**j)
|
cfg = AppConfig(**j)
|
||||||
return cfg
|
return cfg
|
||||||
except:
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
return None
|
return None
|
||||||
58
database.py
58
database.py
@@ -84,20 +84,33 @@ class Database:
|
|||||||
wait_task = asyncio.create_task(
|
wait_task = asyncio.create_task(
|
||||||
asyncio.sleep(self.BACKGROUND_ROUTINE_PERIOD)
|
asyncio.sleep(self.BACKGROUND_ROUTINE_PERIOD)
|
||||||
)
|
)
|
||||||
|
if self._connection is None:
|
||||||
|
continue
|
||||||
|
# get current time
|
||||||
|
current_time = time.time()
|
||||||
# delete old bans
|
# delete old bans
|
||||||
try:
|
try:
|
||||||
if self._connection is not None:
|
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
||||||
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
async with self._connection.execute(statement, (current_time,)) as cursor:
|
||||||
async with self._connection.execute(statement, (time.time(),)) as cursor:
|
async for row in cursor:
|
||||||
async for row in cursor:
|
ip = row["ip"]
|
||||||
ip = row["ip"]
|
if ip in self._bans:
|
||||||
if ip in self._bans:
|
del self._bans[ip]
|
||||||
del self._bans[ip]
|
self._logger.debug(f"IP {ip} is not banned anymore")
|
||||||
self._logger.debug(f"IP {ip} is not banned anymore")
|
await self._connection.commit()
|
||||||
await self._connection.commit()
|
|
||||||
self._logger.debug("Performed banned IPs cleanup")
|
self._logger.debug("Performed banned IPs cleanup")
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
|
# delete expired fails
|
||||||
|
try:
|
||||||
|
for ip in dict(self._fails):
|
||||||
|
self._fails[ip] = list(
|
||||||
|
filter(lambda x: x > current_time, self._fails[ip])
|
||||||
|
)
|
||||||
|
if not self._fails[ip]:
|
||||||
|
del self._fails[ip]
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
self._logger.debug("Stopped background worker")
|
self._logger.debug("Stopped background worker")
|
||||||
|
|
||||||
async def _load_bans_from_database(self) -> None:
|
async def _load_bans_from_database(self) -> None:
|
||||||
@@ -115,6 +128,12 @@ class Database:
|
|||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
return
|
return
|
||||||
|
|
||||||
|
def _get_fails_for_ip(self, ip) -> int:
|
||||||
|
if ip not in self._fails:
|
||||||
|
return 0
|
||||||
|
t = time.time()
|
||||||
|
return sum([(1 if expires_at > t else 0) for expires_at in self._fails[ip]])
|
||||||
|
|
||||||
#
|
#
|
||||||
# PUBLIC
|
# PUBLIC
|
||||||
#
|
#
|
||||||
@@ -127,6 +146,7 @@ class Database:
|
|||||||
self._background_stop_event: asyncio.Event | None = None
|
self._background_stop_event: asyncio.Event | None = None
|
||||||
self._background_task: asyncio.Task | None = None
|
self._background_task: asyncio.Task | None = None
|
||||||
self._bans: dict = {}
|
self._bans: dict = {}
|
||||||
|
self._fails: dict[str, list[float]] = {}
|
||||||
|
|
||||||
async def connect(self) -> bool:
|
async def connect(self) -> bool:
|
||||||
"""Connect to the database. Returns False on failure."""
|
"""Connect to the database. Returns False on failure."""
|
||||||
@@ -137,6 +157,7 @@ class Database:
|
|||||||
self._connection.row_factory = Row
|
self._connection.row_factory = Row
|
||||||
await self._setup_tables(self._connection)
|
await self._setup_tables(self._connection)
|
||||||
self._bans = {}
|
self._bans = {}
|
||||||
|
self._fails = {}
|
||||||
await self._load_bans_from_database()
|
await self._load_bans_from_database()
|
||||||
self._background_stop_event = asyncio.Event()
|
self._background_stop_event = asyncio.Event()
|
||||||
self._background_task = asyncio.create_task(
|
self._background_task = asyncio.create_task(
|
||||||
@@ -346,4 +367,21 @@ class Database:
|
|||||||
if ip in self._bans:
|
if ip in self._bans:
|
||||||
del self._bans[ip]
|
del self._bans[ip]
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
|
||||||
|
def fail_create(self, ip: str, expires_at: float) -> int:
|
||||||
|
"""Adds failed access attempt. Returns count of failed attempts for IP."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
if ip not in self._fails:
|
||||||
|
self._fails[ip] = []
|
||||||
|
self._fails[ip].append(expires_at)
|
||||||
|
return self._get_fails_for_ip(ip)
|
||||||
|
|
||||||
|
def fail_clear(self, ip: str) -> None:
|
||||||
|
"""Clears failed attempts counter."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
if ip in self._fails:
|
||||||
|
del self._fails[ip]
|
||||||
@@ -10,6 +10,10 @@ class AppConfig:
|
|||||||
matrix_homeserver: str
|
matrix_homeserver: str
|
||||||
matrix_user: str
|
matrix_user: str
|
||||||
store_dir: str
|
store_dir: str
|
||||||
|
web_ip: str
|
||||||
|
web_port: int
|
||||||
|
fails_to_ban: int
|
||||||
|
ban_duration: float
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ObjectRoom:
|
class ObjectRoom:
|
||||||
|
|||||||
5
main.py
5
main.py
@@ -14,6 +14,7 @@ import database
|
|||||||
import config
|
import config
|
||||||
import util
|
import util
|
||||||
import bot_callbacks
|
import bot_callbacks
|
||||||
|
import web
|
||||||
|
|
||||||
bot: MatrixBot = None # type: ignore
|
bot: MatrixBot = None # type: ignore
|
||||||
|
|
||||||
@@ -47,17 +48,21 @@ async def main() -> None:
|
|||||||
db = database.Database(Path("database.sqlite"))
|
db = database.Database(Path("database.sqlite"))
|
||||||
# setup the callbacks
|
# setup the callbacks
|
||||||
bot_callbacks.setup(bot, db)
|
bot_callbacks.setup(bot, db)
|
||||||
|
# setup the web server
|
||||||
|
web_server = web.Web(cfg, bot, db)
|
||||||
|
|
||||||
# start the app
|
# start the app
|
||||||
if not await db.connect():
|
if not await db.connect():
|
||||||
util.log_error("Can't connect to the database!")
|
util.log_error("Can't connect to the database!")
|
||||||
return
|
return
|
||||||
bot.start()
|
bot.start()
|
||||||
|
await web_server.start()
|
||||||
|
|
||||||
# wait for Ctrl+C
|
# wait for Ctrl+C
|
||||||
await util.get_app_stop_event().wait()
|
await util.get_app_stop_event().wait()
|
||||||
|
|
||||||
# stop the app
|
# stop the app
|
||||||
|
await web_server.stop()
|
||||||
await bot.stop()
|
await bot.stop()
|
||||||
await db.disconnect()
|
await db.disconnect()
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
agent-detector==1.1.0
|
||||||
aioconsole==0.8.2
|
aioconsole==0.8.2
|
||||||
aiofiles==25.1.0
|
aiofiles==25.1.0
|
||||||
aiohappyeyeballs==2.7.1
|
aiohappyeyeballs==2.7.1
|
||||||
@@ -5,29 +6,70 @@ aiohttp==3.14.3
|
|||||||
aiohttp_socks==0.12.0
|
aiohttp_socks==0.12.0
|
||||||
aiosignal==1.4.0
|
aiosignal==1.4.0
|
||||||
aiosqlite==0.22.1
|
aiosqlite==0.22.1
|
||||||
|
annotated-doc==0.0.5
|
||||||
|
annotated-types==0.8.0
|
||||||
|
anyio==4.14.2
|
||||||
atomicwrites==1.4.1
|
atomicwrites==1.4.1
|
||||||
attrs==26.1.0
|
attrs==26.1.0
|
||||||
build==1.5.0
|
build==1.5.0
|
||||||
cachetools==7.1.7
|
cachetools==7.1.7
|
||||||
|
certifi==2026.7.22
|
||||||
|
click==8.5.0
|
||||||
|
detect-installer==0.1.0
|
||||||
|
dnspython==2.8.0
|
||||||
|
email-validator==2.3.0
|
||||||
|
fastapi==0.141.1
|
||||||
|
fastapi-cli==0.0.32
|
||||||
|
fastapi-cloud-cli==0.24.0
|
||||||
|
fastar==0.12.0
|
||||||
frozenlist==1.8.0
|
frozenlist==1.8.0
|
||||||
h11==0.16.0
|
h11==0.16.0
|
||||||
h2==4.4.1
|
h2==4.4.1
|
||||||
hpack==4.2.0
|
hpack==4.2.0
|
||||||
|
httpcore==1.0.9
|
||||||
|
httptools==0.8.0
|
||||||
|
httpx==0.28.1
|
||||||
hyperframe==6.1.0
|
hyperframe==6.1.0
|
||||||
idna==3.19
|
idna==3.19
|
||||||
|
Jinja2==3.1.6
|
||||||
jsonschema==4.26.0
|
jsonschema==4.26.0
|
||||||
jsonschema-specifications==2025.9.1
|
jsonschema-specifications==2025.9.1
|
||||||
mab @ git+https://git.tyukalov.su/nikita/mab@c3046307c7bc4e65113aab62b07bad2a148b49e5
|
mab @ git+https://git.tyukalov.su/nikita/mab@c3046307c7bc4e65113aab62b07bad2a148b49e5
|
||||||
|
markdown-it-py==4.2.0
|
||||||
|
MarkupSafe==3.0.3
|
||||||
matrix-nio==0.26.0
|
matrix-nio==0.26.0
|
||||||
|
mdurl==0.1.2
|
||||||
multidict==6.7.1
|
multidict==6.7.1
|
||||||
packaging==26.3
|
packaging==26.3
|
||||||
peewee==3.19.0
|
peewee==3.19.0
|
||||||
propcache==0.5.2
|
propcache==0.5.2
|
||||||
pycryptodome==3.23.0
|
pycryptodome==3.23.0
|
||||||
|
pydantic==2.13.5
|
||||||
|
pydantic-extra-types==2.11.1
|
||||||
|
pydantic-settings==2.15.0
|
||||||
|
pydantic_core==2.46.5
|
||||||
|
Pygments==2.21.0
|
||||||
pyproject_hooks==1.2.0
|
pyproject_hooks==1.2.0
|
||||||
|
python-dotenv==1.2.3
|
||||||
|
python-multipart==0.0.32
|
||||||
python-socks==3.0.0
|
python-socks==3.0.0
|
||||||
|
PyYAML==6.0.3
|
||||||
referencing==0.37.0
|
referencing==0.37.0
|
||||||
|
rich==15.0.0
|
||||||
|
rich-toolkit==0.20.3
|
||||||
|
rignore==0.8.1
|
||||||
rpds-py==2026.6.3
|
rpds-py==2026.6.3
|
||||||
|
sentry-sdk==2.68.1
|
||||||
|
shellingham==1.5.4
|
||||||
|
starlette==1.6.0
|
||||||
|
typer==0.27.2
|
||||||
|
typing-inspection==0.4.4
|
||||||
|
typing_extensions==4.16.0
|
||||||
unpaddedbase64==2.1.0
|
unpaddedbase64==2.1.0
|
||||||
|
urllib3==2.7.0
|
||||||
|
uvicorn==0.52.4
|
||||||
|
uvloop==0.22.1
|
||||||
vodozemac==0.10.0
|
vodozemac==0.10.0
|
||||||
|
watchfiles==1.2.0
|
||||||
|
websockets==17.1
|
||||||
yarl==1.24.5
|
yarl==1.24.5
|
||||||
|
|||||||
73
web.py
Normal file
73
web.py
Normal file
@@ -0,0 +1,73 @@
|
|||||||
|
"""This module implements API"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import traceback
|
||||||
|
import time
|
||||||
|
import html
|
||||||
|
|
||||||
|
from fastapi import FastAPI, Request
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
|
from mab import MatrixBot
|
||||||
|
import uvicorn
|
||||||
|
|
||||||
|
from datatypes import AppConfig
|
||||||
|
from database import Database
|
||||||
|
|
||||||
|
import web_routes
|
||||||
|
import web_middleware
|
||||||
|
|
||||||
|
class Web:
|
||||||
|
#
|
||||||
|
# PRIVATE
|
||||||
|
#
|
||||||
|
async def _cb_exception(self, request: Request, exc: Exception):
|
||||||
|
traceback.print_exception(exc)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={"detail": "Internal Server Error"}
|
||||||
|
)
|
||||||
|
|
||||||
|
#
|
||||||
|
# PUBLIC
|
||||||
|
#
|
||||||
|
def __init__(self, app_config: AppConfig, bot: MatrixBot, db: Database):
|
||||||
|
self._api = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
self._api.state.database = db
|
||||||
|
self._api.state.matrix = bot
|
||||||
|
self._api.state.app_config = app_config
|
||||||
|
|
||||||
|
self._api.include_router(web_routes.router)
|
||||||
|
self._api.add_middleware(web_middleware.AuthMiddleware, ["/api"])
|
||||||
|
self._api.add_middleware(web_middleware.BanCheckMiddleware, ["/api"])
|
||||||
|
self._api.add_middleware(web_middleware.RealIpResolver)
|
||||||
|
self._api.add_exception_handler(Exception, self._cb_exception)
|
||||||
|
|
||||||
|
self._server_config = uvicorn.Config(
|
||||||
|
self._api,
|
||||||
|
host=app_config.web_ip,
|
||||||
|
port=app_config.web_port
|
||||||
|
)
|
||||||
|
self._server: uvicorn.Server | None = None
|
||||||
|
self._server_task: asyncio.Task | None = None
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
"""Start the API."""
|
||||||
|
if self._server is not None or self._server_task is not None:
|
||||||
|
raise RuntimeError("The server is already started")
|
||||||
|
self._server = uvicorn.Server(self._server_config)
|
||||||
|
self._server_task = asyncio.create_task(self._server.serve())
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Stop the API server."""
|
||||||
|
if self._server is None or self._server_task is None:
|
||||||
|
raise RuntimeError("The server is not started yet")
|
||||||
|
self._server_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._server_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
|
self._server_task = None
|
||||||
|
self._server = None
|
||||||
135
web_middleware.py
Normal file
135
web_middleware.py
Normal file
@@ -0,0 +1,135 @@
|
|||||||
|
import time
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
from fastapi import status
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
from database import Database
|
||||||
|
from datatypes import ObjectToken, AppConfig
|
||||||
|
|
||||||
|
|
||||||
|
class AuthMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware for authorization"""
|
||||||
|
def _is_request_protected(self, request: Request) -> bool:
|
||||||
|
"""Returns True if request URL starts with one of private prefixes"""
|
||||||
|
for pp in self.private_prefixes:
|
||||||
|
if request.url.path.startswith(pp):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _fail_and_ban_if_needed(self, request: Request) -> None:
|
||||||
|
"""Adds a fail and bans the IP if too much failures are recorded."""
|
||||||
|
app_config: AppConfig = request.app.state.app_config
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
current_time: float = time.time()
|
||||||
|
ip: str = request.state.ip
|
||||||
|
total_fails = database.fail_create(ip, current_time + 10.0)
|
||||||
|
if total_fails >= app_config.fails_to_ban:
|
||||||
|
await database.ban_create(
|
||||||
|
ip,
|
||||||
|
current_time + app_config.ban_duration,
|
||||||
|
"Banned by AuthMiddleware"
|
||||||
|
)
|
||||||
|
database.fail_clear(ip)
|
||||||
|
|
||||||
|
def __init__(self, app, private_prefixes: list[str]):
|
||||||
|
super().__init__(app)
|
||||||
|
self.private_prefixes = private_prefixes
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# do not check authorization for public endpoints
|
||||||
|
if not self._is_request_protected(request):
|
||||||
|
return await call_next(request)
|
||||||
|
token = None
|
||||||
|
# look for the token in headers
|
||||||
|
if "Authorization" in request.headers:
|
||||||
|
v = request.headers["Authorization"]
|
||||||
|
if not v.lower().startswith("bearer"):
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "Invalid authorization scheme"}
|
||||||
|
)
|
||||||
|
parts = [p for p in v.split(" ") if p.strip()]
|
||||||
|
if len(parts) != 2:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No credentials provided in HTTP header"}
|
||||||
|
)
|
||||||
|
token = parts[1]
|
||||||
|
# look for the token in query parameters
|
||||||
|
elif "token" in request.query_params:
|
||||||
|
token = request.query_params.get("token")
|
||||||
|
if not token:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "The token is empry"}
|
||||||
|
)
|
||||||
|
# no token found
|
||||||
|
else:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No token provided"}
|
||||||
|
)
|
||||||
|
# try to authorize
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
token_data: ObjectToken | None = await database.token_get(token)
|
||||||
|
if not token_data:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No such token"}
|
||||||
|
)
|
||||||
|
request.state.token = token_data
|
||||||
|
return await call_next(request)
|
||||||
|
|
||||||
|
|
||||||
|
class BanCheckMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware for checking if the IP is banned."""
|
||||||
|
def _is_request_protected(self, request: Request) -> bool:
|
||||||
|
"""Returns True if request URL starts with one of private prefixes"""
|
||||||
|
for pp in self.private_prefixes:
|
||||||
|
if request.url.path.startswith(pp):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def __init__(self, app, private_prefixes: list[str]):
|
||||||
|
super().__init__(app)
|
||||||
|
self.private_prefixes = private_prefixes
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# do not check authorization for public endpoints
|
||||||
|
if not self._is_request_protected(request):
|
||||||
|
return await call_next(request)
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
if database.ban_check(request.state.ip):
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "You are temporarily banned"}
|
||||||
|
)
|
||||||
|
return await call_next(request)
|
||||||
|
|
||||||
|
|
||||||
|
class RealIpResolver(BaseHTTPMiddleware):
|
||||||
|
"""Middleware that gets real IP address of the client."""
|
||||||
|
def __init__(self, app):
|
||||||
|
super().__init__(app)
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# assign None by default
|
||||||
|
request.state.ip = None
|
||||||
|
# look for IP address
|
||||||
|
x_forwarded_for = request.headers.get("x-forwarded-for")
|
||||||
|
if x_forwarded_for:
|
||||||
|
request.state.ip = x_forwarded_for.split(",")[0].strip()
|
||||||
|
else:
|
||||||
|
x_real_ip = request.headers.get("x-real-ip")
|
||||||
|
if x_real_ip:
|
||||||
|
request.state.ip = x_real_ip
|
||||||
|
else:
|
||||||
|
request.state.ip = request.client.host if request.client else None
|
||||||
|
return await call_next(request)
|
||||||
70
web_routes.py
Normal file
70
web_routes.py
Normal file
@@ -0,0 +1,70 @@
|
|||||||
|
"""API Routes"""
|
||||||
|
|
||||||
|
import html
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Request, Depends
|
||||||
|
from fastapi import HTTPException, status
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from mab import MatrixBot
|
||||||
|
from database import Database
|
||||||
|
from datatypes import *
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api")
|
||||||
|
|
||||||
|
#
|
||||||
|
# MODELS
|
||||||
|
#
|
||||||
|
class ModelGetNotify(BaseModel):
|
||||||
|
channel: str = Field(
|
||||||
|
...,
|
||||||
|
min_length=11,
|
||||||
|
max_length=11,
|
||||||
|
description="Notification channel (use `!info` bot command)"
|
||||||
|
)
|
||||||
|
text: str = Field(
|
||||||
|
...,
|
||||||
|
min_length=1,
|
||||||
|
max_length=1024,
|
||||||
|
description="Text to use as notification body"
|
||||||
|
)
|
||||||
|
service: str | None = Field(
|
||||||
|
None,
|
||||||
|
min_length=1,
|
||||||
|
max_length=1024,
|
||||||
|
description="Service name to use (ignored if token has name)"
|
||||||
|
)
|
||||||
|
|
||||||
|
#
|
||||||
|
# ENDPOINTS
|
||||||
|
#
|
||||||
|
@router.get("/notify")
|
||||||
|
async def _get_notify(request: Request, params: ModelGetNotify = Depends()):
|
||||||
|
# improve readability
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
matrix: MatrixBot = request.app.state.matrix
|
||||||
|
token: ObjectToken = request.state.token
|
||||||
|
# check if room exists
|
||||||
|
room = await database.room_get(params.channel)
|
||||||
|
if room is None:
|
||||||
|
raise HTTPException(status.HTTP_200_OK, "No such channel")
|
||||||
|
# prepare service name
|
||||||
|
service = token.name or params.service or "Unnamed"
|
||||||
|
# prepare text of the notification
|
||||||
|
text = f"<strong>Уведомление от службы <code>{html.escape(service)}</code></strong>"
|
||||||
|
text += "<br>" * 2
|
||||||
|
text += html.escape(params.text)
|
||||||
|
# try to send the notification
|
||||||
|
try:
|
||||||
|
notification_id = await matrix.send_text_to_room(
|
||||||
|
room.matrix_id,
|
||||||
|
text,
|
||||||
|
is_html=True
|
||||||
|
)
|
||||||
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
|
raise HTTPException(status.HTTP_200_OK, "Failed to send matrix message")
|
||||||
|
# success
|
||||||
|
return {"notification_id": notification_id}
|
||||||
Reference in New Issue
Block a user