350 lines
12 KiB
Python
350 lines
12 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)
|
|
if "limit" in signature.parameters or "offset" in signature.parameters:
|
|
return endpoint
|
|
|
|
@wraps(endpoint)
|
|
async def wrapper(*args, **kwargs):
|
|
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
|
|
return Page(
|
|
items=result[offset : offset + limit],
|
|
total=len(result),
|
|
limit=limit,
|
|
offset=offset,
|
|
)
|
|
|
|
parameters = list(signature.parameters.values())
|
|
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)
|