diff --git a/src/faststream_fastapi/_internal/get_dependant.py b/src/faststream_fastapi/_internal/get_dependant.py index b44771f..b221f8e 100644 --- a/src/faststream_fastapi/_internal/get_dependant.py +++ b/src/faststream_fastapi/_internal/get_dependant.py @@ -1,7 +1,9 @@ import inspect from collections.abc import Callable, Iterable +from dataclasses import fields from typing import Annotated, Any, Final, cast, get_args, get_origin +from fast_depends.library.serializer import OptionItem from fast_depends.utils import get_typed_annotation from fastapi.dependencies.models import Dependant from fastapi.dependencies.utils import ( @@ -10,15 +12,30 @@ get_typed_signature, ) from fastapi.params import Depends -from pydantic import Field +from pydantic import Field, create_model from faststream_fastapi._internal.fs_re_exports._compat import PYDANTIC_V2, PydanticUndefined +class _FastStreamDependant(Dependant): + """FastAPI dependant extended with fields required by FastStream.""" + + model: type[Any] + custom_fields: dict[str, Any] + flat_params: list[OptionItem] + + +def _extend_fastapi_dependant(dependant: Dependant) -> _FastStreamDependant: + field_values = { + field.name: getattr(dependant, field.name) for field in fields(dependant) if field.init + } + return _FastStreamDependant(**field_values) + + def get_fastapi_dependant( orig_call: Callable[..., Any], dependencies: Iterable[Depends], -) -> Dependant: +) -> _FastStreamDependant: dependent = get_fastapi_native_dependant(orig_call=orig_call, dependencies=dependencies) return _patch_fastapi_dependent(dependent) @@ -41,7 +58,8 @@ def get_fastapi_native_dependant( return dependent -def _patch_fastapi_dependent(dependant: Dependant) -> Dependant: +def _patch_fastapi_dependent(dependant: Dependant) -> _FastStreamDependant: + dependant = _extend_fastapi_dependant(dependant) params = dependant.query_params + dependant.body_params for d in dependant.dependencies: @@ -117,6 +135,15 @@ def _patch_fastapi_dependent(dependant: Dependant) -> Dependant: f, ) + dependant.model = create_model( + getattr(call, "__name__", type(call).__name__), + ) + dependant.custom_fields = {} + dependant.flat_params = [ + OptionItem(field_name=name, field_type=type_, default_value=default) + for name, (type_, default) in params_unique.items() + ] + return dependant diff --git a/tests/test_get_fastapi_dependant.py b/tests/test_get_fastapi_dependant.py new file mode 100644 index 0000000..b1ef984 --- /dev/null +++ b/tests/test_get_fastapi_dependant.py @@ -0,0 +1,25 @@ +from typing import Annotated + +from fastapi import Depends + +from faststream_fastapi._internal.get_dependant import get_fastapi_dependant + + +def test_get_fastapi_dependant_has_faststream_fields() -> None: + async def dependency() -> str: + return "dependency" + + async def handler( + message: str, + dependency_value: Annotated[str, Depends(dependency)], + ) -> None: + pass + + dependant = get_fastapi_dependant(handler, ()) + + assert dependant.call is handler + assert len(dependant.dependencies) == 1 + assert dependant.dependencies[0].call is dependency + assert dependant.model.__name__ == "handler" + assert dependant.custom_fields == {} + assert [field.field_name for field in dependant.flat_params] == ["message"]