Files
TJWaterServerBinary/app/api/v1/rest_router.py
T
jiang 1d88f8efbe
Server CI/CD / docker-image (push) Failing after 28s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
fix(api): wrap pre-paginated list responses
2026-07-30 21:50:21 +08:00

372 lines
13 KiB
Python

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)