diff --git a/web.py b/web.py
index 54d91d5..b29b030 100644
--- a/web.py
+++ b/web.py
@@ -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"Уведомление от службы {html.escape(service_name)}"
- full_text += "
"
- 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,
diff --git a/web_middleware.py b/web_middleware.py
new file mode 100644
index 0000000..f980ea0
--- /dev/null
+++ b/web_middleware.py
@@ -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)
\ No newline at end of file
diff --git a/web_routes.py b/web_routes.py
new file mode 100644
index 0000000..e9f9edd
--- /dev/null
+++ b/web_routes.py
@@ -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"Уведомление от службы {html.escape(service)}"
+ text += "
" * 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}
\ No newline at end of file