diff --git a/tests/providers/test_selector.py b/tests/providers/test_selector.py index fe7b784..865d05a 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 3f0e02a..9fa5aff 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: diff --git a/that_depends/injection.py b/that_depends/injection.py index 15ef581..67ed574 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): @@ -295,14 +296,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: @@ -346,14 +345,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: @@ -406,13 +403,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] @@ -465,33 +455,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( @@ -500,34 +500,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 9ec2c48..975869f 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):