Files
govoplan-core/src/govoplan_core/server/fastapi.py

185 lines
7.3 KiB
Python

from __future__ import annotations
import logging
import os
import time
from collections.abc import AsyncIterator, Callable, Iterable
from contextlib import AbstractAsyncContextManager
from typing import Any
from fastapi import APIRouter, FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from starlette.middleware.trustedhost import TrustedHostMiddleware
from govoplan_core.core.events import event_context, new_event_id, normalize_trace_id
from govoplan_core.core.install_config import validate_runtime_configuration
from govoplan_core.core.registry import PlatformRegistry
from govoplan_core.db.query_metrics import collect_query_metrics
from govoplan_core.server.conditional_requests import conditional_json_get_middleware
from govoplan_core.server.request_limits import RequestBodyLimitMiddleware
LifespanFactory = Callable[[FastAPI], AbstractAsyncContextManager[None] | AsyncIterator[None]]
logger = logging.getLogger("govoplan.request")
_CONTENT_SECURITY_POLICY = "base-uri 'self'; object-src 'none'; frame-ancestors 'none'"
_PRODUCTION_LIKE_ENVIRONMENTS = frozenset(
{"prod", "production", "self-hosted", "staging", "production-like", "production-like-dev"}
)
def _slow_request_threshold_ms() -> float:
raw = os.getenv("GOVOPLAN_SLOW_REQUEST_MS", "500").strip()
try:
return float(raw)
except ValueError:
return 500.0
def _hsts_seconds() -> int:
default = "31536000" if os.getenv("APP_ENV", "dev").strip().lower() in {"prod", "production"} else "0"
raw = os.getenv("GOVOPLAN_HTTP_HSTS_SECONDS", default).strip()
try:
return max(0, int(raw))
except ValueError:
return int(default)
def _max_request_body_bytes() -> int:
default = 512 * 1024 * 1024
raw = os.getenv("GOVOPLAN_HTTP_MAX_REQUEST_BODY_BYTES", str(default)).strip()
try:
value = int(raw)
except ValueError:
return default
return value if value > 0 else default
def _trusted_hosts() -> tuple[str, ...]:
return tuple(item.strip() for item in os.getenv("GOVOPLAN_TRUSTED_HOSTS", "").split(",") if item.strip())
def _validate_production_startup() -> None:
app_env = os.getenv("APP_ENV", "").strip().lower().replace("_", "-")
install_profile = os.getenv("GOVOPLAN_INSTALL_PROFILE", "").strip().lower().replace("_", "-")
if app_env not in _PRODUCTION_LIKE_ENVIRONMENTS and install_profile not in _PRODUCTION_LIKE_ENVIRONMENTS:
return
validation = validate_runtime_configuration()
if validation.errors:
raise RuntimeError(validation.to_text())
throttle_enabled = os.getenv("AUTH_LOGIN_THROTTLE_ENABLED", "true").strip().lower() not in {
"0",
"false",
"no",
"off",
}
if throttle_enabled and not os.getenv("REDIS_URL", "").strip():
logger.warning(
"Redis is not configured; login throttling is process-local and will not coordinate horizontally"
)
def create_govoplan_app(
*,
title: str,
version: str,
registry: PlatformRegistry,
api_router: APIRouter | None = None,
lifespan: LifespanFactory | None = None,
cors_origins: Iterable[str] = (),
health_payload: dict[str, Any] | None = None,
) -> FastAPI:
_validate_production_startup()
app = FastAPI(title=title, version=version, lifespan=lifespan)
app.add_middleware(RequestBodyLimitMiddleware, max_bytes=_max_request_body_bytes())
trusted_hosts = _trusted_hosts()
if trusted_hosts:
app.add_middleware(TrustedHostMiddleware, allowed_hosts=list(trusted_hosts))
app.state.govoplan_registry = registry
slow_request_threshold_ms = _slow_request_threshold_ms()
hsts_seconds = _hsts_seconds()
@app.middleware("http")
async def security_response_headers(request: Request, call_next):
response = await call_next(request)
response.headers.setdefault("X-Content-Type-Options", "nosniff")
response.headers.setdefault("X-Frame-Options", "DENY")
response.headers.setdefault("Referrer-Policy", "strict-origin-when-cross-origin")
response.headers.setdefault("Permissions-Policy", "camera=(), microphone=(), geolocation=()")
response.headers.setdefault("Content-Security-Policy", _CONTENT_SECURITY_POLICY)
if hsts_seconds and request.url.scheme == "https":
response.headers.setdefault("Strict-Transport-Security", f"max-age={hsts_seconds}")
return response
@app.middleware("http")
async def request_correlation_context(request: Request, call_next):
correlation_id = (
normalize_trace_id(request.headers.get("x-correlation-id"))
or normalize_trace_id(request.headers.get("x-request-id"))
or new_event_id()
)
with event_context(correlation_id=correlation_id):
response = await call_next(request)
response.headers["X-Correlation-ID"] = correlation_id
return response
@app.middleware("http")
async def slow_request_logging(request: Request, call_next):
started_at = time.perf_counter()
with collect_query_metrics() as db_metrics:
try:
response = await call_next(request)
except Exception:
elapsed_ms = (time.perf_counter() - started_at) * 1000
if slow_request_threshold_ms > 0 and elapsed_ms >= slow_request_threshold_ms:
logger.warning(
"slow request failed method=%s path=%s duration_ms=%.1f db_query_count=%s db_time_ms=%.1f db_slowest_ms=%.1f db_error_count=%s",
request.method,
request.url.path,
elapsed_ms,
db_metrics.query_count,
db_metrics.total_ms,
db_metrics.slowest_ms,
db_metrics.error_count,
exc_info=True,
)
raise
elapsed_ms = (time.perf_counter() - started_at) * 1000
if slow_request_threshold_ms > 0 and elapsed_ms >= slow_request_threshold_ms:
logger.warning(
"slow request method=%s path=%s status=%s duration_ms=%.1f db_query_count=%s db_time_ms=%.1f db_slowest_ms=%.1f db_error_count=%s",
request.method,
request.url.path,
response.status_code,
elapsed_ms,
db_metrics.query_count,
db_metrics.total_ms,
db_metrics.slowest_ms,
db_metrics.error_count,
)
return response
app.middleware("http")(conditional_json_get_middleware)
origins = [item.strip() for item in cors_origins if item.strip()]
if origins:
if "*" in origins:
raise ValueError("CORS wildcard origin '*' cannot be used with credentialed requests; configure explicit origins.")
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
if api_router is not None:
app.include_router(api_router)
@app.get("/health")
def health():
payload = {"status": "ok"}
if health_payload:
payload.update(health_payload)
return payload
return app