From b01f414b3739ec6b128d16a580ec541a100e6273 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Mon, 5 Oct 2026 12:35:05 +0300 Subject: [PATCH 1/2] refactor: tidy registry, wiring, factory and container internals (#575) --- AGENTS.md | 2 +- modern_di/container.py | 25 ++-- modern_di/dependency_graph.py | 150 ++++++++++---------- modern_di/group.py | 4 +- modern_di/providers/alias.py | 14 +- modern_di/providers/factory.py | 115 ++++++++------- modern_di/registries/cache_registry.py | 25 ++-- modern_di/registries/overrides_registry.py | 2 +- modern_di/registries/providers_registry.py | 41 +++--- modern_di/types_parser.py | 25 +++- modern_di/wiring.py | 35 ++--- tests/providers/test_factory.py | 2 +- tests/registries/test_cache_registry.py | 6 +- tests/registries/test_providers_registry.py | 11 +- tests/test_dependency_graph.py | 25 ++-- tests/test_dependency_graph_contract.py | 12 +- tests/test_runtime_cycle_guard.py | 2 +- tests/test_types_parser.py | 4 +- tests/test_wiring.py | 50 +++---- 19 files changed, 265 insertions(+), 285 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index c111f39c..39ebd1ae 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -46,7 +46,7 @@ Every module under `modern_di/` is named for what it does; read it. What a singl class here, never at the raise site. Submodules split by family; `__init__` re-exports every public name. - `registries/` — `providers_registry` (type → provider, plus the shared plan/resolver memos) and `overrides_registry` are shared tree-wide; `cache_registry` and `context_registry` are per-container. -- `dependency_graph.py` walks `WiringPlan.edges`, so what `validate()` traverses is exactly what +- `dependency_graph.py` walks `WiringPlan.provider_kwargs`, so what `validate()` traverses is exactly what `resolve()` follows. Explicit-stack, never recursive: a caller runs it inside a `RecursionError` handler near CPython's stack limit. - `types.py` — `UNSET` is load-bearing on the resolve path: the miss marker for both the override diff --git a/modern_di/container.py b/modern_di/container.py index 70220fdd..446600b5 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -3,9 +3,8 @@ import threading import typing -from modern_di import exceptions, types +from modern_di import dependency_graph, exceptions, types from modern_di._scope_algebra import next_deeper -from modern_di.dependency_graph import DependencyGraph, build_cycle_error, collect_errors, redirect_hops from modern_di.group import Group from modern_di.providers.abstract import AbstractProvider from modern_di.providers.container_provider import container_provider @@ -20,17 +19,13 @@ def _handle_recursion_error( provider: AbstractProvider[typing.Any], container: "Container", registry: ProvidersRegistry, exc: RecursionError ) -> typing.NoReturn: - """Convert an escaped `RecursionError` to `CircularDependencyError`, or re-raise it unchanged. - - A separate call, not inlined into `resolve_provider`: the coverage tracer re-arms on the - fresh call boundary before this raises. - """ + """Convert an escaped `RecursionError` to `CircularDependencyError`, or re-raise it unchanged.""" if registry.is_validated(): raise exc # validated => acyclic static graph => genuine self-recursion - cycle = DependencyGraph().find_cycle_from(provider, container) + cycle = dependency_graph.find_cycle_from(provider, container) if cycle is None: raise exc - raise build_cycle_error(cycle, container) from exc + raise dependency_graph.build_cycle_error(cycle, container) from exc class Container: @@ -82,9 +77,7 @@ def __init__( self._closed = False self.scope = scope self.parent_container = parent_container - # Ancestors only, never self: a `scope: self` entry is a reference cycle, so no container - # would ever be freed by refcounting. - # SLF001 exempts `self`/`cls` only, so it flags this same-class read; no boundary is crossed. + # Ancestors only: a `scope: self` entry is a reference cycle that refcounting never frees. self._scope_map: dict[enum.IntEnum, typing.Self] = ( {**parent_container._scope_map, parent_container.scope: parent_container} # noqa: SLF001 if parent_container @@ -93,9 +86,7 @@ def __init__( self._cache_registry = CacheRegistry() self._context_registry = ContextRegistry(copy.copy(context) if context is not None else {}) self._providers_registry: ProvidersRegistry - # Inlined rather than a helper: this runs per child build (benchmark `test_g6_build_child_container`). if parent_container: - # SLF001 exempts `self`/`cls` only, so it flags this same-class read; no boundary is crossed. self._lock = parent_container._lock # noqa: SLF001 self._providers_registry = parent_container._providers_registry # noqa: SLF001 else: @@ -168,7 +159,7 @@ def resolve(self, dependency_type: type[types.T]) -> types.T: except STEP_ERRORS as exc: provider = registry.find_provider(dependency_type) if provider is not None: - exc.prepend_step(*redirect_hops(provider, self)) + exc.prepend_step(*dependency_graph.redirect_hops(provider, self)) raise def resolve_dependency(self, dependency: "AbstractProvider[types.T] | type[types.T]") -> types.T: @@ -194,7 +185,7 @@ def resolve_provider(self, provider: "AbstractProvider[types.T]") -> types.T: except RecursionError as exc: _handle_recursion_error(provider, self, registry, exc) except STEP_ERRORS as exc: - exc.prepend_step(*redirect_hops(provider, self)) + exc.prepend_step(*dependency_graph.redirect_hops(provider, self)) raise def validate(self) -> None: @@ -209,7 +200,7 @@ def validate(self) -> None: if reg.is_validated(): return - if errors := collect_errors(self, reg): + if errors := dependency_graph.collect_errors(self, reg): raise exceptions.ValidationFailedError(errors=errors) reg.mark_validated() diff --git a/modern_di/dependency_graph.py b/modern_di/dependency_graph.py index f0c602e4..273e619b 100644 --- a/modern_di/dependency_graph.py +++ b/modern_di/dependency_graph.py @@ -101,89 +101,85 @@ def build_cycle_error( ) -class DependencyGraph: - """Stateless walker over the static provider graph rooted at a container's registry.""" - - def walk( - self, - roots: "typing.Iterable[AbstractProvider[typing.Any]]", - container: "Container", - ) -> "typing.Iterator[Event]": - """Pre-order DFS from each root; bookkeeping is shared across roots, keyed on ``provider_id``.""" - visiting: set[int] = set() - visited: set[int] = set() - for root in roots: - yield from self._walk_from(root, container, visiting, visited) - - def find_cycle_from( - self, - start: "AbstractProvider[typing.Any]", - container: "Container", - ) -> "list[AbstractProvider[typing.Any]] | None": - """Return the first cycle reachable from ``start``, or None when that subgraph is acyclic.""" - for event in self.walk([start], container): - if isinstance(event, Cycle): - return event.providers - return None - - def _walk_from( - self, - start: "AbstractProvider[typing.Any]", - container: "Container", - visiting: set[int], - visited: set[int], - ) -> "typing.Iterator[Event]": - """Explicit-stack DFS from ``start``; skip immediately if already seen.""" - if start.provider_id in visited or start.provider_id in visiting: - return - - path: list[AbstractProvider[typing.Any]] = [] - stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = [] - yield from self._enter(start, container, visiting, path, stack) - - while stack: - try: - name, dep = next(stack[-1]) - except StopIteration: - finished = path.pop() - stack.pop() - visiting.discard(finished.provider_id) - visited.add(finished.provider_id) - continue - - yield Edge(path[-1], name, dep) - if dep.provider_id in visiting: - cycle_start = next(i for i, p in enumerate(path) if p.provider_id == dep.provider_id) - yield Cycle([*path[cycle_start:], path[cycle_start]]) - continue - if dep.provider_id in visited: - continue - yield from self._enter(dep, container, visiting, path, stack) - - def _enter( - self, - provider: "AbstractProvider[typing.Any]", - container: "Container", - visiting: set[int], - path: "list[AbstractProvider[typing.Any]]", - stack: "list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]]", - ) -> "typing.Iterator[Event]": - """Push ``provider`` onto the active path; a ``ResolutionError`` from it becomes ``DependenciesError``.""" - visiting.add(provider.provider_id) - path.append(provider) - yield NodeEntered(provider) +def walk( + roots: "typing.Iterable[AbstractProvider[typing.Any]]", + container: "Container", +) -> "typing.Iterator[Event]": + """Pre-order DFS from each root; bookkeeping is shared across roots, keyed on ``provider_id``.""" + visiting: set[int] = set() + visited: set[int] = set() + for root in roots: + yield from _walk_from(root, container, visiting, visited) + + +def find_cycle_from( + start: "AbstractProvider[typing.Any]", + container: "Container", +) -> "list[AbstractProvider[typing.Any]] | None": + """Return the first cycle reachable from ``start``, or None when that subgraph is acyclic.""" + for event in walk([start], container): + if isinstance(event, Cycle): + return event.providers + return None + + +def _walk_from( + start: "AbstractProvider[typing.Any]", + container: "Container", + visiting: set[int], + visited: set[int], +) -> "typing.Iterator[Event]": + """Explicit-stack DFS from ``start``; skip immediately if already seen.""" + if start.provider_id in visited or start.provider_id in visiting: + return + + path: list[AbstractProvider[typing.Any]] = [] + stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = [] + yield from _enter(start, container, visiting, path, stack) + + while stack: try: - dependencies = provider._get_dependencies(container) # noqa: SLF001 - except exceptions.ResolutionError as exc: - yield DependenciesError(provider, exc) - dependencies = {} - stack.append(iter(dependencies.items())) + name, dep = next(stack[-1]) + except StopIteration: + finished = path.pop() + stack.pop() + visiting.discard(finished.provider_id) + visited.add(finished.provider_id) + continue + + yield Edge(path[-1], name, dep) + if dep.provider_id in visiting: + cycle_start = next(i for i, p in enumerate(path) if p.provider_id == dep.provider_id) + yield Cycle([*path[cycle_start:], path[cycle_start]]) + continue + if dep.provider_id in visited: + continue + yield from _enter(dep, container, visiting, path, stack) + + +def _enter( + provider: "AbstractProvider[typing.Any]", + container: "Container", + visiting: set[int], + path: "list[AbstractProvider[typing.Any]]", + stack: "list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]]", +) -> "typing.Iterator[Event]": + """Push ``provider`` onto the active path; a ``ResolutionError`` from it becomes ``DependenciesError``.""" + visiting.add(provider.provider_id) + path.append(provider) + yield NodeEntered(provider) + try: + dependencies = provider._get_dependencies(container) # noqa: SLF001 + except exceptions.ResolutionError as exc: + yield DependenciesError(provider, exc) + dependencies = {} + stack.append(iter(dependencies.items())) def collect_errors(container: "Container", registry: "ProvidersRegistry") -> list[Exception]: """Walk the graph rooted at ``registry``'s providers once; return every wiring error in walk order.""" errors: list[Exception] = [] - for event in DependencyGraph().walk(registry, container): + for event in walk(registry, container): match event: case NodeEntered(provider): errors.extend(provider._iter_validation_issues(container)) # noqa: SLF001 diff --git a/modern_di/group.py b/modern_di/group.py index 898bd80d..9eea3791 100644 --- a/modern_di/group.py +++ b/modern_di/group.py @@ -6,12 +6,12 @@ class Group: - def __new__(cls, *_: typing.Any, **__: typing.Any) -> typing.Self: # noqa: ANN401 + def __new__(cls, *_: object, **__: object) -> typing.Self: raise exceptions.GroupInstantiationError(group_name=cls.__name__) _default_scope: typing.ClassVar["enum.IntEnum | None"] = None - def __init_subclass__(cls, scope: "enum.IntEnum | None" = None, **kwargs: typing.Any) -> None: # noqa: ANN401 + def __init_subclass__(cls, scope: "enum.IntEnum | None" = None, **kwargs: object) -> None: """Record a group-default scope and stamp it onto scope-defaulted providers in this class body.""" super().__init_subclass__(**kwargs) if scope is not None: diff --git a/modern_di/providers/alias.py b/modern_di/providers/alias.py index 01a7a29d..5a0b9489 100644 --- a/modern_di/providers/alias.py +++ b/modern_di/providers/alias.py @@ -27,17 +27,11 @@ def __init__( def __repr__(self) -> str: return f"Alias(source_type={self._source_type!r}, bound_type={self.bound_type!r}, scope={self.scope!r})" - def _find_source(self, container: "Container") -> "AbstractProvider[types.T_co]": - source = container.find_provider(self._source_type) + def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: + source = self._redirect_target(container) if source is None: raise exceptions.AliasSourceNotRegisteredError(source_type=self._source_type) - return source - - def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: - return {"source": self._find_source(container)} + return {"source": source} def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None": - try: - return self._find_source(container) - except exceptions.AliasSourceNotRegisteredError: - return None + return container.find_provider(self._source_type) diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index 37a73881..14e8fe4c 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -46,7 +46,7 @@ class Factory(AbstractProvider[types.T_co]): "_creator", "_has_positional_only_gap", "_kwargs", - "_parsed_kwargs", + "_params", "cache_settings", ) @@ -60,40 +60,10 @@ def __init__( # noqa: PLR0913 cache: bool | CacheSettings[types.T_co] = False, skip_creator_parsing: bool = False, ) -> None: - if skip_creator_parsing: - if bound_type is types.UNSET: - warnings.warn( - "skip_creator_parsing=True without an explicit bound_type means this provider " - "cannot be resolved by type. Pass bound_type=MyClass if you need type resolution.", - UserWarning, - stacklevel=2, - ) - parsed = ParsedCreator(return_type=SignatureItem(), params={}, has_positional_only_gap=False) - else: - parsed = parse_creator(creator) - if kwargs: - self._validate_kwargs_against_signature(creator, kwargs, parsed.params) - for param_name, item in parsed.params.items(): - if item.raw_annotation is None or item.default is not types.UNSET or (kwargs and param_name in kwargs): - continue - raise exceptions.UnsupportedCreatorParameterError( - creator=creator, - parameter_name=param_name, - reason=( - f"parameterized generic annotation {item.raw_annotation!r} cannot be resolved by type; " - "pass the value via the kwargs parameter or give the parameter a default" - ), - ) - if parsed.return_type.args and isinstance(bound_type, types.UnsetType): - members = " | ".join(getattr(t, "__name__", str(t)) for t in parsed.return_type.args) - warnings.warn( - f"The return annotation of {creator!r} is a union of {members}, so no bound_type can be " - "inferred and this provider cannot be resolved by type. Pass bound_type=OneOfThem, or " - "bound_type=None to silence this warning.", - UserWarning, - stacklevel=2, - ) - self._parsed_kwargs = parsed.params + parsed = self._parse_creator( + creator, bound_type=bound_type, kwargs=kwargs, skip_creator_parsing=skip_creator_parsing + ) + self._params = parsed.params self._has_positional_only_gap = parsed.has_positional_only_gap super().__init__( scope=scope, @@ -105,18 +75,47 @@ def __init__( # noqa: PLR0913 self._cached_definition_site: str | types.UnsetType | None = types.UNSET @staticmethod - def _validate_kwargs_against_signature( + def _parse_creator( creator: typing.Callable[..., typing.Any], - kwargs: dict[str, typing.Any], - parsed_kwargs: dict[str, SignatureItem], + *, + bound_type: type | types.UnsetType | None, + kwargs: dict[str, typing.Any] | None, + skip_creator_parsing: bool, + ) -> ParsedCreator: + """Parse and check ``creator`` for ``__init__``; warnings point at the ``Factory(...)`` call.""" + if skip_creator_parsing: + if bound_type is types.UNSET: + warnings.warn( + "skip_creator_parsing=True without an explicit bound_type means this provider " + "cannot be resolved by type. Pass bound_type=MyClass if you need type resolution.", + UserWarning, + stacklevel=3, + ) + return ParsedCreator( + return_type=SignatureItem(), params={}, has_positional_only_gap=False, accepts_any_kwargs=True + ) + parsed = parse_creator(creator) + if kwargs: + Factory._reject_unknown_kwargs(creator, kwargs, parsed) + Factory._reject_unresolvable_generics(creator, kwargs, parsed) + if parsed.return_type.member_types and bound_type is types.UNSET: + members = " | ".join(getattr(t, "__name__", str(t)) for t in parsed.return_type.member_types) + warnings.warn( + f"The return annotation of {creator!r} is a union of {members}, so no bound_type can be " + "inferred and this provider cannot be resolved by type. Pass bound_type=OneOfThem, or " + "bound_type=None to silence this warning.", + UserWarning, + stacklevel=3, + ) + return parsed + + @staticmethod + def _reject_unknown_kwargs( + creator: typing.Callable[..., typing.Any], kwargs: dict[str, typing.Any], parsed: ParsedCreator ) -> None: - try: - sig = inspect.signature(creator) - except (ValueError, TypeError): - return - if any(p.kind is inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()): + if parsed.accepts_any_kwargs: return - known = set(parsed_kwargs) + known = set(parsed.params) unknown = sorted(set(kwargs) - known) if not unknown: return @@ -126,6 +125,22 @@ def _validate_kwargs_against_signature( known_keys=sorted(known), ) + @staticmethod + def _reject_unresolvable_generics( + creator: typing.Callable[..., typing.Any], kwargs: dict[str, typing.Any] | None, parsed: ParsedCreator + ) -> None: + for param_name, item in parsed.params.items(): + if item.raw_annotation is None or item.default is not types.UNSET or (kwargs and param_name in kwargs): + continue + raise exceptions.UnsupportedCreatorParameterError( + creator=creator, + parameter_name=param_name, + reason=( + f"parameterized generic annotation {item.raw_annotation!r} cannot be resolved by type; " + "pass the value via the kwargs parameter or give the parameter a default" + ), + ) + def __repr__(self) -> str: return f"Factory(creator={self._creator!r}, scope={self.scope!r}, cached={self.cache_settings is not None})" @@ -172,12 +187,12 @@ def _argument_resolution_error( bound_type=self.bound_type, creator=self._creator, suggestions=suggestions, - member_types=item.args, + member_types=item.member_types, ) def _wiring_plan(self, registry: "ProvidersRegistry") -> WiringPlan: """Return this factory's wiring plan, memoized on the tree-wide providers registry.""" - return registry.plan_for(self, self._parsed_kwargs, self._kwargs) + return registry.plan_for(self) def _can_call_positionally(self, plan: WiringPlan) -> bool: """Whether this creator can be called positionally under `plan`. @@ -185,18 +200,18 @@ def _can_call_positionally(self, plan: WiringPlan) -> bool: True when every parsed parameter is a positional-or-keyword provider dependency, in signature order, with nothing omitted, added, keyword-only or positional-only. """ - if not plan.pure_provider: + if plan.static_kwargs: return False - names = tuple(self._parsed_kwargs) + names = tuple(self._params) if tuple(plan.provider_kwargs) != names: return False - if any(item.is_keyword_only for item in self._parsed_kwargs.values()): + if any(item.is_keyword_only for item in self._params.values()): return False return not (names and self._has_positional_only_gap) def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: """Return parameter name → dependency provider: a pure registry lookup, no scope or cache touched.""" - return self._wiring_plan(container._providers_registry).edges # noqa: SLF001 + return self._wiring_plan(container._providers_registry).provider_kwargs # noqa: SLF001 def _iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]: """Yield ArgumentResolutionError for parameters with no provider, no default, no static kwarg.""" diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index 7889ccb0..9ba91a95 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -16,12 +16,12 @@ @dataclasses.dataclass(kw_only=True, slots=True) class CacheItem: - settings: CacheSettings[typing.Any] | None + settings: CacheSettings[typing.Any] cache: typing.Any = types.UNSET finalized: bool = False def clear(self) -> None: - if self.settings and self.settings.clear_cache: + if self.settings.clear_cache: self.cache = types.UNSET self.finalized = False @@ -46,10 +46,14 @@ def get_or_create( self.cache = value return value, True + def _pending_finalizer(self) -> typing.Callable[[typing.Any], typing.Awaitable[None] | None] | None: + """Return the finalizer still owed to the cached value, or None when nothing is owed.""" + return None if self.cache is types.UNSET or self.finalized else self.settings.finalizer + async def close_async(self) -> None: - if self.cache is not types.UNSET and not self.finalized and self.settings and self.settings.finalizer: + if (finalizer := self._pending_finalizer()) is not None: try: - result = self.settings.finalizer(self.cache) + result = finalizer(self.cache) if inspect.isawaitable(result): await result except Exception: @@ -60,11 +64,11 @@ async def close_async(self) -> None: self.clear() def close_sync(self) -> None: - if self.cache is not types.UNSET and not self.finalized and self.settings and self.settings.finalizer: + if (finalizer := self._pending_finalizer()) is not None: if self.settings.is_async_finalizer: raise exceptions.AsyncFinalizerInSyncCloseError(finalizer_type=type(self.cache)) try: - result = self.settings.finalizer(self.cache) + result = finalizer(self.cache) except Exception: self.clear() raise @@ -87,12 +91,14 @@ def __init__(self) -> None: def cached_count(self) -> int: return sum(1 for item in self._items.values() if item.cache is not types.UNSET) - def fetch_cache_item(self, provider: Factory[types.T_co]) -> CacheItem: + def fetch_cache_item(self, provider: Factory[typing.Any]) -> CacheItem: + """Return the cache slot for a cached ``provider``, creating it on first use.""" # Get before setdefault: a bare setdefault builds a throwaway CacheItem on every hit. item = self._items.get(provider.provider_id) if item is not None: return item - return self._items.setdefault(provider.provider_id, CacheItem(settings=provider.cache_settings)) + settings = typing.cast("CacheSettings[typing.Any]", provider.cache_settings) + return self._items.setdefault(provider.provider_id, CacheItem(settings=settings)) def mark_created(self, cache_item: CacheItem) -> None: """Record creation completion; close finalizes in reverse of this order (LIFO).""" @@ -101,8 +107,7 @@ def mark_created(self, cache_item: CacheItem) -> None: async def close_async(self) -> None: finalizer_errors: list[Exception] = [] for cache_item in reversed(self._creation_order): - settings = cache_item.settings - if settings is None or settings.finalizer is None: + if cache_item.settings.finalizer is None: cache_item.clear() continue try: diff --git a/modern_di/registries/overrides_registry.py b/modern_di/registries/overrides_registry.py index fff07296..8551c9e2 100644 --- a/modern_di/registries/overrides_registry.py +++ b/modern_di/registries/overrides_registry.py @@ -60,7 +60,7 @@ def __exit__( exc: BaseException | None, tb: TracebackType | None, ) -> None: - if isinstance(self._prior, types.UnsetType): + if self._prior is types.UNSET: self._registry.reset_override(self._provider_id) else: self._registry.override(self._provider_id, self._prior) diff --git a/modern_di/registries/providers_registry.py b/modern_di/registries/providers_registry.py index 98041751..8a6dcc24 100644 --- a/modern_di/registries/providers_registry.py +++ b/modern_di/registries/providers_registry.py @@ -11,7 +11,6 @@ if typing.TYPE_CHECKING: from modern_di import Container from modern_di.providers.factory import Factory - from modern_di.types_parser import SignatureItem class ProvidersRegistry: @@ -61,12 +60,7 @@ def mark_validated(self) -> None: def find_provider(self, dependency_type: type[types.T]) -> AbstractProvider[types.T] | None: return self._providers.get(dependency_type) - def plan_for( - self, - provider: "Factory[typing.Any]", - parsed_kwargs: "dict[str, SignatureItem]", - kwargs: dict[str, typing.Any] | None, - ) -> "WiringPlan": + def plan_for(self, provider: "Factory[typing.Any]") -> "WiringPlan": """Return `provider`'s memoized wiring plan, building it on a miss. The memo is tree-wide and dropped on every registry mutation. @@ -76,11 +70,15 @@ def plan_for( if cached is not None: return cached generation = self._generation - plan = WiringPlan.build(parsed_kwargs=parsed_kwargs, kwargs=kwargs, registry=self, owner=provider) + plan = WiringPlan.build(provider, registry=self) + self._publish(self._plans, provider_id, plan, generation) + return plan + + def _publish(self, memo: dict[typing.Any, typing.Any], key: object, value: object, generation: int) -> None: + """Store `value` in `memo` unless a mutation bumped the generation since `generation` was read.""" with self._lock: if self._generation == generation: - self._plans[provider_id] = plan - return plan + memo[key] = value def _building_set(self) -> set[int]: """Return this thread's in-flight-compile set; per-thread, so a concurrent compile is not a cycle.""" @@ -109,11 +107,7 @@ def resolver_for(self, provider: "AbstractProvider[typing.Any]") -> "typing.Call resolver = compile_resolver(provider, self) finally: building.discard(pid) - with self._lock: - # Publish only if no mutation landed while we compiled; memoizing a resolver built - # against the old registry would strand it past the `_invalidate()` meant to drop it. - if self._generation == generation: - self._resolvers[pid] = resolver + self._publish(self._resolvers, pid, resolver, generation) return resolver def resolver_for_type(self, dependency_type: type) -> "typing.Callable[[Container], typing.Any]": @@ -125,17 +119,13 @@ def resolver_for_type(self, dependency_type: type) -> "typing.Callable[[Containe provider_type=dependency_type, suggestions=suggester.suggest(dependency_type, self) ) resolver = self.resolver_for(provider) - with self._lock: - if self._generation == generation: - self._resolvers_by_type[dependency_type] = resolver + self._publish(self._resolvers_by_type, dependency_type, resolver, generation) return resolver def drop_resolvers(self) -> None: """Drop the compiled resolvers — the overrides changed. Plans and the validation flag survive.""" with self._lock: - self._resolvers.clear() - self._resolvers_by_type.clear() - self._generation += 1 + self._drop_resolvers() def register(self, provider_type: type, provider: AbstractProvider[typing.Any]) -> None: with self._lock: @@ -159,8 +149,7 @@ def add_providers(self, *args: AbstractProvider[typing.Any]) -> None: if provider_type in self._providers: raise exceptions.DuplicateProviderTypeError(provider_type=provider_type) self._providers.update(new_providers) - # Over `args`, not `new_providers`: a reference-only provider never enters - # `_providers`, but its resolver is still compiled and still captures its scope. + # Over `args`: a reference-only provider never enters `_providers` but is still compiled. for provider in args: provider.mark_registered() self._invalidate() @@ -168,7 +157,11 @@ def add_providers(self, *args: AbstractProvider[typing.Any]) -> None: def _invalidate(self) -> None: """Drop every memo and the validation flag — the registry changed. Called under `self._lock`.""" self._plans.clear() + self._validated = False + self._drop_resolvers() + + def _drop_resolvers(self) -> None: + """Drop both resolver memos and bump the generation. Called under `self._lock`.""" self._resolvers.clear() self._resolvers_by_type.clear() - self._validated = False self._generation += 1 diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 8e0f7593..18ee09d8 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -15,7 +15,7 @@ @dataclasses.dataclass(kw_only=True, slots=True, frozen=True) class SignatureItem: arg_type: type | None = None - args: list[type] = dataclasses.field(default_factory=list) + member_types: list[type] = dataclasses.field(default_factory=list) is_nullable: bool = False default: object = UNSET raw_annotation: object = None @@ -44,7 +44,7 @@ def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem": result["is_nullable"] = True if len(non_none_members) > 1: - result["args"] = non_none_members + result["member_types"] = non_none_members else: result["arg_type"] = non_none_members[0] @@ -92,11 +92,14 @@ class ParsedCreator: ``has_positional_only_gap`` is True when a positional-only-with-default parameter was dropped from ``params``, so the map is no longer a faithful positional prefix of the signature. + ``accepts_any_kwargs`` is True when the creator takes ``**kwargs`` or its signature cannot be + read, so no ``kwargs={...}`` key can be called unknown. """ return_type: SignatureItem params: dict[str, SignatureItem] has_positional_only_gap: bool + accepts_any_kwargs: bool def _class_type_hints(creator: type) -> dict[str, typing.Any]: @@ -114,7 +117,10 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: sig = inspect.signature(creator) except (ValueError, TypeError): return ParsedCreator( - return_type=SignatureItem.from_type(typing.cast(type, creator)), params={}, has_positional_only_gap=False + return_type=SignatureItem.from_type(typing.cast(type, creator)), + params={}, + has_positional_only_gap=False, + accepts_any_kwargs=True, ) is_class = isinstance(creator, type) @@ -131,8 +137,12 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: param_hints = {} has_positional_only_gap = False + accepts_any_kwargs = False for param_name, param in sig.parameters.items(): - if param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD): + if param.kind is inspect.Parameter.VAR_KEYWORD: + accepts_any_kwargs = True + continue + if param.kind is inspect.Parameter.VAR_POSITIONAL: continue item = _parse_parameter(creator, param_name, param, type_hints) if item is None: @@ -149,4 +159,9 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: else: return_sig = SignatureItem() - return ParsedCreator(return_type=return_sig, params=param_hints, has_positional_only_gap=has_positional_only_gap) + return ParsedCreator( + return_type=return_sig, + params=param_hints, + has_positional_only_gap=has_positional_only_gap, + accepts_any_kwargs=accepts_any_kwargs, + ) diff --git a/modern_di/wiring.py b/modern_di/wiring.py index 0886f8ed..0c4bd946 100644 --- a/modern_di/wiring.py +++ b/modern_di/wiring.py @@ -24,7 +24,7 @@ def find_dep_provider( if provider is owner: return None return provider - for x in item.args: + for x in item.member_types: provider = registry.find_provider(x) if provider is not None and provider is not owner: return provider @@ -35,33 +35,21 @@ def find_dep_provider( class WiringPlan: """Immutable result of partitioning a creator's parameters into wiring buckets. - ``pure_provider`` means no static kwargs, so the call can be built from ``provider_kwargs`` - alone. ``unwireable`` holds records rather than pre-built exceptions: a plan is memoized, and - ``prepend_step`` mutates the error it is called on. + ``provider_kwargs`` holds every provider the plan resolves, so it is also the edge set the + dependency graph walks. ``unwireable`` holds records rather than pre-built exceptions: a plan + is memoized, and ``prepend_step`` mutates the error it is called on. """ provider_kwargs: dict[str, "AbstractProvider[typing.Any]"] static_kwargs: dict[str, typing.Any] unwireable: "list[tuple[str, SignatureItem]]" - pure_provider: bool - - @property - def edges(self) -> dict[str, "AbstractProvider[typing.Any]"]: - """Every provider this plan resolves: the bucket ``resolve()`` reads.""" - return self.provider_kwargs @classmethod - def build( - cls, - *, - parsed_kwargs: dict[str, SignatureItem], - kwargs: dict[str, typing.Any] | None, - registry: "ProvidersRegistry", - owner: "Factory[typing.Any]", - ) -> "WiringPlan": - """Partition *parsed_kwargs* by type, then overlay ``kwargs={...}``. Never raises.""" + def build(cls, owner: "Factory[typing.Any]", *, registry: "ProvidersRegistry") -> "WiringPlan": + """Partition ``owner``'s parameters by type, then overlay its ``kwargs={...}``. Never raises.""" + kwargs = owner._kwargs # noqa: SLF001 provider_kwargs, static_kwargs, unwireable = cls._wire_by_type( - parsed_kwargs=parsed_kwargs, + params=owner._params, # noqa: SLF001 kwargs=kwargs, registry=registry, owner=owner, @@ -77,13 +65,12 @@ def build( provider_kwargs=provider_kwargs, static_kwargs=static_kwargs, unwireable=unwireable, - pure_provider=not static_kwargs, ) @staticmethod def _wire_by_type( *, - parsed_kwargs: dict[str, SignatureItem], + params: dict[str, SignatureItem], kwargs: dict[str, typing.Any] | None, registry: "ProvidersRegistry", owner: "Factory[typing.Any]", @@ -92,7 +79,7 @@ def _wire_by_type( dict[str, typing.Any], "list[tuple[str, SignatureItem]]", ]: - """Bucket each parsed parameter by type; a name in ``kwargs={...}`` is left to the overlay. + """Bucket each parameter by type; a name in ``kwargs={...}`` is left to the overlay. A parameter with no provider is omitted when it has a default, gets ``None`` when nullable, and is unwireable otherwise. @@ -101,7 +88,7 @@ def _wire_by_type( static_kwargs: dict[str, typing.Any] = {} unwireable: list[tuple[str, SignatureItem]] = [] - for name, item in parsed_kwargs.items(): + for name, item in params.items(): if kwargs and name in kwargs: continue provider = find_dep_provider(registry, owner, item) diff --git a/tests/providers/test_factory.py b/tests/providers/test_factory.py index 2656e1f6..52dfe5f3 100644 --- a/tests/providers/test_factory.py +++ b/tests/providers/test_factory.py @@ -980,7 +980,7 @@ def _cov_pos_only_creator(prefix: str = "P", /, dep: _CovLeaf = None) -> _CovPos def test_positional_only_with_default_stays_on_kwargs_path() -> None: - # `prefix` is positional-only WITH a default: the parser drops it from _parsed_kwargs, leaving + # `prefix` is positional-only WITH a default: the parser drops it from _params, leaving # names == ("dep",) -- a clean-looking prefix. The positional-only guard in _positional_names # must reject it, or `creator(dep_instance)` would bind dep to `prefix` and swallow the "P". assert _cov_pos_only_creator(dep=_CovLeaf()) == _CovPosOnlyResult(prefix="P", dep=_CovLeaf()) # exercise body diff --git a/tests/registries/test_cache_registry.py b/tests/registries/test_cache_registry.py index 68a1120d..33d56caa 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/registries/test_cache_registry.py @@ -9,7 +9,7 @@ def _item() -> CacheItem: - return CacheItem(settings=None) + return CacheItem(settings=CacheSettings()) def test_get_or_create_miss_calls_resolve_and_create_once_and_caches() -> None: @@ -108,9 +108,8 @@ async def _recording(self: CacheItem) -> None: registry = CacheRegistry() plain = CacheItem(settings=CacheSettings(), cache="plain") persistent = CacheItem(settings=CacheSettings(clear_cache=False), cache="persistent") - bare = CacheItem(settings=None, cache="bare") with_finalizer = CacheItem(settings=CacheSettings(finalizer=finalized.append), cache="finalized") - for item in (plain, persistent, bare, with_finalizer): + for item in (plain, persistent, with_finalizer): registry.mark_created(item) await registry.close_async() @@ -119,5 +118,4 @@ async def _recording(self: CacheItem) -> None: assert finalized == ["finalized"] assert plain.cache is UNSET assert persistent.cache == "persistent" - assert bare.cache == "bare" assert registry._creation_order == [] diff --git a/tests/registries/test_providers_registry.py b/tests/registries/test_providers_registry.py index dcd6a7fb..59e33997 100644 --- a/tests/registries/test_providers_registry.py +++ b/tests/registries/test_providers_registry.py @@ -10,7 +10,6 @@ from modern_di.registries import providers_registry as pr_mod from modern_di.registries.providers_registry import ProvidersRegistry from modern_di.scope import Scope -from modern_di.types_parser import SignatureItem def test_providers_registry_find_provider_not_found() -> None: @@ -278,14 +277,8 @@ class G(Group): may_publish = threading.Event() real_build = pr_mod.WiringPlan.build - def hold_open( - *, - parsed_kwargs: "dict[str, SignatureItem]", - kwargs: "dict[str, typing.Any] | None", - registry: ProvidersRegistry, - owner: "providers.Factory[typing.Any]", - ) -> pr_mod.WiringPlan: - plan = real_build(parsed_kwargs=parsed_kwargs, kwargs=kwargs, registry=registry, owner=owner) + def hold_open(owner: "providers.Factory[typing.Any]", *, registry: ProvidersRegistry) -> pr_mod.WiringPlan: + plan = real_build(owner, registry=registry) if owner is svc: built.set() may_publish.wait(5) diff --git a/tests/test_dependency_graph.py b/tests/test_dependency_graph.py index 8b68ac3f..5d530e7b 100644 --- a/tests/test_dependency_graph.py +++ b/tests/test_dependency_graph.py @@ -1,15 +1,16 @@ -"""Event-stream tests for ``DependencyGraph.walk`` — the module's test surface is the event SEQUENCE.""" +"""Event-stream tests for ``dependency_graph.walk`` — the module's test surface is the event SEQUENCE.""" from modern_di import Container, Scope from modern_di.dependency_graph import ( Cycle, DependenciesError, - DependencyGraph, Edge, NodeEntered, build_cycle_error, effective_scope, + find_cycle_from, terminal_chain, + walk, ) from modern_di.group import Group from modern_di.providers import Alias, Factory @@ -37,7 +38,7 @@ class G(Group): leaf = Factory(scope=Scope.APP, creator=Leaf) c = Container(scope=Scope.APP, groups=[G]) - events = list(DependencyGraph().walk([G.root, G.leaf], c)) + events = list(walk([G.root, G.leaf], c)) kinds = [type(e).__name__ for e in events] assert kinds[0] == "NodeEntered" assert "Edge" in kinds @@ -51,7 +52,7 @@ class G(Group): leaf = Factory(scope=Scope.APP, creator=Leaf) c = Container(scope=Scope.APP, groups=[G]) - events = list(DependencyGraph().walk([G.root], c)) + events = list(walk([G.root], c)) assert events == [ NodeEntered(G.root), Edge(G.root, "leaf", G.leaf), @@ -65,7 +66,7 @@ class G(Group): b = Factory(scope=Scope.APP, creator=CycB) c = Container(scope=Scope.APP, groups=[G]) - cycles = [e for e in DependencyGraph().walk([G.a], c) if isinstance(e, Cycle)] + cycles = [e for e in walk([G.a], c) if isinstance(e, Cycle)] assert cycles assert cycles[0].providers[0].provider_id == cycles[0].providers[-1].provider_id @@ -76,7 +77,7 @@ class G(Group): b = Factory(scope=Scope.APP, creator=CycB) c = Container(scope=Scope.APP, groups=[G]) - events = list(DependencyGraph().walk([G.a], c)) + events = list(walk([G.a], c)) assert events == [ NodeEntered(G.a), Edge(G.a, "b", G.b), @@ -101,7 +102,7 @@ class G(Group): shared = Factory(scope=Scope.APP, creator=Shared) c = Container(scope=Scope.APP, groups=[G]) - events = list(DependencyGraph().walk([G.left, G.right], c)) + events = list(walk([G.left, G.right], c)) # Shared is a dep of both roots but entered exactly once. assert sum(isinstance(e, NodeEntered) and e.provider is G.shared for e in events) == 1 # Both roots still emit the Edge to the shared dep; the second finds it visited, no re-descent. @@ -116,7 +117,7 @@ class G(Group): c = Container(scope=Scope.APP, groups=[G]) # leaf appears as a dep of root (first root) AND as a later root; the later root is skipped. - events = list(DependencyGraph().walk([G.root, G.leaf], c)) + events = list(walk([G.root, G.leaf], c)) assert sum(isinstance(e, NodeEntered) and e.provider is G.leaf for e in events) == 1 @@ -125,7 +126,7 @@ class G(Group): leaf = Factory(scope=Scope.APP, creator=Leaf) c = Container(scope=Scope.APP, groups=[G]) - assert DependencyGraph().find_cycle_from(G.leaf, c) is None + assert find_cycle_from(G.leaf, c) is None def test_find_cycle_from_returns_loop() -> None: @@ -134,7 +135,7 @@ class G(Group): b = Factory(scope=Scope.APP, creator=CycB) c = Container(scope=Scope.APP, groups=[G]) - cycle = DependencyGraph().find_cycle_from(G.a, c) + cycle = find_cycle_from(G.a, c) assert cycle == [G.a, G.b, G.a] @@ -179,7 +180,7 @@ class G(Group): alias = Alias(Missing, bound_type=Marker) c = Container(scope=Scope.APP, groups=[G]) - events = list(DependencyGraph().walk([G.alias], c)) + events = list(walk([G.alias], c)) assert isinstance(events[0], NodeEntered) assert events[0].provider is G.alias assert isinstance(events[1], DependenciesError) @@ -207,7 +208,7 @@ def test_walk_emits_cycle_closed_through_kwargs_overlay() -> None: c = Container(scope=Scope.APP) c._providers_registry.add_providers(a, b) - events = list(DependencyGraph().walk([a], c)) + events = list(walk([a], c)) cycles = [e for e in events if isinstance(e, Cycle)] assert len(cycles) == 1 assert [p.display_name for p in cycles[0].providers] == ["KwCycA", "KwCycB", "KwCycA"] diff --git a/tests/test_dependency_graph_contract.py b/tests/test_dependency_graph_contract.py index 8525662d..120d29d3 100644 --- a/tests/test_dependency_graph_contract.py +++ b/tests/test_dependency_graph_contract.py @@ -82,7 +82,7 @@ def _explode(*_: object, **__: object) -> object: # pragma: no cover - a valida msg = "re-walked" raise AssertionError(msg) - monkeypatch.setattr(dependency_graph.DependencyGraph, "walk", _explode) + monkeypatch.setattr(dependency_graph, "walk", _explode) container.validate() # short-circuited on the registry's validated flag -> no walk @@ -100,10 +100,10 @@ class G(Group): def test_validate_walks_the_same_edges_resolve_follows() -> None: """INVARIANT: the graph validate() walks is the graph resolve() follows. - Edges come from `WiringPlan.edges`, a view derived from the same buckets resolve() reads, so a - provider named in a declaration-time `kwargs={...}` is an edge like any type-matched one. - Assembling the validation edge set separately would let the two drift, and a cycle routed - through a `kwargs=` provider would surface as a bare RecursionError instead. + Edges come from `WiringPlan.provider_kwargs`, the same bucket resolve() reads, so a provider + named in a declaration-time `kwargs={...}` is an edge like any type-matched one. Assembling the + validation edge set separately would let the two drift, and a cycle routed through a `kwargs=` + provider would surface as a bare RecursionError instead. """ class _Leaf: ... @@ -115,7 +115,7 @@ def __init__(self, leaf: _Leaf) -> None: class G(Group): leaf = Factory(scope=Scope.REQUEST, creator=_Leaf) # Named via kwargs, not type-matched: the by-type pass skips a name present in kwargs, - # so this edge exists only if the overlay pass feeds it into WiringPlan.edges. + # so this edge exists only if the overlay pass feeds it into WiringPlan.provider_kwargs. root = Factory(scope=Scope.APP, creator=_Root, kwargs={"leaf": leaf}) container = Container(scope=Scope.APP, groups=[G]) diff --git a/tests/test_runtime_cycle_guard.py b/tests/test_runtime_cycle_guard.py index 9e3883b1..b052861b 100644 --- a/tests/test_runtime_cycle_guard.py +++ b/tests/test_runtime_cycle_guard.py @@ -174,7 +174,7 @@ def _explode(*_: object, **__: object) -> object: # pragma: no cover - validate msg = "walked" raise AssertionError(msg) - monkeypatch.setattr(dependency_graph.DependencyGraph, "find_cycle_from", _explode) + monkeypatch.setattr(dependency_graph, "find_cycle_from", _explode) with pytest.raises(RecursionError): container.resolve(SelfRec) diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index b92bc8df..33fc8607 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -26,8 +26,8 @@ class GenericClass(typing.Generic[types.T]): ... (dict[str, typing.Any], SignatureItem(raw_annotation=dict[str, typing.Any])), (typing.Optional[str], SignatureItem(arg_type=str, is_nullable=True)), # noqa: UP045 (str | None, SignatureItem(arg_type=str, is_nullable=True)), - (str | int, SignatureItem(args=[str, int])), - (typing.Union[str | int], SignatureItem(args=[str, int])), # noqa: UP007 + (str | int, SignatureItem(member_types=[str, int])), + (typing.Union[str | int], SignatureItem(member_types=[str, int])), # noqa: UP007 (list[str] | None, SignatureItem(arg_type=list, is_nullable=True)), (GenericClass[str], SignatureItem(raw_annotation=GenericClass[str])), (GenericClass[str] | None, SignatureItem(arg_type=GenericClass, is_nullable=True)), diff --git a/tests/test_wiring.py b/tests/test_wiring.py index 608bd5be..6ce81cca 100644 --- a/tests/test_wiring.py +++ b/tests/test_wiring.py @@ -76,10 +76,8 @@ def test_wiring_plan_partitioning() -> None: ) plan = WiringPlan.build( - parsed_kwargs=owner._parsed_kwargs, - kwargs=owner._kwargs, + owner, registry=registry, - owner=owner, ) # a) type-matched provider → provider_kwargs @@ -91,7 +89,7 @@ def test_wiring_plan_partitioning() -> None: # c) ContextProvider param → provider_kwargs, and an edge like any other assert plan.provider_kwargs["req"] is ctx_req - assert plan.edges["req"] is ctx_req + assert plan.provider_kwargs["req"] is ctx_req # d) defaulted param → omitted from both buckets assert "with_default" not in plan.provider_kwargs @@ -107,10 +105,8 @@ def test_wiring_plan_nullable_no_default_goes_to_static_kwargs() -> None: owner = providers.Factory(scope=Scope.APP, creator=_NullableNoDefaultCreator) plan = WiringPlan.build( - parsed_kwargs=owner._parsed_kwargs, - kwargs=None, + owner, registry=registry, - owner=owner, ) assert "nullable" in plan.static_kwargs @@ -134,10 +130,8 @@ def test_wiring_plan_unwireable_no_raise() -> None: owner = providers.Factory(scope=Scope.APP, creator=_UnwirableCreator) plan = WiringPlan.build( - parsed_kwargs=owner._parsed_kwargs, - kwargs=None, + owner, registry=registry, - owner=owner, ) # build returns normally (no raise) @@ -176,21 +170,19 @@ def test_wiring_plan_edges_include_static_supplied_providers() -> None: ) plan = WiringPlan.build( - parsed_kwargs=owner._parsed_kwargs, - kwargs=owner._kwargs, + owner, registry=registry, - owner=owner, ) # `x` is supplied via the kwargs overlay: resolved live AND visible to validate(). assert "x" in plan.provider_kwargs - assert plan.edges["x"] is factory_a + assert plan.provider_kwargs["x"] is factory_a # `y` is type-matched → an edge like any other. - assert plan.edges["y"] is factory_b + assert plan.provider_kwargs["y"] is factory_b # The edge set is exactly what the runtime resolves — however the edge was declared. - assert set(plan.edges) == {"x", "y"} + assert set(plan.provider_kwargs) == {"x", "y"} assert plan.unwireable == [] @@ -215,7 +207,8 @@ def test_wiring_plan_edges_include_static_supplied_providers() -> None: ) def test_parameter_without_provider_precedence(item: SignatureItem, expected: str) -> None: owner = providers.Factory(scope=Scope.APP, creator=_ServiceA) - plan = WiringPlan.build(parsed_kwargs={"p": item}, kwargs=None, registry=ProvidersRegistry(), owner=owner) + owner._params = {"p": item} + plan = WiringPlan.build(owner, registry=ProvidersRegistry()) outcome = { "omitted": ({}, {}, []), "none": ({}, {"p": None}, []), @@ -232,14 +225,15 @@ def test_parameter_without_provider_precedence(item: SignatureItem, expected: st def test_parameter_without_provider_default_wins_over_nullable() -> None: owner = providers.Factory(scope=Scope.APP, creator=_ServiceA) item = SignatureItem(default=None, is_nullable=True) - plan = WiringPlan.build(parsed_kwargs={"p": item}, kwargs=None, registry=ProvidersRegistry(), owner=owner) + owner._params = {"p": item} + plan = WiringPlan.build(owner, registry=ProvidersRegistry()) # default is not UNSET (it is None), so omitted regardless of is_nullable assert plan.static_kwargs == {} assert plan.unwireable == [] # --------------------------------------------------------------------------- -# Test: find_dep_provider — union-args branch (arg_type is None) +# Test: find_dep_provider — union-members branch (arg_type is None) # --------------------------------------------------------------------------- @@ -252,9 +246,9 @@ class _UnionTypeB: def test_find_dep_provider_union_args_matches_first_registered() -> None: - """When arg_type is None (union member), find_dep_provider falls through to args list. + """When arg_type is None (union member), find_dep_provider falls through to member_types. - A SignatureItem with args=[_UnionTypeA, _UnionTypeB] and no arg_type (the + A SignatureItem with member_types=[_UnionTypeA, _UnionTypeB] and no arg_type (the union-member branch) should resolve to the provider registered for the first matching member — and appear in provider_kwargs/dependencies accordingly. """ @@ -264,8 +258,8 @@ def test_find_dep_provider_union_args_matches_first_registered() -> None: owner = providers.Factory(scope=Scope.APP, creator=_UnionTypeA) - # Manually craft a SignatureItem with args but no arg_type (union member scenario) - item = SignatureItem(arg_type=None, args=[_UnionTypeA, _UnionTypeB]) + # Manually craft a SignatureItem with member_types but no arg_type (union member scenario) + item = SignatureItem(arg_type=None, member_types=[_UnionTypeA, _UnionTypeB]) result = find_dep_provider(registry, owner, item) # factory_a is registered for _UnionTypeA; it is not `owner`, so it must be returned @@ -279,7 +273,7 @@ def test_find_dep_provider_union_args_skips_owner() -> None: registry.add_providers(factory_a) # owner IS factory_a — should be skipped - item = SignatureItem(arg_type=None, args=[_UnionTypeA]) + item = SignatureItem(arg_type=None, member_types=[_UnionTypeA]) result = find_dep_provider(registry, factory_a, item) assert result is None @@ -292,7 +286,7 @@ def __init__(self, first: _ServiceA, second: _ServiceB, third: _Request) -> None def test_provider_kwargs_preserves_signature_order() -> None: """provider_kwargs iterates in signature order — the invariant the positional fast path depends on. - _positional_names gates on tuple(provider_kwargs) == tuple(parsed_kwargs), then the resolver + _positional_names gates on tuple(provider_kwargs) == tuple(params), then the resolver builds its positional tuple from provider_kwargs. If build stopped preserving order, the gate would silently de-select the positional path. """ @@ -305,11 +299,9 @@ def test_provider_kwargs_preserves_signature_order() -> None: owner = providers.Factory(scope=Scope.APP, creator=_OrderedDeps) plan = WiringPlan.build( - parsed_kwargs=owner._parsed_kwargs, - kwargs=owner._kwargs, + owner, registry=registry, - owner=owner, ) assert tuple(plan.provider_kwargs) == ("first", "second", "third") - assert tuple(plan.provider_kwargs) == tuple(owner._parsed_kwargs) + assert tuple(plan.provider_kwargs) == tuple(owner._params) From 139bc34d00204fdfaa7238939037ebe2fa7e42ed Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Mon, 5 Oct 2026 12:46:32 +0300 Subject: [PATCH 2/2] refactor: address review on registry and wiring cleanup (#575) --- modern_di/registries/providers_registry.py | 6 +++++- tests/providers/test_factory.py | 21 ++++++++++++++++++++- tests/test_dependency_graph.py | 2 +- tests/test_resolver_compiler.py | 2 +- tests/test_wiring.py | 11 +++++------ 5 files changed, 32 insertions(+), 10 deletions(-) diff --git a/modern_di/registries/providers_registry.py b/modern_di/registries/providers_registry.py index 8a6dcc24..5e0e5daf 100644 --- a/modern_di/registries/providers_registry.py +++ b/modern_di/registries/providers_registry.py @@ -13,6 +13,10 @@ from modern_di.providers.factory import Factory +_K = typing.TypeVar("_K") +_V = typing.TypeVar("_V") + + class ProvidersRegistry: """Type → provider, plus the tree-wide plan and resolver memos. @@ -74,7 +78,7 @@ def plan_for(self, provider: "Factory[typing.Any]") -> "WiringPlan": self._publish(self._plans, provider_id, plan, generation) return plan - def _publish(self, memo: dict[typing.Any, typing.Any], key: object, value: object, generation: int) -> None: + def _publish(self, memo: dict[_K, _V], key: _K, value: _V, generation: int) -> None: """Store `value` in `memo` unless a mutation bumped the generation since `generation` was read.""" with self._lock: if self._generation == generation: diff --git a/tests/providers/test_factory.py b/tests/providers/test_factory.py index 52dfe5f3..ccc4eecd 100644 --- a/tests/providers/test_factory.py +++ b/tests/providers/test_factory.py @@ -275,6 +275,25 @@ def test_factory_skip_creator_parsing_without_bound_type_warns() -> None: providers.Factory(creator=str, skip_creator_parsing=True) +def _union_return_creator() -> int | str: + return 0 # pragma: no cover - never called; only its return annotation is read + + +@pytest.mark.parametrize( + "build", + [ + pytest.param(lambda: providers.Factory(creator=str, skip_creator_parsing=True), id="skip_creator_parsing"), + pytest.param(lambda: providers.Factory(creator=_union_return_creator), id="union_return"), + ], +) +def test_factory_warning_points_at_the_factory_call(build: typing.Callable[[], object]) -> None: + with pytest.warns(UserWarning, match="bound_type") as record: + build() + assert len(record) == 1 + assert record[0].filename == __file__ + assert record[0].lineno == inspect.getsourcelines(build)[1] + + def test_factory_skip_creator_parsing_with_bound_type_no_warning() -> None: with warnings.catch_warnings(): warnings.simplefilter("error") @@ -981,7 +1000,7 @@ def _cov_pos_only_creator(prefix: str = "P", /, dep: _CovLeaf = None) -> _CovPos def test_positional_only_with_default_stays_on_kwargs_path() -> None: # `prefix` is positional-only WITH a default: the parser drops it from _params, leaving - # names == ("dep",) -- a clean-looking prefix. The positional-only guard in _positional_names + # names == ("dep",) -- a clean-looking prefix. The positional-only guard in _can_call_positionally # must reject it, or `creator(dep_instance)` would bind dep to `prefix` and swallow the "P". assert _cov_pos_only_creator(dep=_CovLeaf()) == _CovPosOnlyResult(prefix="P", dep=_CovLeaf()) # exercise body diff --git a/tests/test_dependency_graph.py b/tests/test_dependency_graph.py index 5d530e7b..58d34b6d 100644 --- a/tests/test_dependency_graph.py +++ b/tests/test_dependency_graph.py @@ -1,4 +1,4 @@ -"""Event-stream tests for ``dependency_graph.walk`` — the module's test surface is the event SEQUENCE.""" +"""Event-stream tests for ``dependency_graph.walk``: the module's test surface is the event SEQUENCE.""" from modern_di import Container, Scope from modern_di.dependency_graph import ( diff --git a/tests/test_resolver_compiler.py b/tests/test_resolver_compiler.py index f5becccc..7bed80ed 100644 --- a/tests/test_resolver_compiler.py +++ b/tests/test_resolver_compiler.py @@ -466,7 +466,7 @@ def test_can_call_positionally_rejects_positional_only_param() -> None: a correctness bug, not a slow path. """ - # rule 4: `prefix` is positional-only WITH a default, dropped from parsed_kwargs so the + # rule 4: `prefix` is positional-only WITH a default, dropped from params so the # remaining names look like a clean prefix ("dep",) -- but a positional call would bind # `dep` to the `prefix` slot. The parser's has_positional_only_gap flag must reject it. def creator(prefix: str = "P", /, dep: _A = None) -> _Ordered: # ty: ignore[invalid-parameter-default] diff --git a/tests/test_wiring.py b/tests/test_wiring.py index 6ce81cca..2a5a9291 100644 --- a/tests/test_wiring.py +++ b/tests/test_wiring.py @@ -87,8 +87,7 @@ def test_wiring_plan_partitioning() -> None: # b) static kwarg literal → static_kwargs assert plan.static_kwargs.get("svc_b") == "static-literal" - # c) ContextProvider param → provider_kwargs, and an edge like any other - assert plan.provider_kwargs["req"] is ctx_req + # c) ContextProvider param → provider_kwargs, like any other dependency assert plan.provider_kwargs["req"] is ctx_req # d) defaulted param → omitted from both buckets @@ -233,7 +232,7 @@ def test_parameter_without_provider_default_wins_over_nullable() -> None: # --------------------------------------------------------------------------- -# Test: find_dep_provider — union-members branch (arg_type is None) +# Test: find_dep_provider, union-members branch (arg_type is None) # --------------------------------------------------------------------------- @@ -245,7 +244,7 @@ class _UnionTypeB: pass -def test_find_dep_provider_union_args_matches_first_registered() -> None: +def test_find_dep_provider_union_members_matches_first_registered() -> None: """When arg_type is None (union member), find_dep_provider falls through to member_types. A SignatureItem with member_types=[_UnionTypeA, _UnionTypeB] and no arg_type (the @@ -266,7 +265,7 @@ def test_find_dep_provider_union_args_matches_first_registered() -> None: assert result is factory_a -def test_find_dep_provider_union_args_skips_owner() -> None: +def test_find_dep_provider_union_members_skips_owner() -> None: """When the only union member matching a registered provider IS the owner, returns None.""" factory_a = providers.Factory(scope=Scope.APP, creator=_UnionTypeA) registry = ProvidersRegistry() @@ -286,7 +285,7 @@ def __init__(self, first: _ServiceA, second: _ServiceB, third: _Request) -> None def test_provider_kwargs_preserves_signature_order() -> None: """provider_kwargs iterates in signature order — the invariant the positional fast path depends on. - _positional_names gates on tuple(provider_kwargs) == tuple(params), then the resolver + _can_call_positionally gates on tuple(provider_kwargs) == tuple(params), then the resolver builds its positional tuple from provider_kwargs. If build stopped preserving order, the gate would silently de-select the positional path. """