from __future__ import annotations from functools import partial from typing import Type from marshmallow import Schema __all__ = ["request_schema", "querystring_schema", "path_schema", "headers_schema", "cookies_schema"] def _get_wrapper(schema: Schema | list[Schema], location: str): def wrapper(func): if not hasattr(func, "__apispec__"): setattr(func, "__apispec__", {"request": {}, "responses": {}}) func.__apispec__["request"][location] = schema return func return wrapper def _params_schema(schema: Schema | Type[Schema], location: str): if not isinstance(schema, Schema): schema = schema() return _get_wrapper(schema, location) def request_schema(*schema: Schema | Type[Schema]): schemas = [] for schema_ in schema: if not isinstance(schema_, Schema): schema_ = schema_() schemas.append(schema_) return _get_wrapper(schemas, "body") querystring_schema = partial( _params_schema, location="query", ) path_schema = partial(_params_schema, location="path") headers_schema = partial(_params_schema, location="headers") cookies_schema = partial(_params_schema, location="cookies")