183 lines
6.4 KiB
Python
183 lines
6.4 KiB
Python
"""This module implements API"""
|
|
|
|
import asyncio
|
|
import traceback
|
|
import time
|
|
import html
|
|
|
|
from fastapi import FastAPI, Request
|
|
from fastapi.responses import JSONResponse
|
|
|
|
from mab import MatrixBot
|
|
import uvicorn
|
|
|
|
from datatypes import AppConfig, ObjectToken
|
|
from database import Database
|
|
|
|
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(
|
|
status_code=500,
|
|
content={"detail": "Internal Server Error"}
|
|
)
|
|
|
|
#
|
|
# 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.middleware("http")(self._middleware_auth)
|
|
self._api.middleware("http")(self._middleware_banned)
|
|
self._api.add_exception_handler(Exception, self._cb_exception)
|
|
self._server_config = uvicorn.Config(
|
|
self._api,
|
|
host=app_config.web_ip,
|
|
port=app_config.web_port
|
|
)
|
|
self._server: uvicorn.Server | None = None
|
|
self._server_task: asyncio.Task | None = None
|
|
|
|
async def start(self) -> None:
|
|
"""Start the API."""
|
|
if self._server is not None or self._server_task is not None:
|
|
raise RuntimeError("The server is already started")
|
|
self._server = uvicorn.Server(self._server_config)
|
|
self._server_task = asyncio.create_task(self._server.serve())
|
|
|
|
async def stop(self) -> None:
|
|
"""Stop the API server."""
|
|
if self._server is None or self._server_task is None:
|
|
raise RuntimeError("The server is not started yet")
|
|
self._server_task.cancel()
|
|
try:
|
|
await self._server_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
except:
|
|
traceback.print_exc()
|
|
self._server_task = None
|
|
self._server = None |