from __future__ import annotations import inspect import re from collections.abc import Iterable from copy import copy from functools import wraps from typing import Any, Generic, TypeVar, get_args, get_origin from fastapi import APIRouter, Depends, Query from fastapi.routing import APIRoute from pydantic import BaseModel, JsonValue, create_model from app.api.problem_details import ProblemDetails from app.api.v1.router import api_router as handler_api_router from app.auth.metadata_dependencies import get_current_metadata_user from app.auth.project_dependencies import ProjectContext, get_project_context T = TypeVar("T") class Page(BaseModel, Generic[T]): items: list[T] total: int limit: int offset: int _NAME_IS_NETWORK = { "pressure_sensor_placement_sensitivity_endpoint", "pressure_sensor_placement_kmeans_endpoint", } _DERIVE_USERNAME = { "pressure_sensor_placement_sensitivity_endpoint": "username", "pressure_sensor_placement_kmeans_endpoint": "username", "fastapi_pressure_sensor_placement": "user_name", } _PUBLIC_PARAMETER_RENAMES = { "burst_ID": "burst_id", "drainage_node_ID": "drainage_node_id", } _MODEL_NAME_IS_NETWORK = {"RunSimulationManuallyByDate"} _MODEL_USERNAME_FROM_AUTH: set[str] = set() def _clean_name(name: str) -> str: for prefix in ("fastapi_", "fast_"): if name.startswith(prefix): name = name[len(prefix) :] break if name.endswith("_endpoint"): name = name[: -len("_endpoint")] return name def _rest_body_model(annotation): if not inspect.isclass(annotation) or not issubclass(annotation, BaseModel): return None project_fields = { name for name in ("network", "network_name") if name in annotation.model_fields } if annotation.__name__ in _MODEL_NAME_IS_NETWORK and "name" in annotation.model_fields: project_fields.add("name") username_fields = ( { name for name in ("username", "user_name") if name in annotation.model_fields } if annotation.__name__ in _MODEL_USERNAME_FROM_AUTH else set() ) excluded_fields = project_fields | username_fields if not excluded_fields: return None public_fields = { name: (field.annotation, copy(field)) for name, field in annotation.model_fields.items() if name not in excluded_fields } public_model = create_model( f"{annotation.__name__}Rest", __module__=annotation.__module__, **public_fields, ) return annotation, public_model, project_fields, username_fields def _with_header_project_context(endpoint, route_name: str): signature = inspect.signature(endpoint) network_parameters = [ name for name in ("network", "network_name") if name in signature.parameters ] if route_name in _NAME_IS_NETWORK and "name" in signature.parameters: network_parameters.append("name") username_parameter = _DERIVE_USERNAME.get(route_name) parameter_renames = { internal: public for internal, public in _PUBLIC_PARAMETER_RENAMES.items() if internal in signature.parameters } body_models = { name: body_model for name, parameter in signature.parameters.items() if (body_model := _rest_body_model(parameter.annotation)) is not None } model_has_username = any(model[3] for model in body_models.values()) if ( not network_parameters and not username_parameter and not parameter_renames and not body_models ): return endpoint existing_context_parameter = next( ( name for name, parameter in signature.parameters.items() if parameter.annotation is ProjectContext ), None, ) injected_context_name = existing_context_parameter or "_rest_project_context" injected_user_name = "_rest_current_user" @wraps(endpoint) async def wrapper(*args, **kwargs): project_context = kwargs.get(injected_context_name) if not isinstance(project_context, ProjectContext): raise RuntimeError("REST project context was not resolved") if not existing_context_parameter: kwargs.pop(injected_context_name, None) for parameter_name in network_parameters: kwargs[parameter_name] = project_context.project_code if username_parameter: kwargs[username_parameter] = kwargs[injected_user_name].username kwargs.pop(injected_user_name, None) for internal_name, public_name in parameter_renames.items(): kwargs[internal_name] = kwargs.pop(public_name) for parameter_name, ( original_model, _public_model, project_fields, username_fields, ) in body_models.items(): data = kwargs[parameter_name].model_dump() data.update( {field_name: project_context.project_code for field_name in project_fields} ) if username_fields: current_user = kwargs[injected_user_name] data.update( {field_name: current_user.username for field_name in username_fields} ) kwargs[parameter_name] = original_model.model_validate(data) if model_has_username: kwargs.pop(injected_user_name, None) result = endpoint(*args, **kwargs) if inspect.isawaitable(result): return await result return result parameters = [] for name, parameter in signature.parameters.items(): if name in network_parameters or name == username_parameter: continue public_name = parameter_renames.get(name, name) if public_name != name: default = copy(parameter.default) default.alias = public_name default.validation_alias = public_name default.serialization_alias = public_name parameter = parameter.replace(name=public_name, default=default) if name in body_models: parameter = parameter.replace(annotation=body_models[name][1]) parameters.append(parameter) if not existing_context_parameter: parameters.append( inspect.Parameter( injected_context_name, kind=inspect.Parameter.KEYWORD_ONLY, annotation=ProjectContext, default=Depends(get_project_context), ) ) if username_parameter or model_has_username: parameters.append( inspect.Parameter( injected_user_name, kind=inspect.Parameter.KEYWORD_ONLY, default=Depends(get_current_metadata_user), ) ) wrapper.__signature__ = signature.replace(parameters=parameters) return wrapper def _with_pagination(endpoint): signature = inspect.signature(endpoint) handler_limit_parameter = "limit" if "limit" in signature.parameters else None handler_offset_parameter = next( ( parameter_name for parameter_name in ("offset", "skip") if parameter_name in signature.parameters ), None, ) handler_handles_pagination = bool( handler_limit_parameter or handler_offset_parameter ) @wraps(endpoint) async def wrapper(*args, **kwargs): if handler_handles_pagination: limit = kwargs.get(handler_limit_parameter, 0) offset = kwargs.get(handler_offset_parameter, 0) else: limit = kwargs.pop("_rest_limit") offset = kwargs.pop("_rest_offset") result = endpoint(*args, **kwargs) if inspect.isawaitable(result): result = await result if not isinstance(result, list): return result if handler_handles_pagination: return Page( items=result, total=offset + len(result), limit=limit or len(result), offset=offset, ) return Page( items=result[offset : offset + limit], total=len(result), limit=limit, offset=offset, ) parameters = list(signature.parameters.values()) if not handler_handles_pagination: parameters.extend( [ inspect.Parameter( "_rest_limit", kind=inspect.Parameter.KEYWORD_ONLY, annotation=int, default=Query(100, ge=1, le=1000, alias="limit"), ), inspect.Parameter( "_rest_offset", kind=inspect.Parameter.KEYWORD_ONLY, annotation=int, default=Query(0, ge=0, alias="offset"), ), ] ) wrapper.__signature__ = signature.replace(parameters=parameters) return wrapper def _adapt_route(route: APIRoute) -> APIRoute: methods = route.methods or set() if len(methods) != 1: raise RuntimeError( f"REST route {route.name!r} must declare exactly one HTTP method" ) method = next(iter(methods)) responses = dict(route.responses or {}) for status_code, description in ( (401, "Authentication required"), (403, "Insufficient permission"), (404, "Resource not found"), (409, "Resource conflict"), (422, "Validation error"), (503, "Dependency unavailable"), ): responses.setdefault( status_code, {"model": ProblemDetails, "description": description}, ) endpoint = _with_header_project_context(route.endpoint, route.name) response_model = route.response_model if get_origin(response_model) is list: item_type = get_args(response_model)[0] if get_args(response_model) else JsonValue response_model = Page[item_type] endpoint = _with_pagination(endpoint) clean_name = _clean_name(route.name) creates_resource = clean_name.startswith( ("add_", "create_", "copy_", "import_", "insert_", "store_", "take_", "upload_") ) or route.name == "fastapi_pressure_sensor_placement" status_code = ( 204 if method == "DELETE" else 201 if method == "POST" and creates_resource else route.status_code ) if status_code == 204: response_model = None elif response_model is None: response_model = JsonValue return APIRoute( path=route.path, endpoint=endpoint, response_model=response_model, status_code=status_code, tags=route.tags, dependencies=route.dependencies, summary=route.summary, description=route.description, response_description=route.response_description, responses=responses, deprecated=False, name=route.name, methods={method}, operation_id=f"{method.lower()}_{re.sub(r'[^a-z0-9]+', '_', route.path).strip('_')}", response_model_include=route.response_model_include, response_model_exclude=route.response_model_exclude, response_model_by_alias=route.response_model_by_alias, response_model_exclude_unset=route.response_model_exclude_unset, response_model_exclude_defaults=route.response_model_exclude_defaults, response_model_exclude_none=route.response_model_exclude_none, include_in_schema=route.include_in_schema, response_class=route.response_class, callbacks=route.callbacks, openapi_extra=route.openapi_extra, ) def build_rest_router(routes: Iterable[Any]) -> APIRouter: router = APIRouter() seen: dict[tuple[str, str], APIRoute] = {} operation_ids: set[str] = set() for route in routes: if not isinstance(route, APIRoute): continue methods = route.methods or set() if len(methods) != 1: raise RuntimeError( f"REST route {route.name!r} must declare exactly one HTTP method" ) method = next(iter(methods)) key = (method, route.path) if key in seen: previous = seen[key] raise RuntimeError( "REST route collision for " f"{method} {route.path}: {previous.name!r} and {route.name!r}." ) adapted = _adapt_route(route) if adapted.operation_id in operation_ids: adapted.operation_id = f"{adapted.operation_id}_{route.name}" seen[key] = route operation_ids.add(adapted.operation_id or "") router.routes.append(adapted) return router api_router = build_rest_router(handler_api_router.routes)