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