First API implementation
This commit is contained in:
52
README.md
52
README.md
@@ -64,33 +64,45 @@ python main.py
|
||||
> `nginx`. Это также позволит вам использовать защищённое соединение, что
|
||||
> исключит возможность применения атаки Man-in-the-Middle для перехвата токена.
|
||||
|
||||
Все запросы к веб-серверу требуют авторизации, используя HTTP заголовок
|
||||
`Authorization` и схему `Bearer`. Токен авторизации генерируется посредством
|
||||
взаимодействия с ботом в Matrix.
|
||||
Все запросы к веб-серверу являются GET-запросами и требуют авторизации.
|
||||
Авторизоваться можно двумя путями:
|
||||
1. Использовать HTTP-заголовок `Authorization` и схему `Bearer`. Например:
|
||||
```plain
|
||||
Authorization: Bearer 1234567890abcdef
|
||||
```
|
||||
2. Использовать URL параметр `token`. Например:
|
||||
```plain
|
||||
https://csonac.su/notify?token=f1829e94d...
|
||||
```
|
||||
Рекомендуется использовать первый способ (HTTP-заголовок), так как это позволяет
|
||||
избежать раскрытия токена в логах сервера и других местах, где можно посмотреть
|
||||
URL прошлых запросов.
|
||||
|
||||
> Для каждого сервиса, использующего бота, рекомендуется генерировать свой
|
||||
> собственный токен. Это позволит отозвать токен только для одного серсива, если
|
||||
> токен будет украден.
|
||||
|
||||
Доступные эндпоинты:
|
||||
- `POST /<channel>/notify`
|
||||
- **Описание.** Используется, чтобы отправить уведомление в указанный канал.
|
||||
Вместо `<channel>` указывается код канала, получаемый при помощи команды
|
||||
`!info`, выполненной в комнате Matrix.
|
||||
- **Тело запроса.** Тело запроса представляет собой `json` объект:
|
||||
```json
|
||||
- `GET /notify`
|
||||
- **Описание.** Используется, чтобы отправить уведомление в канал.
|
||||
- **Параметры запроса**
|
||||
- `channel` - код канала, получаемый при помощи команды `!info`,
|
||||
выполненной в комнате Matrix
|
||||
- `service` - имя сервиса (учитывается, только если для токена не было
|
||||
настроено имя сервиса через бота)
|
||||
- `text` - текст уведомления (форматирование не поддерживается)
|
||||
- **Ответ**
|
||||
- В случае успеха сервер вернёт `200` и JSON следующего формата:
|
||||
```json
|
||||
{
|
||||
"service": "<название сервиса; указывается, если токен это позволяет>",
|
||||
"text": "<текст уведомления>"
|
||||
"notification_id": "Notification ID will be here"
|
||||
}
|
||||
```
|
||||
- **Тело ответа.** Тело ответа представляет собой `json` объект. В случае
|
||||
успеха в нём будут все поля, перечисляемые ниже. В случае провала - только
|
||||
поле `error`, содержащее текстовое описание ошибки.
|
||||
```json
|
||||
```
|
||||
> В текущей версии `notification_id` не имеет практической пользы и его
|
||||
> формат будет меняться.
|
||||
- В случае ошибки сервер вернёт JSON следующего формата:
|
||||
```json
|
||||
{
|
||||
"error": null,
|
||||
"notification_id": "<здесь будет Notification ID>"
|
||||
"detail": "Error description in English"
|
||||
}
|
||||
```
|
||||
- ``
|
||||
```
|
||||
@@ -11,7 +11,11 @@ import util
|
||||
DEFAULT_CONFIG = {
|
||||
"matrix_homeserver": "https://matrix.domain.net",
|
||||
"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)
|
||||
return cfg
|
||||
except:
|
||||
traceback.print_exc()
|
||||
return None
|
||||
58
database.py
58
database.py
@@ -84,20 +84,33 @@ class Database:
|
||||
wait_task = asyncio.create_task(
|
||||
asyncio.sleep(self.BACKGROUND_ROUTINE_PERIOD)
|
||||
)
|
||||
if self._connection is None:
|
||||
continue
|
||||
# get current time
|
||||
current_time = time.time()
|
||||
# delete old bans
|
||||
try:
|
||||
if self._connection is not None:
|
||||
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
||||
async with self._connection.execute(statement, (time.time(),)) as cursor:
|
||||
async for row in cursor:
|
||||
ip = row["ip"]
|
||||
if ip in self._bans:
|
||||
del self._bans[ip]
|
||||
self._logger.debug(f"IP {ip} is not banned anymore")
|
||||
await self._connection.commit()
|
||||
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
||||
async with self._connection.execute(statement, (current_time,)) as cursor:
|
||||
async for row in cursor:
|
||||
ip = row["ip"]
|
||||
if ip in self._bans:
|
||||
del self._bans[ip]
|
||||
self._logger.debug(f"IP {ip} is not banned anymore")
|
||||
await self._connection.commit()
|
||||
self._logger.debug("Performed banned IPs cleanup")
|
||||
except:
|
||||
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")
|
||||
|
||||
async def _load_bans_from_database(self) -> None:
|
||||
@@ -115,6 +128,12 @@ class Database:
|
||||
self._logger.error(traceback.format_exc())
|
||||
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
|
||||
#
|
||||
@@ -127,6 +146,7 @@ class Database:
|
||||
self._background_stop_event: asyncio.Event | None = None
|
||||
self._background_task: asyncio.Task | None = None
|
||||
self._bans: dict = {}
|
||||
self._fails: dict[str, list[float]] = {}
|
||||
|
||||
async def connect(self) -> bool:
|
||||
"""Connect to the database. Returns False on failure."""
|
||||
@@ -137,6 +157,7 @@ class Database:
|
||||
self._connection.row_factory = Row
|
||||
await self._setup_tables(self._connection)
|
||||
self._bans = {}
|
||||
self._fails = {}
|
||||
await self._load_bans_from_database()
|
||||
self._background_stop_event = asyncio.Event()
|
||||
self._background_task = asyncio.create_task(
|
||||
@@ -346,4 +367,21 @@ class Database:
|
||||
if ip in self._bans:
|
||||
del self._bans[ip]
|
||||
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_user: str
|
||||
store_dir: str
|
||||
web_ip: str
|
||||
web_port: int
|
||||
fails_to_ban: int
|
||||
ban_duration: float
|
||||
|
||||
@dataclass
|
||||
class ObjectRoom:
|
||||
|
||||
5
main.py
5
main.py
@@ -14,6 +14,7 @@ import database
|
||||
import config
|
||||
import util
|
||||
import bot_callbacks
|
||||
import web
|
||||
|
||||
bot: MatrixBot = None # type: ignore
|
||||
|
||||
@@ -47,17 +48,21 @@ async def main() -> None:
|
||||
db = database.Database(Path("database.sqlite"))
|
||||
# setup the callbacks
|
||||
bot_callbacks.setup(bot, db)
|
||||
# setup the web server
|
||||
web_server = web.Web(cfg, bot, db)
|
||||
|
||||
# start the app
|
||||
if not await db.connect():
|
||||
util.log_error("Can't connect to the database!")
|
||||
return
|
||||
bot.start()
|
||||
await web_server.start()
|
||||
|
||||
# wait for Ctrl+C
|
||||
await util.get_app_stop_event().wait()
|
||||
|
||||
# stop the app
|
||||
await web_server.stop()
|
||||
await bot.stop()
|
||||
await db.disconnect()
|
||||
|
||||
|
||||
183
web.py
Normal file
183
web.py
Normal file
@@ -0,0 +1,183 @@
|
||||
"""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, ObjectToken
|
||||
from database import Database
|
||||
|
||||
class Web:
|
||||
#
|
||||
# PRIVATE
|
||||
#
|
||||
@staticmethod
|
||||
def _get_client_ip(request: Request) -> str | None:
|
||||
"""Returns IP address of the client"""
|
||||
x_forwarded_for = request.headers.get("x-forwarded-for")
|
||||
if x_forwarded_for:
|
||||
return x_forwarded_for.split(",")[0].strip()
|
||||
else:
|
||||
x_real_ip = request.headers.get("x-real-ip")
|
||||
if x_real_ip:
|
||||
return x_real_ip
|
||||
else:
|
||||
return request.client.host if request.client else None
|
||||
|
||||
def _is_banned(self, request: Request) -> bool:
|
||||
# get IP
|
||||
ip = self._get_client_ip(request)
|
||||
if ip is None:
|
||||
# treat requests without source IP as banned
|
||||
return True
|
||||
# check if banned in database
|
||||
return self._db.ban_check(ip)
|
||||
|
||||
async def _handle_failed_request(self, request: Request, possible_ban_reason: str) -> None:
|
||||
"""Adds the request to list of failed."""
|
||||
ip: str = self._get_client_ip(request) # type: ignore
|
||||
fails_total = self._db.fail_create(
|
||||
ip,
|
||||
time.time() + 60
|
||||
)
|
||||
if fails_total >= self._app_config.fails_to_ban:
|
||||
await self._db.ban_create(
|
||||
ip,
|
||||
time.time() + self._app_config.ban_duration,
|
||||
possible_ban_reason
|
||||
)
|
||||
self._db.fail_clear(ip)
|
||||
return None
|
||||
|
||||
async def _middleware_banned(self, request: Request, call_next):
|
||||
if self._is_banned(request):
|
||||
return JSONResponse(
|
||||
status_code=403,
|
||||
content={"detail": "Try again later"}
|
||||
)
|
||||
return await call_next(request)
|
||||
|
||||
async def _middleware_auth(self, request: Request, call_next):
|
||||
# ignore favicon
|
||||
if request.url.path.endswith("favicon.ico"):
|
||||
return JSONResponse(
|
||||
status_code=404,
|
||||
content={"detail": "Not found"}
|
||||
)
|
||||
# token that will be used
|
||||
token: str | None = None
|
||||
# check if token is provided via header
|
||||
auth_header = request.headers.get("Authorization")
|
||||
if auth_header and auth_header.lower().startswith("bearer"):
|
||||
parts = auth_header.split(" ")
|
||||
if len(parts) >= 2:
|
||||
token = auth_header.split(" ")[1]
|
||||
# get token from URL if not found in headers
|
||||
if not token:
|
||||
params = dict(request.query_params)
|
||||
if "token" in params:
|
||||
token = params["token"]
|
||||
# still no token
|
||||
if not token:
|
||||
await self._handle_failed_request(request, "Banned by the web server")
|
||||
return JSONResponse(
|
||||
status_code=403,
|
||||
content={"detail": "No token provided"}
|
||||
)
|
||||
# try to authorize
|
||||
auth_data = await self._db.token_get(token)
|
||||
if not auth_data:
|
||||
await self._handle_failed_request(request, "Banned by the web server")
|
||||
return JSONResponse(
|
||||
status_code=403,
|
||||
content={"detail": "Invalid token"}
|
||||
)
|
||||
# authorized
|
||||
request.state.user = auth_data
|
||||
return await call_next(request)
|
||||
|
||||
async def _cb_get_notify(self, request: Request) -> dict:
|
||||
params = dict(request.query_params)
|
||||
token: ObjectToken = request.state.user
|
||||
if "channel" not in params or len(params["channel"]) < 1:
|
||||
return { "detail": "Missing `channel`" }
|
||||
if params["channel"][0] != "q":
|
||||
return { "detail": "Invalid `channel`" }
|
||||
if "text" not in params:
|
||||
return { "detail": "Missing `text`" }
|
||||
if not token.name:
|
||||
service_name = params["service"] if "service" in params else "Unnamed"
|
||||
else:
|
||||
service_name = token.name
|
||||
# get the room
|
||||
room = await self._db.room_get(params["channel"])
|
||||
if not room:
|
||||
return { "detail": "Unknown `channel`" }
|
||||
# send the notification
|
||||
full_text = f"<strong>Уведомление от службы <code>{html.escape(service_name)}</code></strong>"
|
||||
full_text += "<br><br>"
|
||||
full_text += html.escape(params["text"])
|
||||
try:
|
||||
notification_id = await self._bot.send_text_to_room(room.matrix_id, full_text, is_html=True)
|
||||
except:
|
||||
traceback.print_exc()
|
||||
return { "detail": "Failed to send the notification" }
|
||||
return { "notification_id": notification_id }
|
||||
|
||||
|
||||
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._bot = bot
|
||||
self._db = db
|
||||
self._app_config = app_config
|
||||
|
||||
self._api = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||
self._api.add_api_route("/notify", self._cb_get_notify, methods=["GET"])
|
||||
|
||||
self._api.middleware("http")(self._middleware_auth)
|
||||
self._api.middleware("http")(self._middleware_banned)
|
||||
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
|
||||
Reference in New Issue
Block a user