chore: consolidate platform split checks
This commit is contained in:
@@ -9,6 +9,7 @@ from govoplan_core.server.config import GovoplanServerConfig, load_server_config
|
||||
from govoplan_core.server.fastapi import create_govoplan_app
|
||||
from govoplan_core.server.platform import create_platform_router
|
||||
from govoplan_core.server.registry import available_module_manifests, build_platform_registry
|
||||
from govoplan_core.server.route_validation import validate_no_route_collisions
|
||||
|
||||
|
||||
def create_app(config: GovoplanServerConfig | str | None = None):
|
||||
@@ -19,7 +20,8 @@ def create_app(config: GovoplanServerConfig | str | None = None):
|
||||
|
||||
database_url = getattr(server_config.settings, "database_url", None) if server_config.settings is not None else None
|
||||
if database_url:
|
||||
configure_database(str(database_url))
|
||||
dispose_previous = getattr(server_config.settings, "app_env", None) == "test"
|
||||
configure_database(str(database_url), dispose_previous=dispose_previous)
|
||||
|
||||
raw_enabled_modules = load_startup_enabled_modules(server_config.enabled_modules)
|
||||
candidate_modules = startup_candidate_module_ids(server_config.enabled_modules, raw_enabled_modules)
|
||||
@@ -55,6 +57,8 @@ def create_app(config: GovoplanServerConfig | str | None = None):
|
||||
if contribution.should_include(server_config.settings, registry):
|
||||
api_router.include_router(contribution.router)
|
||||
|
||||
validate_no_route_collisions(api_router, owner="server startup routes")
|
||||
|
||||
app = create_govoplan_app(
|
||||
title=server_config.title,
|
||||
version=server_config.version,
|
||||
|
||||
123
src/govoplan_core/server/conditional_requests.py
Normal file
123
src/govoplan_core/server/conditional_requests.py
Normal file
@@ -0,0 +1,123 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from fastapi import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
JSON_CACHE_CONTROL = "private, no-cache"
|
||||
JSON_ETAG_VARY_HEADERS = ("Authorization", "Cookie", "X-API-Key", "Accept-Language")
|
||||
|
||||
|
||||
async def conditional_json_get_middleware(
|
||||
request: Request,
|
||||
call_next: Callable[[Request], Awaitable[Response]],
|
||||
) -> Response:
|
||||
"""Attach ETags to JSON GET responses and return 304 when unchanged.
|
||||
|
||||
The middleware deliberately works after route handling. That keeps the
|
||||
contract platform-wide without requiring every module router to learn about
|
||||
conditional requests, while still limiting buffering to successful JSON GET
|
||||
responses.
|
||||
"""
|
||||
|
||||
response = await call_next(request)
|
||||
if not _eligible_for_conditional_json_get(request, response):
|
||||
return response
|
||||
|
||||
body = b"".join([chunk async for chunk in response.body_iterator])
|
||||
etag = response.headers.get("etag") or json_response_etag(body)
|
||||
headers = dict(response.headers)
|
||||
headers["etag"] = etag
|
||||
headers["cache-control"] = _conditional_cache_control(headers.get("cache-control"))
|
||||
headers["vary"] = _merge_vary(headers.get("vary"), JSON_ETAG_VARY_HEADERS)
|
||||
headers.pop("content-length", None)
|
||||
|
||||
if if_none_match_matches(request.headers.get("if-none-match"), etag):
|
||||
return Response(status_code=304, headers=_not_modified_headers(headers), background=response.background)
|
||||
|
||||
return Response(content=body, status_code=response.status_code, headers=headers, background=response.background)
|
||||
|
||||
|
||||
def json_response_etag(body: bytes) -> str:
|
||||
digest = hashlib.sha256(body).hexdigest()
|
||||
return f'W/"sha256-{digest}"'
|
||||
|
||||
|
||||
def if_none_match_matches(header_value: str | None, etag: str) -> bool:
|
||||
if not header_value:
|
||||
return False
|
||||
normalized_etag = _normalize_etag_for_weak_compare(etag)
|
||||
for raw_candidate in header_value.split(","):
|
||||
candidate = raw_candidate.strip()
|
||||
if candidate == "*":
|
||||
return True
|
||||
if _normalize_etag_for_weak_compare(candidate) == normalized_etag:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _eligible_for_conditional_json_get(request: Request, response: Response) -> bool:
|
||||
if request.method.upper() != "GET":
|
||||
return False
|
||||
if response.status_code != 200:
|
||||
return False
|
||||
if "set-cookie" in response.headers:
|
||||
return False
|
||||
if "content-encoding" in response.headers:
|
||||
return False
|
||||
if "content-disposition" in response.headers:
|
||||
return False
|
||||
if "no-store" in request.headers.get("cache-control", "").lower():
|
||||
return False
|
||||
if "no-store" in response.headers.get("cache-control", "").lower():
|
||||
return False
|
||||
return _is_json_content_type(response.headers.get("content-type"))
|
||||
|
||||
|
||||
def _is_json_content_type(value: str | None) -> bool:
|
||||
media_type = (value or "").split(";", 1)[0].strip().lower()
|
||||
return media_type == "application/json" or media_type.endswith("+json")
|
||||
|
||||
|
||||
def _conditional_cache_control(current: str | None) -> str:
|
||||
if not current:
|
||||
return JSON_CACHE_CONTROL
|
||||
directives = [item.strip() for item in current.split(",") if item.strip()]
|
||||
normalized = {item.split("=", 1)[0].strip().lower() for item in directives}
|
||||
if "private" not in normalized and "public" not in normalized:
|
||||
directives.append("private")
|
||||
if "no-cache" not in normalized:
|
||||
directives.append("no-cache")
|
||||
return ", ".join(directives)
|
||||
|
||||
|
||||
def _merge_vary(current: str | None, additions: tuple[str, ...]) -> str:
|
||||
tokens: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for value in [current or "", *additions]:
|
||||
for token in value.split(","):
|
||||
stripped = token.strip()
|
||||
if not stripped:
|
||||
continue
|
||||
normalized = stripped.lower()
|
||||
if normalized in seen:
|
||||
continue
|
||||
seen.add(normalized)
|
||||
tokens.append(stripped)
|
||||
return ", ".join(tokens)
|
||||
|
||||
|
||||
def _not_modified_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||
allowed = {"cache-control", "content-location", "date", "etag", "expires", "vary", "x-correlation-id"}
|
||||
return {key: value for key, value in headers.items() if key.lower() in allowed}
|
||||
|
||||
|
||||
def _normalize_etag_for_weak_compare(value: str) -> str:
|
||||
normalized = value.strip()
|
||||
if normalized.startswith("W/"):
|
||||
normalized = normalized[2:].strip()
|
||||
if len(normalized) >= 2 and normalized[0] == '"' and normalized[-1] == '"':
|
||||
normalized = normalized[1:-1]
|
||||
return normalized
|
||||
@@ -3,8 +3,9 @@ from __future__ import annotations
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import Depends, FastAPI
|
||||
from sqlalchemy.engine import make_url
|
||||
|
||||
from govoplan_access.backend.auth.dependencies import ApiPrincipal, require_scope
|
||||
from govoplan_access.auth import ApiPrincipal, require_scope
|
||||
from govoplan_core.core.registry import PlatformRegistry
|
||||
from govoplan_core.db.bootstrap import bootstrap_dev_data, create_all_tables
|
||||
from govoplan_core.db.session import get_database
|
||||
@@ -12,6 +13,13 @@ from govoplan_core.server.config import GovoplanServerConfig
|
||||
from govoplan_core.settings import Settings, settings
|
||||
|
||||
|
||||
def _dev_bootstrap_needs_create_all(database_url: str) -> bool:
|
||||
try:
|
||||
return make_url(database_url).get_backend_name() == "sqlite"
|
||||
except Exception:
|
||||
return database_url.startswith("sqlite:")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
if settings.app_env.lower() == "dev" and settings.dev_auto_migrate_enabled:
|
||||
@@ -19,7 +27,8 @@ async def lifespan(app: FastAPI):
|
||||
|
||||
migrate_database(database_url=settings.database_url)
|
||||
if settings.app_env.lower() == "dev" and settings.dev_bootstrap_enabled:
|
||||
create_all_tables()
|
||||
if _dev_bootstrap_needs_create_all(settings.database_url):
|
||||
create_all_tables()
|
||||
with get_database().SessionLocal() as session:
|
||||
bootstrap_dev_data(
|
||||
session,
|
||||
|
||||
@@ -4,10 +4,12 @@ from collections.abc import AsyncIterator, Callable, Iterable
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, FastAPI
|
||||
from fastapi import APIRouter, FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from govoplan_core.core.events import event_context, new_event_id, normalize_trace_id
|
||||
from govoplan_core.core.registry import PlatformRegistry
|
||||
from govoplan_core.server.conditional_requests import conditional_json_get_middleware
|
||||
|
||||
LifespanFactory = Callable[[FastAPI], AbstractAsyncContextManager[None] | AsyncIterator[None]]
|
||||
|
||||
@@ -25,6 +27,20 @@ def create_govoplan_app(
|
||||
app = FastAPI(title=title, version=version, lifespan=lifespan)
|
||||
app.state.govoplan_registry = registry
|
||||
|
||||
@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")(conditional_json_get_middleware)
|
||||
|
||||
origins = [item.strip() for item in cors_origins if item.strip()]
|
||||
if origins:
|
||||
app.add_middleware(
|
||||
|
||||
@@ -3,10 +3,13 @@ from __future__ import annotations
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from govoplan_core.admin.models import SystemSettings
|
||||
from govoplan_core.admin.settings import SYSTEM_SETTINGS_ID
|
||||
from govoplan_core.core.maintenance import saved_maintenance_mode
|
||||
from govoplan_core.core.modules import FrontendModule, FrontendRoute, NavItem
|
||||
from govoplan_core.core.registry import PlatformRegistry
|
||||
from govoplan_core.db.session import get_database
|
||||
from govoplan_core.i18n import system_i18n_payload
|
||||
|
||||
|
||||
def _registry(request: Request) -> PlatformRegistry:
|
||||
@@ -72,10 +75,13 @@ def create_platform_router(settings: object | None = None) -> APIRouter:
|
||||
try:
|
||||
with get_database().session() as session:
|
||||
maintenance_mode = saved_maintenance_mode(session)
|
||||
settings_item = session.get(SystemSettings, SYSTEM_SETTINGS_ID)
|
||||
except (RuntimeError, SQLAlchemyError):
|
||||
maintenance_mode = None
|
||||
settings_item = None
|
||||
return {
|
||||
"maintenance_mode": maintenance_mode.as_dict() if maintenance_mode is not None else {"enabled": False, "message": None},
|
||||
"i18n": system_i18n_payload(settings_item),
|
||||
}
|
||||
|
||||
@router.get("/modules")
|
||||
|
||||
73
src/govoplan_core/server/route_validation.py
Normal file
73
src/govoplan_core/server/route_validation.py
Normal file
@@ -0,0 +1,73 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class RouteCollisionError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteSignature:
|
||||
method: str
|
||||
path: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteRecord:
|
||||
signature: RouteSignature
|
||||
owner: str
|
||||
|
||||
|
||||
def validate_no_route_collisions(routes_or_app: object, *, owner: str = "router") -> None:
|
||||
seen: dict[RouteSignature, RouteRecord] = {}
|
||||
for record in iter_route_records(routes_or_app, owner=owner):
|
||||
previous = seen.get(record.signature)
|
||||
if previous is not None:
|
||||
raise RouteCollisionError(_collision_message(record.signature, previous.owner, record.owner))
|
||||
seen[record.signature] = record
|
||||
|
||||
|
||||
def validate_router_can_mount(existing_routes_or_app: object, candidate_router: object, *, prefix: str = "", owner: str) -> None:
|
||||
validate_no_route_collisions(candidate_router, owner=owner)
|
||||
existing = {record.signature: record for record in iter_route_records(existing_routes_or_app, owner="existing application")}
|
||||
for record in iter_route_records(candidate_router, prefix=prefix, owner=owner):
|
||||
previous = existing.get(record.signature)
|
||||
if previous is not None:
|
||||
raise RouteCollisionError(_collision_message(record.signature, previous.owner, record.owner))
|
||||
|
||||
|
||||
def iter_route_records(routes_or_app: object, *, prefix: str = "", owner: str = "router") -> Iterable[RouteRecord]:
|
||||
routes = getattr(routes_or_app, "routes", routes_or_app)
|
||||
for route in routes or ():
|
||||
path = _join_route_path(prefix, str(getattr(route, "path", "") or ""))
|
||||
include_context = getattr(route, "include_context", None)
|
||||
nested_prefix = _join_route_path(prefix, str(getattr(include_context, "prefix", "") or ""))
|
||||
nested_router = getattr(route, "original_router", None)
|
||||
if nested_router is not None:
|
||||
yield from iter_route_records(nested_router, prefix=nested_prefix, owner=owner)
|
||||
|
||||
methods = getattr(route, "methods", None)
|
||||
if methods:
|
||||
for method in sorted(str(item).upper() for item in methods):
|
||||
yield RouteRecord(RouteSignature(method=method, path=path), owner)
|
||||
|
||||
nested_routes = getattr(route, "routes", None)
|
||||
if nested_routes is not None and nested_routes is not routes:
|
||||
yield from iter_route_records(nested_routes, prefix=path, owner=owner)
|
||||
|
||||
|
||||
def _join_route_path(prefix: str, path: str) -> str:
|
||||
if not prefix:
|
||||
return path or "/"
|
||||
if not path:
|
||||
return prefix or "/"
|
||||
return f"{prefix.rstrip('/')}/{path.lstrip('/')}" or "/"
|
||||
|
||||
|
||||
def _collision_message(signature: RouteSignature, previous_owner: str, next_owner: str) -> str:
|
||||
return (
|
||||
f"Route collision: {signature.method} {signature.path} is registered by "
|
||||
f"{previous_owner} and {next_owner}"
|
||||
)
|
||||
Reference in New Issue
Block a user