"""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"Уведомление от службы {html.escape(service_name)}" full_text += "

" 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