Files
2026-matrix-csonac/web_middleware.py
nikita 4194f3da48 mab update, !short command, last_access fix
- Updated mab to v0.5.0
- Added `!short` command to add named tokens
- Fixed `last_access_at` field not being updated
2026-09-13 00:12:00 +03:00

136 lines
5.3 KiB
Python

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"}
)
await database.token_set_last_access(token, time.time())
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)