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)