diff --git a/web.py b/web.py index 54d91d5..b29b030 100644 --- a/web.py +++ b/web.py @@ -11,127 +11,16 @@ from fastapi.responses import JSONResponse from mab import MatrixBot import uvicorn -from datatypes import AppConfig, ObjectToken +from datatypes import AppConfig from database import Database +import web_routes +import web_middleware + 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( @@ -143,16 +32,17 @@ class Web: # 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.state.database = db + self._api.state.matrix = bot + self._api.state.app_config = app_config - self._api.middleware("http")(self._middleware_auth) - self._api.middleware("http")(self._middleware_banned) + 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, diff --git a/web_middleware.py b/web_middleware.py new file mode 100644 index 0000000..f980ea0 --- /dev/null +++ b/web_middleware.py @@ -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) \ No newline at end of file diff --git a/web_routes.py b/web_routes.py new file mode 100644 index 0000000..e9f9edd --- /dev/null +++ b/web_routes.py @@ -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"Уведомление от службы {html.escape(service)}" + text += "
" * 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} \ No newline at end of file