feat(api): standardize REST contracts and auth
This commit is contained in:
@@ -0,0 +1,349 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user