mirror of
https://github.com/vladmandic/automatic
synced 2026-08-25 22:20:46 +02:00
a882ce945b
Signed-off-by: Vladimir Mandic <mandic00@live.com>
183 lines
8.0 KiB
Python
183 lines
8.0 KiB
Python
import ssl
|
|
import time
|
|
import logging
|
|
import asyncio
|
|
from asyncio.exceptions import CancelledError
|
|
import anyio
|
|
import starlette
|
|
import uvicorn
|
|
import fastapi
|
|
from starlette.responses import JSONResponse
|
|
from starlette.websockets import WebSocket, WebSocketDisconnect
|
|
from fastapi import FastAPI, Request, Response
|
|
from fastapi.exceptions import HTTPException
|
|
from fastapi.responses import RedirectResponse
|
|
from fastapi.encoders import jsonable_encoder
|
|
from modules.logger import log
|
|
import modules.errors as errors
|
|
from modules.api.validate import validate_request, validate_log
|
|
|
|
errors.install()
|
|
|
|
|
|
def validate_subpath(endpoint: str, subpath: str | None):
|
|
if (subpath is not None) and (len(subpath) > 0) and (not endpoint.startswith(subpath)):
|
|
if not subpath.endswith('/'):
|
|
subpath = f'{subpath}/'
|
|
if endpoint.startswith('/'):
|
|
url = f'{subpath}{endpoint[1:]}'
|
|
else:
|
|
url = f'{subpath}{endpoint}'
|
|
log.trace(f'API: redirect subpath={subpath} url="{endpoint}" redirect="{url}"')
|
|
return RedirectResponse(url=url, status_code=308)
|
|
return None
|
|
|
|
class LoopInstrumentorMiddleware:
|
|
def __init__(self, app):
|
|
self.app = app
|
|
self.instrumented = False
|
|
|
|
async def __call__(self, scope, receive, send):
|
|
if not self.instrumented:
|
|
loop = asyncio.get_running_loop()
|
|
def verbose_task_factory(loop, coro, context=None):
|
|
coro_name = getattr(coro, '__qualname__', str(coro))
|
|
frame = getattr(coro, 'cr_frame', None)
|
|
origin = f"{frame.f_code.co_filename}:{frame.f_lineno}" if frame else "unknown"
|
|
log.trace(f"HTTP: coro={coro_name} fn={origin}")
|
|
if context is not None:
|
|
return asyncio.Task(coro, loop=loop, name=coro_name, context=context)
|
|
return asyncio.Task(coro, loop=loop, name=coro_name)
|
|
|
|
loop.set_task_factory(verbose_task_factory)
|
|
self.instrumented = True
|
|
|
|
await self.app(scope, receive, send)
|
|
|
|
|
|
def setup_logging(debug: bool = False):
|
|
level = logging.DEBUG if debug else logging.WARNING
|
|
logging.getLogger("httpcore").setLevel(level)
|
|
logging.getLogger("httpx").setLevel(level)
|
|
logging.getLogger("uvicorn.access").setLevel(level)
|
|
logging.getLogger("asyncio").setLevel(level)
|
|
if not debug:
|
|
logging.getLogger("uvicorn.error").disabled = True
|
|
if debug:
|
|
asyncio_logger = logging.getLogger("asyncio")
|
|
if not asyncio_logger.handlers:
|
|
handler = logging.StreamHandler()
|
|
handler.setFormatter(logging.Formatter("[asyncio] %(message)s"))
|
|
asyncio_logger.addHandler(handler)
|
|
|
|
|
|
def setup_middleware(app: FastAPI, cmd_opts):
|
|
ssl._create_default_https_context = ssl._create_unverified_context # pylint: disable=protected-access
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from fastapi.middleware.gzip import GZipMiddleware
|
|
app.user_middleware = [x for x in app.user_middleware if x.cls.__name__ != 'CORSMiddleware']
|
|
app.middleware_stack = None # reset current middleware to allow modifying user provided list
|
|
app.add_middleware(GZipMiddleware, minimum_size=2048)
|
|
if cmd_opts.profile:
|
|
app.add_middleware(LoopInstrumentorMiddleware)
|
|
if cmd_opts.cors_origins and cmd_opts.cors_regex:
|
|
app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*'])
|
|
elif cmd_opts.cors_origins:
|
|
app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_methods=['*'], allow_credentials=True, allow_headers=['*'])
|
|
elif cmd_opts.cors_regex:
|
|
app.add_middleware(CORSMiddleware, allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*'])
|
|
|
|
@app.middleware("http")
|
|
async def api_preprocess(req: Request, call_next):
|
|
try:
|
|
ts = time.time()
|
|
res: Response = await call_next(req)
|
|
duration = str(round(time.time() - ts, 4))
|
|
res.headers["X-Process-Time"] = duration
|
|
endpoint = req.scope.get('path', 'err')
|
|
client = req.scope.get('client', ('0:0.0.0', 0))[0]
|
|
token = req.cookies.get("access-token") or req.cookies.get("access-token-unsecure")
|
|
|
|
redirect = validate_subpath(endpoint, cmd_opts.subpath)
|
|
if redirect:
|
|
return redirect
|
|
|
|
validate_request(client, endpoint)
|
|
if cmd_opts.api_log:
|
|
if not validate_log(client, endpoint):
|
|
return res
|
|
log.info('API user={user} code={code} {prot}/{ver} {method} {endpoint} {client} {duration}'.format( # pylint: disable=consider-using-f-string, logging-format-interpolation
|
|
user = app.tokens.get(token) if hasattr(app, 'tokens') else None,
|
|
code = res.status_code,
|
|
ver = req.scope.get('http_version', '0.0'),
|
|
client = client,
|
|
prot = req.scope.get('scheme', 'err'),
|
|
method = req.scope.get('method', 'err'),
|
|
endpoint = endpoint,
|
|
duration = duration,
|
|
))
|
|
return res
|
|
except CancelledError:
|
|
log.warning('WebSocket closed')
|
|
except BaseException as e:
|
|
return handle_exception(req, e)
|
|
|
|
def handle_exception(req: Request, e: Exception):
|
|
err = {
|
|
"error": type(e).__name__,
|
|
"code": vars(e).get('status_code', 500),
|
|
"detail": vars(e).get('detail', ''),
|
|
"body": vars(e).get('body', ''),
|
|
"errors": str(e),
|
|
}
|
|
if err['code'] == 401 and 'file=' in req.url.path: # dont spam with unauth
|
|
return JSONResponse(status_code=err['code'], content=jsonable_encoder(err))
|
|
if err['code'] == 404 and 'file=ui/' in req.url.path: # dont spam with locales
|
|
return JSONResponse(status_code=err['code'], content=jsonable_encoder(err))
|
|
if err["code"] == 429: # dont spam with rate limit errors
|
|
return JSONResponse(status_code=err["code"], content=jsonable_encoder(err))
|
|
|
|
endpoint = req.scope.get("path", "err")
|
|
client = req.scope.get("client", ("0:0.0.0", 0))[0]
|
|
if not validate_log(client, endpoint):
|
|
log.error(f"API error: {req.method}: {req.url} {err}")
|
|
|
|
if not isinstance(e, HTTPException) and err['error'] != 'TypeError': # do not print backtrace on known httpexceptions
|
|
errors.display(e, 'HTTP API', [anyio, fastapi, uvicorn, starlette])
|
|
elif err['code'] in [404, 401, 400, 403]:
|
|
pass
|
|
else:
|
|
log.debug(e, exc_info=True) # print stack trace
|
|
return JSONResponse(status_code=err['code'], content=jsonable_encoder(err))
|
|
|
|
@app.exception_handler(HTTPException)
|
|
async def http_exception_handler(req: Request, e: HTTPException):
|
|
return handle_exception(req, e)
|
|
|
|
@app.exception_handler(Exception)
|
|
async def general_exception_handler(req: Request, e: Exception):
|
|
if isinstance(e, TypeError):
|
|
return JSONResponse(status_code=500, content=jsonable_encoder(str(e)))
|
|
else:
|
|
return handle_exception(req, e)
|
|
|
|
app.build_middleware_stack() # rebuild middleware stack on-the-fly
|
|
log.debug(f'API middleware: {[m.cls.__name__ for m in app.user_middleware]}')
|
|
|
|
@app.websocket("/internal/monitor")
|
|
async def ws_monitor(ws: WebSocket):
|
|
await ws.accept()
|
|
try:
|
|
while True:
|
|
await asyncio.sleep(1.0)
|
|
await ws.send_json({"status": "ok"})
|
|
except WebSocketDisconnect:
|
|
pass # Expected when client navigates away or closes tab
|
|
except Exception as e:
|
|
log.error(f'WebSocket monitor: {e}')
|
|
finally:
|
|
try:
|
|
await ws.close()
|
|
except RuntimeError:
|
|
pass # Socket was already closed by client
|