diff --git a/.gitignore b/.gitignore index 246b803..8f17a87 100644 --- a/.gitignore +++ b/.gitignore @@ -21,3 +21,4 @@ dist/ .venv uv.lock .agents +site diff --git a/tests/providers/test_base.py b/tests/providers/test_base.py index 7a0601c..28c972b 100644 --- a/tests/providers/test_base.py +++ b/tests/providers/test_base.py @@ -88,6 +88,25 @@ def test_register_with_mixed_items() -> None: assert parent in child_1._children, "Expected child_1._children to contain parent" +def test_get_resolution_dependencies_registers_provider_arguments() -> None: + dependency = DummyProvider() + provider = Singleton(lambda value: value, dependency) + + resolution_dependencies = provider.get_resolution_dependencies() + + assert resolution_dependencies == frozenset({dependency}) + assert isinstance(resolution_dependencies, frozenset) + + +async def test_default_resolution_contexts_have_no_runtime_dependencies() -> None: + provider = DummyProvider() + + async with provider.resolution_context() as async_dependencies: + assert async_dependencies == () + with provider.resolution_context_sync() as sync_dependencies: + assert sync_dependencies == () + + def test_invalidate_scope_init_order_handles_duplicate_descendants() -> None: root = DummyProvider() left = DummyProvider() @@ -100,13 +119,11 @@ def test_invalidate_scope_init_order_handles_duplicate_descendants() -> None: right.add_child_provider(shared) for provider in (root, left, right, shared): - provider._scope_context_init_order = () provider._scope_init_order = () root._invalidate_scope_init_order() for provider in (root, left, right, shared): - assert provider._scope_context_init_order is None assert provider._scope_init_order is None diff --git a/tests/providers/test_selector.py b/tests/providers/test_selector.py index fe7b784..4246200 100644 --- a/tests/providers/test_selector.py +++ b/tests/providers/test_selector.py @@ -140,6 +140,45 @@ async def test_selector_with_provider_selector_async() -> None: assert (await StringProviderSelectorContainer.selector.resolve()) == "Provider 1" +def test_selector_exposes_only_its_key_provider_as_static_dependency() -> 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) + + assert selector.get_resolution_dependencies() == frozenset({selector_key}) + assert selector.get_resolution_dependencies() == frozenset({selector_key}) + + selector._deregister_arguments() + + assert selector.get_resolution_dependencies() == frozenset({selector_key}) + + +def test_selector_resolution_context_exposes_and_pins_selected_provider() -> None: + expected_selection_count = 2 + selected_key = "one" + selection_count = 0 + one = providers.Object("one") + two = providers.Object("two") + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return selected_key + + selector = providers.Selector(_select, one=one, two=two) + + with selector.resolution_context_sync() as dependencies: + assert dependencies == (one,) + selected_key = "two" + assert selector.resolve_sync() == "one" + + assert selector.resolve_sync() == "two" + assert selection_count == expected_selection_count + + 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..86711d1 100644 --- a/tests/test_injection.py +++ b/tests/test_injection.py @@ -7,6 +7,7 @@ from unittest.mock import Mock import pytest +from typing_extensions import override from tests import container from that_depends import ( @@ -42,6 +43,87 @@ def _sync_creator() -> typing.Iterator[int]: yield 1 +class _TestDynamicProvider(providers.AbstractProvider[str]): + """A third-party-style provider with dependencies chosen at resolution time.""" + + def __init__( + self, + name: str, + events: list[str], + *, + static_dependencies: typing.Collection[providers.AbstractProvider[typing.Any]] = (), + delegate: providers.AbstractProvider[str] | None = None, + ) -> None: + super().__init__() + self._name = name + self._events = events + self._static_dependencies = tuple(static_dependencies) + self._runtime_dependencies: tuple[providers.AbstractProvider[typing.Any], ...] = ( + (delegate,) if delegate is not None else () + ) + self._delegate = delegate + self._active = False + + def set_runtime_dependencies(self, *dependencies: providers.AbstractProvider[typing.Any]) -> None: + self._runtime_dependencies = dependencies + + @override + def get_resolution_dependencies(self) -> typing.Collection[providers.AbstractProvider[typing.Any]]: + return self._static_dependencies + + @asynccontextmanager + @override + async def resolution_context( + self, + ) -> typing.AsyncIterator[typing.Collection[providers.AbstractProvider[typing.Any]]]: + self._events.append(f"enter:{self._name}") + self._active = True + try: + yield self._runtime_dependencies + finally: + self._active = False + self._events.append(f"exit:{self._name}") + + @contextmanager + @override + def resolution_context_sync( + self, + ) -> typing.Iterator[typing.Collection[providers.AbstractProvider[typing.Any]]]: + self._events.append(f"enter:{self._name}") + self._active = True + try: + yield self._runtime_dependencies + finally: + self._active = False + self._events.append(f"exit:{self._name}") + + @override + async def resolve(self) -> str: + assert self._active + self._events.append(f"resolve:{self._name}") + static_value = await self._resolve_static_dependency() + delegate_value = await self._delegate.resolve() if self._delegate is not None else self._name + return f"{static_value}:{delegate_value}" if static_value is not None else delegate_value + + @override + def resolve_sync(self) -> str: + assert self._active + self._events.append(f"resolve:{self._name}") + static_value = self._resolve_static_dependency_sync() + delegate_value = self._delegate.resolve_sync() if self._delegate is not None else self._name + return f"{static_value}:{delegate_value}" if static_value is not None else delegate_value + + async def _resolve_static_dependency(self) -> typing.Any: # noqa: ANN401 + if not self._static_dependencies: + return None + return await next(iter(self._static_dependencies)).resolve() + + def _resolve_static_dependency_sync(self) -> typing.Any: # noqa: ANN401 + if not self._static_dependencies: + return None + return next(iter(self._static_dependencies)).resolve_sync() + + @inject async def test_injection( fixture_one: int, @@ -71,6 +153,473 @@ async def inner( await inner(True, arg2=container.SimpleFactory(dep1="1", dep2=2)) +def test_custom_dynamic_provider_prepares_static_and_runtime_dependencies_sync() -> None: + events: list[str] = [] + + def _resource(value: str) -> typing.Iterator[str]: + events.append(f"resource-enter:{value}") + try: + yield value + finally: + events.append(f"resource-exit:{value}") + + static = providers.ContextResource(_resource, "static").with_config(scope=ContextScopes.INJECT) + runtime = providers.ContextResource(_resource, "runtime").with_config(scope=ContextScopes.INJECT) + dynamic = _TestDynamicProvider("dynamic", events, static_dependencies=(static,), delegate=runtime) + + @inject + def _injected(value: str = Provide[dynamic]) -> str: + events.append("body") + assert static.resolve_sync() == "static" + assert runtime.resolve_sync() == "runtime" + return value + + assert _injected() == "static:runtime" + assert events == [ + "enter:dynamic", + "resolve:dynamic", + "resource-enter:static", + "resource-enter:runtime", + "exit:dynamic", + "body", + "resource-exit:runtime", + "resource-exit:static", + ] + + +async def test_custom_dynamic_provider_prepares_static_and_runtime_dependencies_async() -> None: + events: list[str] = [] + + async def _resource(value: str) -> typing.AsyncIterator[str]: + events.append(f"resource-enter:{value}") + try: + yield value + finally: + events.append(f"resource-exit:{value}") + + static = providers.ContextResource(_resource, "static").with_config(scope=ContextScopes.INJECT) + runtime = providers.ContextResource(_resource, "runtime").with_config(scope=ContextScopes.INJECT) + dynamic = _TestDynamicProvider("dynamic", events, static_dependencies=(static,), delegate=runtime) + + @inject + async def _injected(value: str = Provide[dynamic]) -> str: + events.append("body") + assert await static.resolve() == "static" + assert await runtime.resolve() == "runtime" + return value + + assert await _injected() == "static:runtime" + assert events == [ + "enter:dynamic", + "resolve:dynamic", + "resource-enter:static", + "resource-enter:runtime", + "exit:dynamic", + "body", + "resource-exit:runtime", + "resource-exit:static", + ] + + +def test_nested_custom_dynamic_providers_remain_active_until_root_resolution() -> None: + events: list[str] = [] + leaf = providers.Object("value") + inner = _TestDynamicProvider("inner", events, delegate=leaf) + outer = _TestDynamicProvider("outer", events, delegate=inner) + root = providers.Factory(lambda value: value, outer.cast) + + @inject + def _injected(value: str = Provide[root]) -> str: + return value + + assert _injected() == "value" + assert events == [ + "enter:outer", + "enter:inner", + "resolve:outer", + "resolve:inner", + "exit:inner", + "exit:outer", + ] + + +def test_dynamic_provider_traversal_handles_duplicate_and_cyclic_dependencies() -> None: + events: list[str] = [] + root = _TestDynamicProvider("root", events) + dependency = _TestDynamicProvider("dependency", events) + root.set_runtime_dependencies(root, dependency, dependency) + dependency.set_runtime_dependencies(root) + + @inject + def _injected(value: str = Provide[root]) -> str: + return value + + assert _injected() == "root" + assert events == [ + "enter:root", + "enter:dependency", + "resolve:root", + "exit:dependency", + "exit:root", + ] + + +def test_selector_injection_prepares_only_active_branch_and_reselects_in_body_sync() -> None: + expected_selection_count = 2 + events: list[str] = [] + selection_count = 0 + + def _resource(name: str) -> typing.Iterator[str]: + events.append(f"enter:{name}") + try: + yield name + finally: + events.append(f"exit:{name}") + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return "selected" + + selected = providers.ContextResource(_resource, "selected").with_config(scope=ContextScopes.INJECT) + unselected = providers.ContextResource(_resource, "unselected").with_config(scope=ContextScopes.INJECT) + selector = providers.Selector(_select, selected=selected, unselected=unselected) + + @inject + def _injected(value: str = Provide[selector]) -> str: + assert selection_count == 1 + assert selector.resolve_sync() == "selected" + with pytest.raises(RuntimeError): + unselected.resolve_sync() + return value + + assert _injected() == "selected" + assert selection_count == expected_selection_count + assert events == ["enter:selected", "exit:selected"] + + +async def test_selector_injection_prepares_only_active_branch_and_reselects_in_body_async() -> None: + expected_selection_count = 2 + events: list[str] = [] + selection_count = 0 + + async def _resource(name: str) -> typing.AsyncIterator[str]: + events.append(f"enter:{name}") + try: + yield name + finally: + events.append(f"exit:{name}") + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return "selected" + + selected = providers.ContextResource(_resource, "selected").with_config(scope=ContextScopes.INJECT) + unselected = providers.ContextResource(_resource, "unselected").with_config(scope=ContextScopes.INJECT) + selector = providers.Selector(_select, selected=selected, unselected=unselected) + + @inject + async def _injected(value: str = Provide[selector]) -> str: + assert selection_count == 1 + assert await selector.resolve() == "selected" + with pytest.raises(RuntimeError): + await unselected.resolve() + return value + + assert await _injected() == "selected" + assert selection_count == expected_selection_count + assert events == ["enter:selected", "exit:selected"] + + +def test_nested_selectors_prepare_the_innermost_selected_branch() -> None: + selection_events: list[str] = [] + + def _resource() -> typing.Iterator[str]: + yield "value" + + def _select_outer() -> str: + selection_events.append("outer") + return "inner" + + def _select_inner() -> str: + selection_events.append("inner") + return "resource" + + resource = providers.ContextResource(_resource).with_config(scope=ContextScopes.INJECT) + inner = providers.Selector(_select_inner, resource=resource) + outer = providers.Selector(_select_outer, inner=inner, unused=providers.Object("unused")) + + @inject + def _injected(value: str = Provide[outer]) -> str: + return value + + assert _injected() == "value" + assert selection_events == ["outer", "inner"] + + +def test_selector_selection_state_is_reset_after_resolution_error_sync() -> None: + expected_selection_count = 2 + selected_key = "failing" + selection_count = 0 + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return selected_key + + def _fail() -> typing.NoReturn: + msg = "resolution failed" + raise RuntimeError(msg) + + selector = providers.Selector[str]( + _select, + failing=providers.Factory(_fail), + successful=providers.Object("value"), + ) + + @inject + def _injected(value: str = Provide[selector]) -> str: + return value + + with pytest.raises(RuntimeError, match="resolution failed"): + _injected() + + selected_key = "successful" + assert _injected() == "value" + assert selection_count == expected_selection_count + + +async def test_selector_selection_state_is_reset_after_resolution_error_async() -> None: + expected_selection_count = 2 + selected_key = "failing" + selection_count = 0 + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return selected_key + + async def _fail() -> typing.NoReturn: + msg = "resolution failed" + raise RuntimeError(msg) + + selector = providers.Selector[str]( + _select, + failing=providers.AsyncFactory(_fail), + successful=providers.Object("value"), + ) + + @inject + async def _injected(value: str = Provide[selector]) -> str: + return value + + with pytest.raises(RuntimeError, match="resolution failed"): + await _injected() + + selected_key = "successful" + assert await _injected() == "value" + assert selection_count == expected_selection_count + + +def test_overridden_selector_does_not_prepare_candidate_branch_sync() -> None: + override_value = 2 + candidate = providers.ContextResource(_sync_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector("candidate", candidate=candidate) + selector.override_sync(override_value) + + @inject + def _injected(value: int = Provide[selector]) -> int: + with pytest.raises(RuntimeError): + candidate.resolve_sync() + return value + + try: + assert _injected() == override_value + finally: + selector.reset_override_sync() + + +async def test_overridden_selector_does_not_prepare_candidate_branch_async() -> None: + override_value = 2 + candidate = providers.ContextResource(_async_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector("candidate", candidate=candidate) + selector.override_sync(override_value) + + @inject + async def _injected(value: int = Provide[selector]) -> int: + with pytest.raises(RuntimeError): + await candidate.resolve() + return value + + try: + assert await _injected() == override_value + finally: + selector.reset_override_sync() + + +async def test_dynamic_provider_async_traversal_handles_duplicate_and_cyclic_dependencies() -> None: + events: list[str] = [] + root = _TestDynamicProvider("root", events) + dependency = _TestDynamicProvider("dependency", events) + root.set_runtime_dependencies(root, dependency, dependency) + dependency.set_runtime_dependencies(root) + + @inject + async def _injected(value: str = Provide[root]) -> str: + return value + + assert await _injected() == "root" + assert events == [ + "enter:root", + "enter:dependency", + "resolve:root", + "exit:dependency", + "exit:root", + ] + + +def test_selector_branch_preparation_supports_every_provider_lookup_surface() -> None: + events: list[str] = [] + + def _resource() -> typing.Iterator[str]: + events.append("enter") + try: + yield "value" + finally: + events.append("exit") + + selected = providers.ContextResource(_resource).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector("selected", selected=selected).bind(str) + + class _ResolutionSurfaceContainer(BaseContainer): + dynamic = selector + + @inject + def _direct(value: str = Provide[selector]) -> str: + return value + + @inject + def _string(value: str = Provide["_ResolutionSurfaceContainer.dynamic"]) -> str: + return value + + @inject(container=_ResolutionSurfaceContainer) + def _typed(value: str = Provide()) -> str: + return value + + assert _direct() == "value" + assert _string() == "value" + assert _typed() == "value" + assert events == ["enter", "exit", "enter", "exit", "enter", "exit"] + + +def test_static_resolution_fast_path_handles_dependency_cycle() -> None: + root = providers.Object("value") + dependency = providers.Object("unused") + root._register((dependency,)) + dependency._register((root,)) + + @inject + def _injected(value: str = Provide[root]) -> str: + return value + + assert _injected() == "value" + + +def test_sync_generator_pins_dynamic_selection_without_resource_stack() -> None: + selection_count = 0 + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return "one" if selection_count == 1 else "two" + + selector = providers.Selector(_select, one=providers.Object("one"), two=providers.Object("two")) + + @inject + def _injected(value: str = Provide[selector]) -> typing.Generator[str, None, None]: + yield value + + assert next(_injected()) == "one" + assert selection_count == 1 + + +async def test_async_generator_pins_dynamic_selection_without_resource_stack() -> None: + selection_count = 0 + + def _select() -> str: + nonlocal selection_count + selection_count += 1 + return "one" if selection_count == 1 else "two" + + selector = providers.Selector(_select, one=providers.Object("one"), two=providers.Object("two")) + + @inject + async def _injected(value: str = Provide[selector]) -> typing.AsyncGenerator[str, None]: + yield value + + assert await anext(_injected()) == "one" + assert selection_count == 1 + + +def test_sync_generator_rejects_only_selected_context_resource_branch() -> None: + selected_key = "plain" + resource = providers.ContextResource(_sync_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector( + lambda: selected_key, + plain=providers.Object(1), + resource=resource, + ) + + @inject + def _injected(value: int = Provide[selector]) -> typing.Generator[int, None, None]: + yield value + + assert next(_injected()) == 1 + + selected_key = "resource" + with pytest.raises(ContextProviderError): + next(_injected()) + + +async def test_async_generator_rejects_only_selected_context_resource_branch() -> None: + selected_key = "plain" + resource = providers.ContextResource(_async_creator).with_config(scope=ContextScopes.INJECT) + selector = providers.Selector( + lambda: selected_key, + plain=providers.Object(1), + resource=resource, + ) + + @inject + async def _injected(value: int = Provide[selector]) -> typing.AsyncGenerator[int, None]: + yield value + + assert await anext(_injected()) == 1 + + selected_key = "resource" + with pytest.raises(ContextProviderError): + await anext(_injected()) + + +def test_selector_branch_preserves_context_resource_scope_filtering() -> None: + def _resource() -> typing.Iterator[str]: + yield "value" + + resource = providers.ContextResource(_resource).with_config(scope=ContextScopes.REQUEST) + selector = providers.Selector("resource", resource=resource) + + @inject(scope=ContextScopes.APP) + def _injected(value: str = Provide[selector]) -> str: + return value + + with pytest.raises(RuntimeError): + _injected() + + with resource.context_sync(force=True): + assert _injected() == "value" + + def test_sync_injection_stack_closes_entered_context_managers() -> None: events: list[str] = [] @@ -98,9 +647,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..218ba36 100644 --- a/that_depends/injection.py +++ b/that_depends/injection.py @@ -3,7 +3,7 @@ import re import typing import warnings -from contextlib import AsyncExitStack +from contextlib import AsyncExitStack, ExitStack from types import TracebackType from typing_extensions import Self @@ -12,7 +12,7 @@ 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 class ContextProviderError(Exception): @@ -28,10 +28,18 @@ class ContextProviderError(Exception): ) +class _RuntimeContextResources: + """Mark a provider graph whose context resources require runtime discovery.""" + + +_RUNTIME_CONTEXT_RESOURCES = _RuntimeContextResources() +_ContextResources = tuple[ContextResource[typing.Any], ...] | _RuntimeContextResources + + class _DirectInjectionParameter(typing.NamedTuple): field_name: str provider: AbstractProvider[typing.Any] - scope_context_init_order: tuple[AbstractProvider[typing.Any], ...] + context_resources: _ContextResources class _StringInjectionParameter(typing.NamedTuple): @@ -51,6 +59,19 @@ class _InjectionPlan(typing.NamedTuple): typed_parameters: tuple[_TypedInjectionParameter, ...] +class _ProviderVisits(typing.NamedTuple): + """Track traversal state shared by one provider resolution. + + Attributes: + traversed: Providers whose dependency contexts have already been visited. + initialized_contexts: Context resources already entered by the injection call. + + """ + + traversed: set[AbstractProvider[typing.Any]] + initialized_contexts: set[AbstractProvider[typing.Any]] + + class _SyncInjectionStack: __slots__ = ("_exit_states",) @@ -94,6 +115,56 @@ def close(self) -> None: self._context_manager.__exit__(None, None, None) +@functools.cache +def _get_static_context_resources( + provider: AbstractProvider[typing.Any], +) -> _ContextResources: + """Collect context resources from a provider's static dependency graph. + + The result is cached with the injection plan. A private marker signals that at + least one provider overrides a resolution-context hook, so injection must walk + the graph at runtime to discover dynamic dependencies. + + Args: + provider: Root provider whose dependency graph should be inspected. + + Returns: + The statically reachable context resources, or a marker requesting runtime + traversal. + + """ + resources: list[ContextResource[typing.Any]] = [] + visited: set[AbstractProvider[typing.Any]] = set() + + def _visit(dependency: AbstractProvider[typing.Any]) -> bool: + """Visit a dependency while the graph remains statically discoverable. + + Args: + dependency: Provider whose static dependencies should be inspected. + + Returns: + Whether the dependency and all of its descendants use the default + resolution-context hooks. + + """ + if dependency in visited: + return True + visited.add(dependency) + + if ( + type(dependency).resolution_context is not AbstractProvider.resolution_context + or type(dependency).resolution_context_sync is not AbstractProvider.resolution_context_sync + ): + return False + if not all(_visit(parent) for parent in dependency.get_resolution_dependencies()): + return False + if isinstance(dependency, ContextResource): + resources.append(dependency) + return True + + return tuple(resources) if _visit(provider) else _RUNTIME_CONTEXT_RESOURCES + + @functools.cache def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionPlan: signature = inspect.signature(func) @@ -125,7 +196,7 @@ def _build_injection_plan(func: typing.Callable[..., typing.Any]) -> _InjectionP _DirectInjectionParameter( field_name, default, - default._get_scope_context_init_order(), # noqa: SLF001 + _get_static_context_resources(default), ) ) elif isinstance(default, _Provide): @@ -295,9 +366,18 @@ 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, + context_resources = direct_parameter.context_resources + if context_resources: + if context_resources is _RUNTIME_CONTEXT_RESOURCES: + kwargs[direct_parameter.field_name] = await _resolve_provider_with_scope_async( + direct_parameter.provider, + scope, + stack, + context_providers, + ) + continue + await _prepare_static_context_resources_async( + typing.cast(tuple[ContextResource[typing.Any], ...], context_resources), scope, stack, context_providers, @@ -346,9 +426,18 @@ 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, + context_resources = direct_parameter.context_resources + if context_resources: + if context_resources is _RUNTIME_CONTEXT_RESOURCES: + kwargs[direct_parameter.field_name] = _resolve_provider_with_scope_sync( + direct_parameter.provider, + scope, + stack, + context_providers, + ) + continue + _prepare_static_context_resources_sync( + typing.cast(tuple[ContextResource[typing.Any], ...], context_resources), scope, stack, context_providers, @@ -406,13 +495,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] @@ -448,50 +530,141 @@ async def _resolve_provider_with_scope_async( stack: AsyncExitStack | None, providers: set[AbstractProvider[typing.Any]], ) -> T: - """Resolve a provider with given scope and stack. + """Resolve a provider and initialize its matching asynchronous resources. - Use `stack=None` to ensure ContextResource providers are not allowed. + Static graphs use their cached resource list. Graphs with runtime dependencies + are traversed while their resolution contexts remain active. Passing ``None`` + as the stack explicitly disallows context-resource initialization. Args: - provider: provider to resolve. - scope: scope to resolve provider in. - stack: stack to use for context resources. - providers: providers traversed. + provider: Provider to resolve. + scope: Scope in which matching context resources should be initialized. + stack: Stack that owns initialized context resources, or ``None`` to reject + resources that require initialization. + providers: Context resources already initialized by the injection call. Returns: - resolved value for the provider. + The value resolved by the provider. Raises: - ContextProviderError: if the stack is None. + ContextProviderError: If a matching context resource requires initialization + but no stack was supplied. """ - 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) - return await provider.resolve() + static_resources = _get_static_context_resources(provider) + if static_resources is not _RUNTIME_CONTEXT_RESOURCES: + if static_resources: + await _prepare_static_context_resources_async( + typing.cast(tuple[ContextResource[typing.Any], ...], static_resources), + scope, + stack, + providers, + ) + return await provider.resolve() + async with AsyncExitStack() as resolution_stack: + visits = _ProviderVisits(set(), providers) + await _prepare_provider_contexts_async(provider, scope, stack, resolution_stack, visits) + return await provider.resolve() -async def _setup_scope_contexts_async( - scope_init_order: tuple[AbstractProvider[typing.Any], ...], + +async def _prepare_static_context_resources_async( + resources: tuple[ContextResource[typing.Any], ...], scope: ContextScope | None, stack: AsyncExitStack | None, - providers: set[AbstractProvider[typing.Any]], + initialized: set[AbstractProvider[typing.Any]], ) -> None: - if not scope: + """Enter statically discovered asynchronous context resources once. + + Args: + resources: Context resources reachable from the root provider. + scope: Scope in which matching resources should be initialized. + stack: Stack that owns initialized resources, or ``None`` to reject them. + initialized: Context resources already entered by the injection call. + + Raises: + ContextProviderError: If a matching resource requires initialization but no + stack was supplied. + + """ + if scope is None: return - for provider in scope_init_order: - if provider in providers: + for resource in resources: + if resource in initialized or resource._scope not in (ContextScopes.ANY, scope): # noqa: SLF001 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)) + if stack is None: + msg = ( + f"No stack exists, cannot initialize context for {resource} using scope {scope}.\n" + f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." + ) + raise ContextProviderError(msg) + initialized.add(resource) + await stack.enter_async_context(resource.context_async(force=True)) + + +async def _prepare_provider_contexts_async( + provider: AbstractProvider[typing.Any], + scope: ContextScope | None, + resource_stack: AsyncExitStack | None, + resolution_stack: AsyncExitStack, + visits: _ProviderVisits, +) -> None: + """Prepare one provider's asynchronous static and runtime dependencies. + + Dependencies are visited before the provider itself. Resolution contexts are + kept open on ``resolution_stack`` until the root provider has resolved, while + context resources live on ``resource_stack`` for the entire injection call. + + Args: + provider: Provider whose dependency contexts should be prepared. + scope: Scope in which matching context resources should be initialized. + resource_stack: Stack that owns context resources, or ``None`` to reject + resources that require initialization. + resolution_stack: Stack that owns provider resolution contexts. + visits: Traversal and resource-initialization state for this resolution. + + Raises: + ContextProviderError: If a matching context resource requires initialization + but no resource stack was supplied. + + """ + if provider in visits.traversed: + return + visits.traversed.add(provider) + + for dependency in provider.get_resolution_dependencies(): + await _prepare_provider_contexts_async( + dependency, + scope, + resource_stack, + resolution_stack, + visits, + ) + + runtime_dependencies = await resolution_stack.enter_async_context(provider.resolution_context()) + for dependency in runtime_dependencies: + await _prepare_provider_contexts_async( + dependency, + scope, + resource_stack, + resolution_stack, + visits, + ) + + if ( + scope is not None + and isinstance(provider, ContextResource) + and provider.get_scope() in (ContextScopes.ANY, scope) + and provider not in visits.initialized_contexts + ): + if resource_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) + visits.initialized_contexts.add(provider) + await resource_stack.enter_async_context(provider.context_async(force=True)) def _resolve_provider_with_scope_sync( @@ -500,34 +673,142 @@ 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) - return provider.resolve_sync() + """Resolve a provider and initialize its matching synchronous resources. + Static graphs use their cached resource list. Graphs with runtime dependencies + are traversed while their resolution contexts remain active. Passing ``None`` + as the stack explicitly disallows context-resource initialization. -def _setup_scope_contexts_sync( - scope_init_order: tuple[AbstractProvider[typing.Any], ...], + Args: + provider: Provider to resolve. + scope: Scope in which matching context resources should be initialized. + stack: Stack that owns initialized context resources, or ``None`` to reject + resources that require initialization. + providers: Context resources already initialized by the injection call. + + Returns: + The value resolved by the provider. + + Raises: + ContextProviderError: If a matching context resource requires initialization + but no stack was supplied. + + """ + static_resources = _get_static_context_resources(provider) + if static_resources is not _RUNTIME_CONTEXT_RESOURCES: + if static_resources: + _prepare_static_context_resources_sync( + typing.cast(tuple[ContextResource[typing.Any], ...], static_resources), + scope, + stack, + providers, + ) + return provider.resolve_sync() + + with ExitStack() as resolution_stack: + visits = _ProviderVisits(set(), providers) + _prepare_provider_contexts_sync(provider, scope, stack, resolution_stack, visits) + return provider.resolve_sync() + + +def _prepare_static_context_resources_sync( + resources: tuple[ContextResource[typing.Any], ...], scope: ContextScope | None, stack: _SyncInjectionStack | None, - providers: set[AbstractProvider[typing.Any]], + initialized: set[AbstractProvider[typing.Any]], ) -> None: - if not scope: + """Enter statically discovered synchronous context resources once. + + Args: + resources: Context resources reachable from the root provider. + scope: Scope in which matching resources should be initialized. + stack: Stack that owns initialized resources, or ``None`` to reject them. + initialized: Context resources already entered by the injection call. + + Raises: + ContextProviderError: If a matching resource requires initialization but no + stack was supplied. + + """ + if scope is None: return - for provider in scope_init_order: - if provider in providers: + for resource in resources: + if resource in initialized or resource._scope not in (ContextScopes.ANY, scope): # noqa: SLF001 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) + if stack is None: + msg = ( + f"No stack exists, cannot initialize context for {resource} using scope {scope}.\n" + f"Note: @inject cannot initialize context for ContextResources when wrapping a generator." + ) + raise ContextProviderError(msg) + initialized.add(resource) + _, exit_state = resource._enter_injection_context_sync(force=True) # noqa: SLF001 + stack.push_exit_state(exit_state) + + +def _prepare_provider_contexts_sync( + provider: AbstractProvider[typing.Any], + scope: ContextScope | None, + resource_stack: _SyncInjectionStack | None, + resolution_stack: ExitStack, + visits: _ProviderVisits, +) -> None: + """Prepare one provider's synchronous static and runtime dependencies. + + Dependencies are visited before the provider itself. Resolution contexts are + kept open on ``resolution_stack`` until the root provider has resolved, while + context resources live on ``resource_stack`` for the entire injection call. + + Args: + provider: Provider whose dependency contexts should be prepared. + scope: Scope in which matching context resources should be initialized. + resource_stack: Stack that owns context resources, or ``None`` to reject + resources that require initialization. + resolution_stack: Stack that owns provider resolution contexts. + visits: Traversal and resource-initialization state for this resolution. + + Raises: + ContextProviderError: If a matching context resource requires initialization + but no resource stack was supplied. + + """ + if provider in visits.traversed: + return + visits.traversed.add(provider) + + for dependency in provider.get_resolution_dependencies(): + _prepare_provider_contexts_sync( + dependency, + scope, + resource_stack, + resolution_stack, + visits, + ) + + runtime_dependencies = resolution_stack.enter_context(provider.resolution_context_sync()) + for dependency in runtime_dependencies: + _prepare_provider_contexts_sync( + dependency, + scope, + resource_stack, + resolution_stack, + visits, + ) + + if ( + scope is not None + and isinstance(provider, ContextResource) + and provider.get_scope() in (ContextScopes.ANY, scope) + and provider not in visits.initialized_contexts + ): + if resource_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) + visits.initialized_contexts.add(provider) + resource_stack.enter_context(provider.context_sync(force=True)) class StringProviderDefinition: diff --git a/that_depends/providers/base.py b/that_depends/providers/base.py index 31f5982..6da7e2b 100644 --- a/that_depends/providers/base.py +++ b/that_depends/providers/base.py @@ -95,7 +95,6 @@ def __init__(self) -> None: self._children: set[AbstractProvider[typing.Any]] = set() self._parents: set[AbstractProvider[typing.Any]] = set() self._is_context_resource = False - self._scope_context_init_order: tuple[AbstractProvider[typing.Any], ...] | None = None self._scope_init_order: tuple[AbstractProvider[typing.Any], ...] | None = None self._override: typing.Any = UNSET self._bindings: set[type] = set() @@ -166,7 +165,6 @@ def _invalidate_scope_init_order(self) -> None: if provider in visited: continue visited.add(provider) - provider._scope_context_init_order = None # noqa: SLF001 provider._scope_init_order = None # noqa: SLF001 stack.extend(provider._children) # noqa: SLF001 @@ -192,28 +190,6 @@ def _get_scope_init_order(self) -> tuple["AbstractProvider[typing.Any]", ...]: self._scope_init_order = tuple(ordered) return self._scope_init_order - def _get_scope_context_init_order(self) -> tuple["AbstractProvider[typing.Any]", ...]: - if self._scope_context_init_order is not None: - return self._scope_context_init_order - - if isinstance(self, ProviderWithArguments): - self._register_arguments() - - ordered: list[AbstractProvider[typing.Any]] = [] - seen: set[AbstractProvider[typing.Any]] = set() - - for parent in self._parents: - for ancestor in parent._get_scope_context_init_order(): # noqa: SLF001 - if ancestor not in seen: - seen.add(ancestor) - ordered.append(ancestor) - - if self._is_context_resource and self not in seen: - ordered.append(self) - - self._scope_context_init_order = tuple(ordered) - return self._scope_context_init_order - def add_child_provider(self, provider: "AbstractProvider[typing.Any]") -> None: """Add a child provider to the current provider. @@ -262,6 +238,58 @@ def __getattr__(self, attr_name: str) -> typing.Any: # noqa: ANN401 raise AttributeError(msg) return AttrGetter(provider=self, attr_name=attr_name) + def get_resolution_dependencies(self) -> typing.Collection["AbstractProvider[typing.Any]"]: + """Return providers that must be prepared before resolving this provider. + + Injection evaluates these static dependencies before entering this provider's + resolution context. Providers with arguments are registered lazily, and the + returned collection is an immutable snapshot with no ordering guarantee. + + Dynamic providers can expose dependencies selected at runtime from + :meth:`resolution_context` and :meth:`resolution_context_sync` instead. + + Returns: + An immutable snapshot of the provider's static dependencies. + + """ + if isinstance(self, ProviderWithArguments): + self._register_arguments() + return frozenset(self._parents) + + @asynccontextmanager + async def resolution_context( + self, + ) -> typing.AsyncIterator[typing.Collection["AbstractProvider[typing.Any]"]]: + """Yield dependencies known only while resolving this provider asynchronously. + + Injection prepares the yielded providers and keeps this context active until + the root provider has resolved. The default implementation has no runtime + dependencies. Custom providers should yield a read-only collection and must + not rely on its iteration order. + + Yields: + The dependencies discovered for the current asynchronous resolution. + + """ + yield () + + @contextmanager + def resolution_context_sync( + self, + ) -> typing.Iterator[typing.Collection["AbstractProvider[typing.Any]"]]: + """Yield dependencies known only while resolving this provider synchronously. + + Injection prepares the yielded providers and keeps this context active until + the root provider has resolved. The default implementation has no runtime + dependencies. Custom providers should yield a read-only collection and must + not rely on its iteration order. + + Yields: + The dependencies discovered for the current synchronous resolution. + + """ + yield () + @abc.abstractmethod async def resolve(self) -> T_co: """Resolve dependency asynchronously.""" diff --git a/that_depends/providers/selector.py b/that_depends/providers/selector.py index 9ec2c48..5e5277a 100644 --- a/that_depends/providers/selector.py +++ b/that_depends/providers/selector.py @@ -1,22 +1,27 @@ """Selection based providers.""" import typing +from contextlib import asynccontextmanager, 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 named providers at runtime. The provider key is determined by a - user-supplied selector function. + user-supplied selector function. During injection, only the selected + provider branch is prepared. That selection remains stable for one root + provider resolution and is evaluated again by later resolutions. Examples: ```python @@ -35,7 +40,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 +72,150 @@ 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: + """Register the provider-valued selector as a static dependency. + + Registration is idempotent because providers attach their arguments lazily + when the dependency graph is first inspected. + """ + if not self._mark_arguments_registered(): + return + self._register((self._selector,)) + + def _deregister_arguments(self) -> None: + """Detach the provider-valued selector from this provider's dependency graph.""" + self._deregister((self._selector,)) + self._reset_arguments_registration() + + @contextmanager + def _pin_selected_provider(self, provider: AbstractProvider[T_co]) -> typing.Iterator[None]: + """Keep a selected provider stable for the current resolution context. + + Args: + provider: Provider selected for the active root resolution. + + Yields: + Control while the provider is pinned in the current context. + + """ + token = self._selected_provider.set(provider) + try: + yield + finally: + self._selected_provider.reset(token) + + @asynccontextmanager + @override + async def resolution_context( + self, + ) -> typing.AsyncIterator[typing.Collection[AbstractProvider[typing.Any]]]: + """Expose the selected provider for one asynchronous root resolution. + + Overrides bypass provider selection because resolving the selector returns + the override directly. Otherwise, the selected provider is pinned so the + injection traversal and final resolution use the same branch. + + Yields: + The selected provider, or an empty collection while overridden. + + """ + if is_set(self._override): + yield () + return + + selected_provider = await self._select_provider() + with self._pin_selected_provider(selected_provider): + yield (selected_provider,) + + @contextmanager + @override + def resolution_context_sync( + self, + ) -> typing.Iterator[typing.Collection[AbstractProvider[typing.Any]]]: + """Expose the selected provider for one synchronous root resolution. + + Overrides bypass provider selection because resolving the selector returns + the override directly. Otherwise, the selected provider is pinned so the + injection traversal and final resolution use the same branch. + + Yields: + The selected provider, or an empty collection while overridden. + + """ + if is_set(self._override): + yield () + return + + selected_provider = self._select_provider_sync() + with self._pin_selected_provider(selected_provider): + yield (selected_provider,) @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]: + """Return the provider selected for asynchronous resolution. + + A provider pinned by :meth:`resolution_context` takes precedence over + evaluating the selector again. + + Returns: + The provider associated with the selected key. + + Raises: + TypeError: If the selector is not a supported type. + KeyError: If the selected key has no associated provider. + + """ + 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() + def _select_provider_sync(self) -> AbstractProvider[T_co]: + """Return the provider selected for synchronous resolution. - @override - def resolve_sync(self) -> T_co: - if is_set(self._override): - return typing.cast(T_co, self._override) + A provider pinned by :meth:`resolution_context_sync` takes precedence over + evaluating the selector again. + + Returns: + The provider associated with the selected key. + + Raises: + TypeError: If the selector is not a supported type. + KeyError: If the selected key has no associated provider. + + """ + 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):