# 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)