Improved web server code

This commit is contained in:
Nikita Tyukalov, ASUS, Linux
2026-08-29 19:24:33 +03:00
parent 8094aaab10
commit a803b80d59
3 changed files with 217 additions and 122 deletions

134
web.py
View File

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

135
web_middleware.py Normal file
View File

@@ -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)

70
web_routes.py Normal file
View File

@@ -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"<strong>Уведомление от службы <code>{html.escape(service)}</code></strong>"
text += "<br>" * 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}