From 98c91ff671581725c849c38343ebe0c13fc75cf3 Mon Sep 17 00:00:00 2001 From: alex Date: Mon, 13 Jul 2026 00:38:12 +0200 Subject: [PATCH 1/3] fix: resolve postponed injection annotations --- tests/test_future_annotations.py | 68 ++++++++++++++++++++++++++++++++ that_depends/injection.py | 21 +++++++++- 2 files changed, 87 insertions(+), 2 deletions(-) create mode 100644 tests/test_future_annotations.py diff --git a/tests/test_future_annotations.py b/tests/test_future_annotations.py new file mode 100644 index 00000000..4cb53564 --- /dev/null +++ b/tests/test_future_annotations.py @@ -0,0 +1,68 @@ +from __future__ import annotations +import typing + +import pytest + +from that_depends import BaseContainer, Provide, providers + + +if typing.TYPE_CHECKING: + + class Unresolvable: + pass + + +class Service: + pass + + +class FutureAnnotationsContainer(BaseContainer): + service = providers.Object(Service()).bind(Service) + + +@FutureAnnotationsContainer.inject +def inject_service_sync(service: Service = Provide()) -> Service: + return service + + +@FutureAnnotationsContainer.inject +async def inject_service_async(service: Service = Provide()) -> Service: + return service + + +def test_type_based_injection_resolves_postponed_annotations_sync() -> None: + assert isinstance(inject_service_sync(), Service) + + +async def test_type_based_injection_resolves_postponed_annotations_async() -> None: + assert isinstance(await inject_service_async(), Service) + + +def test_type_based_injection_rejects_unresolvable_local_annotation() -> None: + class LocalService: + pass + + with pytest.raises(TypeError, match="Cannot resolve annotations for injected function"): + + @FutureAnnotationsContainer.inject + def target(service: LocalService = Provide()) -> LocalService: # pragma: no cover + return service + + +def test_type_based_injection_rejects_non_concrete_annotation() -> None: + with pytest.raises(TypeError, match="Type-based injection for 'service' requires a concrete runtime type"): + + @FutureAnnotationsContainer.inject + def target(service: list[str] = Provide()) -> list[str]: # pragma: no cover + return service + + +def test_direct_provider_injection_does_not_resolve_unrelated_annotations() -> None: + provider = providers.Object(1) + + @FutureAnnotationsContainer.inject + def target(value: int = Provide[provider], unrelated: Unresolvable | None = None) -> int: + _ = unrelated + return value + + assert target() == 1 diff --git a/that_depends/injection.py b/that_depends/injection.py index bd020b73..4d8d076e 100644 --- a/that_depends/injection.py +++ b/that_depends/injection.py @@ -48,6 +48,7 @@ class _TypedInjectionParameter(typing.NamedTuple): class _InjectionPlan(typing.NamedTuple): + signature: inspect.Signature direct_parameters: tuple[_DirectInjectionParameter, ...] string_parameters: tuple[_StringInjectionParameter, ...] typed_parameters: tuple[_TypedInjectionParameter, ...] @@ -98,10 +99,21 @@ def close(self) -> None: @functools.cache def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionPlan: + signature = inspect.signature(func) + parameters = tuple(signature.parameters.items()) direct_parameters: list[_DirectInjectionParameter] = [] string_parameters: list[_StringInjectionParameter] = [] typed_parameters: list[_TypedInjectionParameter] = [] - for index, (field_name, param) in enumerate(inspect.signature(func).parameters.items()): + if any(isinstance(param.default, _Provide) for _, param in parameters): + try: + resolved_hints = typing.get_type_hints(func) + except (NameError, TypeError) as exc: + msg = f"Cannot resolve annotations for injected function {func.__qualname__}" + raise TypeError(msg) from exc + else: + resolved_hints = {} + + for index, (field_name, param) in enumerate(parameters): default = param.default if isinstance(default, StringProviderDefinition): string_parameters.append(_StringInjectionParameter(index, field_name, default)) @@ -115,14 +127,19 @@ def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionP ) ) elif isinstance(default, _Provide): + annotation = resolved_hints.get(field_name, param.annotation) + if annotation is inspect.Parameter.empty or annotation is typing.Any or not isinstance(annotation, type): + msg = f"Type-based injection for {field_name!r} requires a concrete runtime type" + raise TypeError(msg) typed_parameters.append( _TypedInjectionParameter( index, field_name, - typing.cast(type[typing.Any], param.annotation), + annotation, ) ) return _InjectionPlan( + signature=signature, direct_parameters=tuple(direct_parameters), string_parameters=tuple(string_parameters), typed_parameters=tuple(typed_parameters), From fa0f380262cf12e3c816d5279d76aec24978d2a4 Mon Sep 17 00:00:00 2001 From: alex Date: Mon, 13 Jul 2026 00:39:30 +0200 Subject: [PATCH 2/3] fix: bind injected arguments by signature --- tests/test_injection.py | 29 ++++++++++++++++++++++++++--- that_depends/injection.py | 38 ++++++++++++++++---------------------- 2 files changed, 42 insertions(+), 25 deletions(-) diff --git a/tests/test_injection.py b/tests/test_injection.py index ac4f3d5b..3f0e02ac 100644 --- a/tests/test_injection.py +++ b/tests/test_injection.py @@ -99,7 +99,7 @@ def _injected(value: providers.Object[int] = provider) -> providers.Object[int]: assert _injected(provider) is provider assert plan.direct_parameters == ( - _DirectInjectionParameter(0, "value", provider, provider._get_scope_context_init_order()), + _DirectInjectionParameter("value", provider, provider._get_scope_context_init_order()), ) @@ -110,7 +110,7 @@ def _injected(value: float = Provide()) -> float: plan = _build_injection_plan(_injected) assert _injected(1.0) == 1.0 - assert plan.typed_parameters == (_TypedInjectionParameter(0, "value", float),) + assert plan.typed_parameters == (_TypedInjectionParameter("value", float),) def test_build_injection_plan_stores_string_provider_separately() -> None: @@ -121,11 +121,34 @@ def _injected(value: int = Provide["Container.provider"]) -> int: assert _injected(1) == 1 assert len(plan.string_parameters) == 1 - assert plan.string_parameters[0].argument_index == 0 assert plan.string_parameters[0].field_name == "value" assert plan.string_parameters[0].definition._definition == "Container.provider" +def test_injection_handles_keyword_only_parameter_after_varargs() -> None: + injected_value = 42 + override_value = 7 + provider = providers.Object(injected_value) + + @inject + def target(*values: str, dependency: int = Provide[provider]) -> int: + _ = values + return dependency + + assert target("one", "two") == injected_value + assert target("one", dependency=override_value) == override_value + + +def test_injection_rejects_positional_only_parameter() -> None: + provider = providers.Object(42) + + with pytest.raises(TypeError, match="Injected parameter 'dependency' cannot be positional-only"): + + @inject + def target(dependency: int = Provide[provider], /) -> int: # pragma: no cover + return dependency + + async def test_empty_injection() -> None: @inject async def inner(_: int) -> None: diff --git a/that_depends/injection.py b/that_depends/injection.py index 4d8d076e..bb8c1965 100644 --- a/that_depends/injection.py +++ b/that_depends/injection.py @@ -29,20 +29,17 @@ class ContextProviderError(Exception): class _DirectInjectionParameter(typing.NamedTuple): - argument_index: int field_name: str provider: AbstractProvider[typing.Any] scope_context_init_order: tuple[AbstractProvider[typing.Any], ...] class _StringInjectionParameter(typing.NamedTuple): - argument_index: int field_name: str definition: "StringProviderDefinition" class _TypedInjectionParameter(typing.NamedTuple): - argument_index: int field_name: str annotation: type[typing.Any] @@ -113,14 +110,19 @@ def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionP else: resolved_hints = {} - for index, (field_name, param) in enumerate(parameters): + for field_name, param in parameters: default = param.default + if ( + isinstance(default, (StringProviderDefinition, AbstractProvider, _Provide)) + and param.kind is inspect.Parameter.POSITIONAL_ONLY + ): + msg = f"Injected parameter {field_name!r} cannot be positional-only" + raise TypeError(msg) if isinstance(default, StringProviderDefinition): - string_parameters.append(_StringInjectionParameter(index, field_name, default)) + string_parameters.append(_StringInjectionParameter(field_name, default)) elif isinstance(default, AbstractProvider): direct_parameters.append( _DirectInjectionParameter( - index, field_name, default, default._get_scope_context_init_order(), # noqa: SLF001 @@ -133,7 +135,6 @@ def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionP raise TypeError(msg) typed_parameters.append( _TypedInjectionParameter( - index, field_name, annotation, ) @@ -284,8 +285,9 @@ async def _resolve_arguments_async( return False, kwargs context_providers: set[AbstractProvider[typing.Any]] = set() + provided_names = plan.signature.bind_partial(*args, **kwargs).arguments for direct_parameter in plan.direct_parameters: - if _is_argument_provided(direct_parameter.argument_index, direct_parameter.field_name, args, kwargs): + if direct_parameter.field_name in provided_names: continue if direct_parameter.scope_context_init_order: @@ -298,7 +300,7 @@ async def _resolve_arguments_async( kwargs[direct_parameter.field_name] = await direct_parameter.provider.resolve() for string_parameter in plan.string_parameters: - if _is_argument_provided(string_parameter.argument_index, string_parameter.field_name, args, kwargs): + if string_parameter.field_name in provided_names: continue kwargs[string_parameter.field_name] = await _resolve_provider_with_scope_async( @@ -309,7 +311,7 @@ async def _resolve_arguments_async( ) for typed_parameter in plan.typed_parameters: - if _is_argument_provided(typed_parameter.argument_index, typed_parameter.field_name, args, kwargs): + if typed_parameter.field_name in provided_names: continue provider = _resolve_typed_provider(typed_parameter.annotation, container) @@ -334,8 +336,9 @@ def _resolve_arguments_sync( return False, kwargs context_providers: set[AbstractProvider[typing.Any]] = set() + provided_names = plan.signature.bind_partial(*args, **kwargs).arguments for direct_parameter in plan.direct_parameters: - if _is_argument_provided(direct_parameter.argument_index, direct_parameter.field_name, args, kwargs): + if direct_parameter.field_name in provided_names: continue if direct_parameter.scope_context_init_order: @@ -348,7 +351,7 @@ def _resolve_arguments_sync( kwargs[direct_parameter.field_name] = direct_parameter.provider.resolve_sync() for string_parameter in plan.string_parameters: - if _is_argument_provided(string_parameter.argument_index, string_parameter.field_name, args, kwargs): + if string_parameter.field_name in provided_names: continue kwargs[string_parameter.field_name] = _resolve_provider_with_scope_sync( @@ -359,7 +362,7 @@ def _resolve_arguments_sync( ) for typed_parameter in plan.typed_parameters: - if _is_argument_provided(typed_parameter.argument_index, typed_parameter.field_name, args, kwargs): + if typed_parameter.field_name in provided_names: continue provider = _resolve_typed_provider(typed_parameter.annotation, container) @@ -377,15 +380,6 @@ def _plan_has_injected_parameters(plan: _InjectionPlan) -> bool: return bool(plan.direct_parameters or plan.string_parameters or plan.typed_parameters) -def _is_argument_provided( - argument_index: int, - field_name: str, - args: tuple[typing.Any, ...], - kwargs: dict[str, typing.Any], -) -> bool: - return argument_index < len(args) or field_name in kwargs - - def _resolve_typed_provider( annotation: type[typing.Any], container: BaseContainerMeta | None, From e0fdd6d31022394fd7a254069f5db22dac072aad Mon Sep 17 00:00:00 2001 From: alex Date: Mon, 13 Jul 2026 11:39:38 +0200 Subject: [PATCH 3/3] feat: selector context registration. --- tests/providers/test_selector.py | 21 ++++ tests/test_injection.py | 185 ++++++++++++++++++++++++++++- that_depends/injection.py | 144 +++++++++++----------- that_depends/providers/selector.py | 58 +++++++-- 4 files changed, 325 insertions(+), 83 deletions(-) diff --git a/tests/providers/test_selector.py b/tests/providers/test_selector.py index fe7b784c..865d05a2 100644 --- a/tests/providers/test_selector.py +++ b/tests/providers/test_selector.py @@ -140,6 +140,27 @@ async def test_selector_with_provider_selector_async() -> None: assert (await StringProviderSelectorContainer.selector.resolve()) == "Provider 1" +def test_selector_registers_only_its_key_provider() -> None: + def _selector_key() -> typing.Iterator[str]: # pragma: no cover + yield "selected" + + selector_key = providers.ContextResource(_selector_key) + selected = providers.Object("value") + selector = providers.Selector(selector_key, selected=selected) + + selector._register_arguments() + selector._register_arguments() + + assert selector in selector_key._children + assert selected not in selector._parents + assert selector._get_scope_context_init_order() == (selector_key,) + assert selector._get_scope_context_init_order() == (selector_key,) + + selector._deregister_arguments() + + assert selector not in selector_key._children + + class InvalidSelectorContainer(BaseContainer): selector = providers.Selector( None, # type: ignore[arg-type] diff --git a/tests/test_injection.py b/tests/test_injection.py index 3f0e02ac..4b324392 100644 --- a/tests/test_injection.py +++ b/tests/test_injection.py @@ -98,9 +98,7 @@ def _injected(value: providers.Object[int] = provider) -> providers.Object[int]: plan = _build_injection_plan(_injected) assert _injected(provider) is provider - assert plan.direct_parameters == ( - _DirectInjectionParameter("value", provider, provider._get_scope_context_init_order()), - ) + assert plan.direct_parameters == (_DirectInjectionParameter("value", provider),) def test_build_injection_plan_stores_annotation_for_type_based_injection() -> None: @@ -756,6 +754,187 @@ def injected(repo: DocumentRepository = Provide[_Container._factory_provider]) - assert isinstance(injected(), float) +def test_injection_enters_only_selected_selector_context_sync() -> None: + events: list[str] = [] + + def _sync_creator() -> typing.Iterator[str]: + events.append("sync enter") + try: + yield "sync" + finally: + events.append("sync exit") + + async def _async_creator() -> typing.AsyncIterator[str]: # pragma: no cover + events.append("async enter") + try: + yield "async" + finally: + events.append("async exit") + + selector = providers.Selector( + lambda: "sync", + sync=providers.ContextResource(_sync_creator).with_config(scope=ContextScopes.INJECT), + async_=providers.ContextResource(_async_creator).with_config(scope=ContextScopes.INJECT), + ) + + @inject + def target(value: str = Provide[selector]) -> str: + return value + + assert target() == "sync" + assert events == ["sync enter", "sync exit"] + + +async def test_injection_enters_only_selected_selector_context_async() -> None: + events: list[str] = [] + + def _sync_creator() -> typing.Iterator[str]: # pragma: no cover + events.append("sync enter") + try: + yield "sync" + finally: + events.append("sync exit") + + async def _async_creator() -> typing.AsyncIterator[str]: + events.append("async enter") + try: + yield "async" + finally: + events.append("async exit") + + selector = providers.Selector( + lambda: "async_", + sync=providers.ContextResource(_sync_creator).with_config(scope=ContextScopes.INJECT), + async_=providers.ContextResource(_async_creator).with_config(scope=ContextScopes.INJECT), + ) + + @inject + async def target(value: str = Provide[selector]) -> str: + return value + + assert await target() == "async" + assert events == ["async enter", "async exit"] + + +def test_injection_pins_nested_selector_branch_sync() -> None: + selected_keys: list[str] = [] + events: list[str] = [] + + def _select_key() -> str: + key = "first" if not selected_keys else "second" + selected_keys.append(key) + return key + + def _creator(value: str) -> typing.Iterator[str]: + events.append(f"{value} enter") + try: + yield value + finally: + events.append(f"{value} exit") + + selector = providers.Selector( + _select_key, + first=providers.ContextResource(_creator, "first").with_config(scope=ContextScopes.INJECT), + second=providers.ContextResource(_creator, "second").with_config(scope=ContextScopes.INJECT), + ) + factory = providers.Factory(lambda value: (value, selector.resolve_sync()), selector.cast) + + @inject + def target(value: tuple[str, str] = Provide[factory]) -> tuple[str, str]: + return value + + assert target() == ("first", "first") + assert selected_keys == ["first"] + assert events == ["first enter", "first exit"] + + +async def test_injection_pins_nested_selector_branch_async() -> None: + selected_keys: list[str] = [] + events: list[str] = [] + + def _select_key() -> str: + key = "first" if not selected_keys else "second" + selected_keys.append(key) + return key + + async def _creator(value: str) -> typing.AsyncIterator[str]: + events.append(f"{value} enter") + try: + yield value + finally: + events.append(f"{value} exit") + + selector = providers.Selector( + _select_key, + first=providers.ContextResource(_creator, "first").with_config(scope=ContextScopes.INJECT), + second=providers.ContextResource(_creator, "second").with_config(scope=ContextScopes.INJECT), + ) + + async def _factory(value: str) -> tuple[str, str]: + return value, await selector.resolve() + + factory = providers.AsyncFactory(_factory, selector.cast) + + @inject + async def target(value: tuple[str, str] = Provide[factory]) -> tuple[str, str]: + return value + + assert await target() == ("first", "first") + assert selected_keys == ["first"] + assert events == ["first enter", "first exit"] + + +def test_string_injection_prepares_selected_selector_context() -> None: + def _creator() -> typing.Iterator[str]: + yield "selected" + + class _Container(BaseContainer): + resource = providers.ContextResource(_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector("selected", selected=resource) + + @inject + def target(value: str = Provide["_Container.selector"]) -> str: + return value + + assert target() == "selected" + + +async def test_type_injection_prepares_selected_selector_context() -> None: + async def _creator() -> typing.AsyncIterator[str]: + yield "selected" + + class _Container(BaseContainer): + resource = providers.ContextResource(_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector("selected", selected=resource).bind(str) + + @_Container.inject + async def target(value: str = Provide()) -> str: + return value + + assert await target() == "selected" + + +def test_injection_scope_none_does_not_enter_selector_context() -> None: + events: list[str] = [] + + def _creator() -> typing.Iterator[str]: # pragma: no cover + events.append("entered") + yield "selected" + + selector = providers.Selector( + "selected", + selected=providers.ContextResource(_creator).with_config(scope=ContextScopes.INJECT), + ) + + @inject(scope=None) + def target(value: str = Provide[selector]) -> str: + return value # pragma: no cover + + with pytest.raises(RuntimeError, match="Context is not set"): + target() + assert events == [] + + def test_simple_injection_into_iterator_sync() -> None: class _Container(BaseContainer): sync_resource = providers.Factory(random.random) diff --git a/that_depends/injection.py b/that_depends/injection.py index bb8c1965..21ec6e17 100644 --- a/that_depends/injection.py +++ b/that_depends/injection.py @@ -12,7 +12,10 @@ from that_depends.exceptions import TypeNotBoundError from that_depends.meta import BaseContainerMeta from that_depends.providers import AbstractProvider -from that_depends.providers.context_resources import ContextScope, ContextScopes, container_context +from that_depends.providers.context_resources import ContextResource, ContextScope, ContextScopes, container_context +from that_depends.providers.mixin import ProviderWithArguments +from that_depends.providers.selector import Selector +from that_depends.utils import is_set class ContextProviderError(Exception): @@ -31,7 +34,6 @@ class ContextProviderError(Exception): class _DirectInjectionParameter(typing.NamedTuple): field_name: str provider: AbstractProvider[typing.Any] - scope_context_init_order: tuple[AbstractProvider[typing.Any], ...] class _StringInjectionParameter(typing.NamedTuple): @@ -125,7 +127,6 @@ def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionP _DirectInjectionParameter( field_name, default, - default._get_scope_context_init_order(), # noqa: SLF001 ) ) elif isinstance(default, _Provide): @@ -290,14 +291,12 @@ async def _resolve_arguments_async( if direct_parameter.field_name in provided_names: continue - if direct_parameter.scope_context_init_order: - await _setup_scope_contexts_async( - direct_parameter.scope_context_init_order, - scope, - stack, - context_providers, - ) - kwargs[direct_parameter.field_name] = await direct_parameter.provider.resolve() + kwargs[direct_parameter.field_name] = await _resolve_provider_with_scope_async( + direct_parameter.provider, + scope, + stack, + context_providers, + ) for string_parameter in plan.string_parameters: if string_parameter.field_name in provided_names: @@ -341,14 +340,12 @@ def _resolve_arguments_sync( if direct_parameter.field_name in provided_names: continue - if direct_parameter.scope_context_init_order: - _setup_scope_contexts_sync( - direct_parameter.scope_context_init_order, - scope, - stack, - context_providers, - ) - kwargs[direct_parameter.field_name] = direct_parameter.provider.resolve_sync() + kwargs[direct_parameter.field_name] = _resolve_provider_with_scope_sync( + direct_parameter.provider, + scope, + stack, + context_providers, + ) for string_parameter in plan.string_parameters: if string_parameter.field_name in provided_names: @@ -401,13 +398,6 @@ def _resolve_sync( *args: P.args, **kwargs: P.kwargs, ) -> T: - if scope is None: - injected, kwargs = _resolve_arguments_sync(plan, scope, container, None, *args, **kwargs) # type: ignore[assignment] - if not injected: - warnings.warn(_INJECTION_WARNING_MESSAGE, RuntimeWarning, stacklevel=3) - - return func(*args, **kwargs) - with _SyncInjectionStack() as stack: injected, kwargs = _resolve_arguments_sync(plan, scope, container, stack, *args, **kwargs) # type: ignore[assignment] @@ -460,33 +450,43 @@ async def _resolve_provider_with_scope_async( ContextProviderError: if the stack is None. """ - scope_context_init_order = provider._get_scope_context_init_order() # noqa: SLF001 - if scope_context_init_order: - await _setup_scope_contexts_async(scope_context_init_order, scope, stack, providers) + await _prepare_provider_contexts_async(provider, scope, stack, providers) return await provider.resolve() -async def _setup_scope_contexts_async( - scope_init_order: tuple[AbstractProvider[typing.Any], ...], +async def _prepare_provider_contexts_async( + provider: AbstractProvider[typing.Any], scope: ContextScope | None, stack: AsyncExitStack | None, - providers: set[AbstractProvider[typing.Any]], + visited: set[AbstractProvider[typing.Any]], ) -> None: - if not scope: + if provider in visited: return - for provider in scope_init_order: - if provider in providers: - continue - providers.add(provider) - provider_scope = provider._scope # noqa: SLF001 - if provider_scope in (ContextScopes.ANY, scope): - if stack is None: - msg = ( - f"No stack exists, cannot initialize context for {provider} using scope {scope}.\n" - f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." - ) - raise ContextProviderError(msg) - await stack.enter_async_context(provider.context_async(force=True)) + visited.add(provider) + + if isinstance(provider, ProviderWithArguments): + provider._register_arguments() # noqa: SLF001 + for parent in provider._parents: # noqa: SLF001 + await _prepare_provider_contexts_async(parent, scope, stack, visited) + + if isinstance(provider, Selector) and not is_set(provider._override): # noqa: SLF001 + selected_provider = await provider._select_provider() # noqa: SLF001 + if stack is not None: + stack.enter_context(provider._pin_selected_provider(selected_provider)) # noqa: SLF001 + await _prepare_provider_contexts_async(selected_provider, scope, stack, visited) + + if ( + scope is not None + and isinstance(provider, ContextResource) + and provider.get_scope() in (ContextScopes.ANY, scope) + ): + if stack is None: + msg = ( + f"No stack exists, cannot initialize context for {provider} using scope {scope}.\n" + f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." + ) + raise ContextProviderError(msg) + await stack.enter_async_context(provider.context_async(force=True)) def _resolve_provider_with_scope_sync( @@ -495,34 +495,44 @@ def _resolve_provider_with_scope_sync( stack: _SyncInjectionStack | None, providers: set[AbstractProvider[typing.Any]], ) -> T: - scope_context_init_order = provider._get_scope_context_init_order() # noqa: SLF001 - if scope_context_init_order: - _setup_scope_contexts_sync(scope_context_init_order, scope, stack, providers) + _prepare_provider_contexts_sync(provider, scope, stack, providers) return provider.resolve_sync() -def _setup_scope_contexts_sync( - scope_init_order: tuple[AbstractProvider[typing.Any], ...], +def _prepare_provider_contexts_sync( + provider: AbstractProvider[typing.Any], scope: ContextScope | None, stack: _SyncInjectionStack | None, - providers: set[AbstractProvider[typing.Any]], + visited: set[AbstractProvider[typing.Any]], ) -> None: - if not scope: + if provider in visited: return - for provider in scope_init_order: - if provider in providers: - continue - providers.add(provider) - provider_scope = provider._scope # noqa: SLF001 - if provider_scope in (ContextScopes.ANY, scope): - if stack is None: - msg = ( - f"No stack exists, cannot initialize context for {provider} using scope {scope}.\n" - f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." - ) - raise ContextProviderError(msg) - _, exit_state = provider._enter_injection_context_sync(force=True) # noqa: SLF001 - stack.push_exit_state(exit_state) + visited.add(provider) + + if isinstance(provider, ProviderWithArguments): + provider._register_arguments() # noqa: SLF001 + for parent in provider._parents: # noqa: SLF001 + _prepare_provider_contexts_sync(parent, scope, stack, visited) + + if isinstance(provider, Selector) and not is_set(provider._override): # noqa: SLF001 + selected_provider = provider._select_provider_sync() # noqa: SLF001 + if stack is not None: + stack.enter_context(provider._pin_selected_provider(selected_provider)) # noqa: SLF001 + _prepare_provider_contexts_sync(selected_provider, scope, stack, visited) + + if ( + scope is not None + and isinstance(provider, ContextResource) + and provider.get_scope() in (ContextScopes.ANY, scope) + ): + if stack is None: + msg = ( + f"No stack exists, cannot initialize context for {provider} using scope {scope}.\n" + f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." + ) + raise ContextProviderError(msg) + _, exit_state = provider._enter_injection_context_sync(force=True) # noqa: SLF001 + stack.push_exit_state(exit_state) class StringProviderDefinition: diff --git a/that_depends/providers/selector.py b/that_depends/providers/selector.py index 9ec2c48b..975869f1 100644 --- a/that_depends/providers/selector.py +++ b/that_depends/providers/selector.py @@ -1,17 +1,20 @@ """Selection based providers.""" import typing +from contextlib import contextmanager +from contextvars import ContextVar from typing_extensions import override from that_depends.providers.base import AbstractProvider -from that_depends.utils import is_set +from that_depends.providers.mixin import ProviderWithArguments +from that_depends.utils import UNSET, Unset, is_set T_co = typing.TypeVar("T_co", covariant=True) -class Selector(AbstractProvider[T_co]): +class Selector(ProviderWithArguments, AbstractProvider[T_co]): """Chooses a provider based on a key returned by a selector function. This class allows you to dynamically select and resolve one of several @@ -35,7 +38,7 @@ def environment_selector(): """ - __slots__ = "_override", "_providers", "_selector" + __slots__ = "_override", "_providers", "_selected_provider", "_selector" def __init__( self, selector: typing.Callable[[], str] | AbstractProvider[str] | str, **providers: AbstractProvider[T_co] @@ -67,34 +70,63 @@ def my_selector(): super().__init__() self._selector: typing.Final[typing.Callable[[], str] | AbstractProvider[str] | str] = selector self._providers: typing.Final = providers + self._selected_provider: typing.Final[ContextVar[AbstractProvider[T_co] | Unset]] = ContextVar( + f"selector-{id(self)}", + default=UNSET, + ) + + def _register_arguments(self) -> None: + if not self._mark_arguments_registered(): + return + self._register((self._selector,)) + + def _deregister_arguments(self) -> None: + self._deregister((self._selector,)) + self._reset_arguments_registration() + + @contextmanager + def _pin_selected_provider(self, provider: AbstractProvider[T_co]) -> typing.Iterator[None]: + token = self._selected_provider.set(provider) + try: + yield + finally: + self._selected_provider.reset(token) @override async def resolve(self) -> T_co: if is_set(self._override): return typing.cast(T_co, self._override) + return await (await self._select_provider()).resolve() + + @override + def resolve_sync(self) -> T_co: + if is_set(self._override): + return typing.cast(T_co, self._override) + return self._select_provider_sync().resolve_sync() + + async def _select_provider(self) -> AbstractProvider[T_co]: + selected_provider = self._selected_provider.get() + if is_set(selected_provider): + return selected_provider if isinstance(self._selector, AbstractProvider): selected_key = await self._selector.resolve() else: selected_key = self._get_selected_key() - self._validate_key(selected_key) + return self._providers[selected_key] - return await self._providers[selected_key].resolve() - - @override - def resolve_sync(self) -> T_co: - if is_set(self._override): - return typing.cast(T_co, self._override) + def _select_provider_sync(self) -> AbstractProvider[T_co]: + selected_provider = self._selected_provider.get() + if is_set(selected_provider): + return selected_provider if isinstance(self._selector, AbstractProvider): selected_key = self._selector.resolve_sync() else: selected_key = self._get_selected_key() - self._validate_key(selected_key) - - return self._providers[selected_key].resolve_sync() + return self._providers[selected_key] def _get_selected_key(self) -> str: if callable(self._selector):