Improved web server code
This commit is contained in:
134
web.py
134
web.py
@@ -11,127 +11,16 @@ from fastapi.responses import JSONResponse
|
|||||||
from mab import MatrixBot
|
from mab import MatrixBot
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
|
||||||
from datatypes import AppConfig, ObjectToken
|
from datatypes import AppConfig
|
||||||
from database import Database
|
from database import Database
|
||||||
|
|
||||||
|
import web_routes
|
||||||
|
import web_middleware
|
||||||
|
|
||||||
class Web:
|
class Web:
|
||||||
#
|
#
|
||||||
# PRIVATE
|
# 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):
|
async def _cb_exception(self, request: Request, exc: Exception):
|
||||||
traceback.print_exception(exc)
|
traceback.print_exception(exc)
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
@@ -143,16 +32,17 @@ class Web:
|
|||||||
# PUBLIC
|
# PUBLIC
|
||||||
#
|
#
|
||||||
def __init__(self, app_config: AppConfig, bot: MatrixBot, db: Database):
|
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 = 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.include_router(web_routes.router)
|
||||||
self._api.middleware("http")(self._middleware_banned)
|
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._api.add_exception_handler(Exception, self._cb_exception)
|
||||||
|
|
||||||
self._server_config = uvicorn.Config(
|
self._server_config = uvicorn.Config(
|
||||||
self._api,
|
self._api,
|
||||||
host=app_config.web_ip,
|
host=app_config.web_ip,
|
||||||
|
|||||||
135
web_middleware.py
Normal file
135
web_middleware.py
Normal 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
70
web_routes.py
Normal 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}
|
||||||
Reference in New Issue
Block a user