- Updated mab to v0.5.0 - Added `!short` command to add named tokens - Fixed `last_access_at` field not being updated
136 lines
5.3 KiB
Python
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) |