|
# retoor <retoor@molodetz.nl>
|
|
"""
|
|
Generic FastAPI dependency that accepts JSON or form-encoded data,
|
|
validated against a Pydantic model.
|
|
"""
|
|
|
|
import json
|
|
import logging
|
|
from typing import Any, TypeVar, get_origin
|
|
|
|
from fastapi import HTTPException, Request
|
|
from fastapi.exceptions import RequestValidationError
|
|
from pydantic import BaseModel, ValidationError
|
|
from starlette.datastructures import FormData
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_TModel = TypeVar("_TModel", bound=BaseModel)
|
|
|
|
# Container origins recognised as sequence fields that may receive
|
|
# multiple values from form data.
|
|
_SEQUENCE_ORIGINS = frozenset({list, set, tuple, frozenset})
|
|
|
|
|
|
def _formdata_to_dict(form: FormData, model: type[BaseModel]) -> dict[str, Any]:
|
|
"""Convert FormData to a dict suitable for Pydantic validation.
|
|
|
|
* Sequence-typed model fields collect every submitted value via
|
|
``getlist()``; a lone empty string is dropped (browsers emit empty
|
|
hidden inputs by default).
|
|
* Scalar fields use ``get()`` (the last value).
|
|
* Fields absent from the form are omitted so that Pydantic applies
|
|
the model default.
|
|
"""
|
|
body: dict[str, Any] = {}
|
|
for field_name, field_info in model.model_fields.items():
|
|
origin = get_origin(field_info.annotation)
|
|
if origin in _SEQUENCE_ORIGINS:
|
|
values = form.getlist(field_name)
|
|
if not values:
|
|
continue
|
|
if values == [""]:
|
|
continue
|
|
body[field_name] = [v for v in values if v != ""] or []
|
|
else:
|
|
value = form.get(field_name)
|
|
if value is not None:
|
|
body[field_name] = value
|
|
return body
|
|
|
|
|
|
class _JsonOrForm:
|
|
"""Internal callable that parses JSON or form data and validates."""
|
|
|
|
def __init__(self, model: type[BaseModel]):
|
|
self.model = model
|
|
|
|
async def __call__(self, request: Request) -> Any:
|
|
content_type = request.headers.get("content-type", "")
|
|
body: Any = None
|
|
try:
|
|
if "application/json" in content_type:
|
|
try:
|
|
body = await request.json()
|
|
except (json.JSONDecodeError, UnicodeDecodeError, ValueError) as exc:
|
|
logger.debug("JSON parse failed: %s", exc)
|
|
raise HTTPException(status_code=400, detail="Invalid JSON body")
|
|
if not isinstance(body, dict):
|
|
raise HTTPException(
|
|
status_code=400, detail="JSON body must be an object"
|
|
)
|
|
return self.model.model_validate(body)
|
|
# Default: form-encoded (multipart or url-encoded)
|
|
try:
|
|
form = await request.form()
|
|
except Exception as exc:
|
|
logger.debug("Form parse failed: %s", exc)
|
|
raise HTTPException(
|
|
status_code=400, detail="Could not parse form data"
|
|
)
|
|
body = _formdata_to_dict(form, self.model)
|
|
return self.model.model_validate(body)
|
|
except ValidationError as exc:
|
|
raise RequestValidationError(errors=exc.errors(), body=body)
|
|
|
|
|
|
def json_or_form(model: type[_TModel]) -> _JsonOrForm:
|
|
"""Dependency factory: accept JSON or form-encoded data for a Pydantic model.
|
|
|
|
Usage:
|
|
@router.post("/create")
|
|
async def create(data: Annotated[PostForm, Depends(json_or_form(PostForm))]):
|
|
...
|
|
"""
|
|
return _JsonOrForm(model)
|