From 91f8bf3d875d281c03a98a93136b5b31e679340d Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 14:55:07 +0300 Subject: [PATCH 01/17] refactor: pass the registry to provider hooks and tidy private internals (#605) - The provider hooks (`_get_dependencies`, `_redirect_target`, `_iter_validation_issues`) and the `dependency_graph` functions take a `ProvidersRegistry` instead of a `Container`; they only ever used its registry. `collect_errors` takes the registry alone. - `_defaultless_context_terminal` reuses `terminal_chain` instead of its own alias walk, so the test that kept the two walks in sync is gone. `Factory._wiring_plan` is gone; callers use `registry.plan_for`. - `ContextRegistry` is folded into a `Container._context` dict. - `STEP_ERRORS` is gone; the container and the resolver template catch `ResolutionError` directly. - `AbstractProvider.__init__` resolves an unset `bound_type` from the subclass's `inferred_bound_type`, once. - `SignatureItem.raw_annotation` is renamed `unresolvable_generic`. - `Container.resolve` looks the provider up with `find_provider` in its `RecursionError` branch too, and an unregistered type there re-raises the `RecursionError`. `register` and `add_providers` share one locked `_add`. No behaviour change. Folding the context dict into the container takes one object off child construction: G6 -11.4%, G6b -9.0%, C6 -5.8% (A/B against ffe365c, 10 interleaved rounds). --- modern_di/container.py | 30 +++++---- modern_di/dependency_graph.py | 47 +++++++------- modern_di/providers/abstract.py | 14 +++-- modern_di/providers/alias.py | 14 ++--- modern_di/providers/context_provider.py | 4 +- modern_di/providers/factory.py | 27 ++++---- modern_di/registries/context_registry.py | 13 ---- modern_di/registries/providers_registry.py | 20 +++--- modern_di/resolver_compiler.py | 40 ++++++------ modern_di/types_parser.py | 14 ++--- tests/providers/test_alias.py | 6 +- tests/providers/test_context_provider.py | 2 +- tests/registries/test_providers_registry.py | 2 +- tests/test_container.py | 7 ++- tests/test_dependency_graph.py | 68 ++++++++++++--------- tests/test_resolver_compiler.py | 40 +----------- tests/test_types_parser.py | 8 +-- 17 files changed, 151 insertions(+), 205 deletions(-) delete mode 100644 modern_di/registries/context_registry.py diff --git a/modern_di/container.py b/modern_di/container.py index 94519093..b82e5bc4 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -8,23 +8,21 @@ from modern_di.providers.abstract import AbstractProvider from modern_di.providers.container_provider import container_provider from modern_di.registries.cache_registry import CacheRegistry -from modern_di.registries.context_registry import ContextRegistry from modern_di.registries.overrides_registry import OverrideHandle from modern_di.registries.providers_registry import ProvidersRegistry -from modern_di.resolver_compiler import STEP_ERRORS from modern_di.scope import Scope def _handle_recursion_error( - provider: AbstractProvider[typing.Any], container: "Container", registry: ProvidersRegistry, exc: RecursionError + provider: AbstractProvider[typing.Any] | None, registry: ProvidersRegistry, exc: RecursionError ) -> typing.NoReturn: """Convert an escaped `RecursionError` to `CircularDependencyError`, or re-raise it unchanged.""" - if registry.is_validated(): + if provider is None or registry.is_validated(): raise exc # validated => acyclic static graph => genuine self-recursion - cycle = dependency_graph.find_cycle_from(provider, container) + cycle = dependency_graph.find_cycle_from(provider, registry) if cycle is None: raise exc - raise dependency_graph.build_cycle_error(cycle, container) from exc + raise dependency_graph.build_cycle_error(cycle, registry) from exc class Container: @@ -38,7 +36,7 @@ class Container: __slots__ = ( "_cache_registry", "_closed", - "_context_registry", + "_context", "_parent_container", "_providers_registry", "_scope", @@ -128,7 +126,7 @@ def _set_state( self._parent_container = parent self._scope_map = scope_map self._cache_registry = CacheRegistry() - self._context_registry = ContextRegistry(copy.copy(context) if context is not None else {}) + self._context = copy.copy(context) if context is not None else {} self._providers_registry = providers_registry def find_container(self, scope: enum.IntEnum) -> typing.Self: @@ -163,11 +161,11 @@ def resolve(self, dependency_type: type[types.T]) -> types.T: resolver = registry.resolver_for_type(dependency_type) return resolver(self) except RecursionError as exc: - _handle_recursion_error(registry._providers[dependency_type], self, registry, exc) # noqa: SLF001 - except STEP_ERRORS as exc: + _handle_recursion_error(registry.find_provider(dependency_type), registry, exc) + except exceptions.ResolutionError as exc: provider = registry.find_provider(dependency_type) if provider is not None: - exc._prepend_step(*dependency_graph.redirect_hops(provider, self)) # noqa: SLF001 + exc._prepend_step(*dependency_graph.redirect_hops(provider, registry)) # noqa: SLF001 raise def resolve_dependency(self, dependency: "AbstractProvider[types.T] | type[types.T]") -> types.T: @@ -191,9 +189,9 @@ def resolve_provider(self, provider: "AbstractProvider[types.T]") -> types.T: resolver = registry.resolver_for(provider) return resolver(self) except RecursionError as exc: - _handle_recursion_error(provider, self, registry, exc) - except STEP_ERRORS as exc: - exc._prepend_step(*dependency_graph.redirect_hops(provider, self)) # noqa: SLF001 + _handle_recursion_error(provider, registry, exc) + except exceptions.ResolutionError as exc: + exc._prepend_step(*dependency_graph.redirect_hops(provider, registry)) # noqa: SLF001 raise def validate(self) -> None: @@ -208,7 +206,7 @@ def validate(self) -> None: if reg.is_validated(): return - if errors := dependency_graph.collect_errors(self, reg): + if errors := dependency_graph.collect_errors(reg): raise exceptions.ValidationFailedError(errors=errors) reg.mark_validated() @@ -285,7 +283,7 @@ def set_context(self, context_type: type[types.T], obj: types.T) -> None: matches the ``ContextProvider``. A cached provider is built once and is not rebuilt by a later ``set_context``; set the context before its first resolve. """ - self._context_registry.set_context(context_type, obj) + self._context[context_type] = obj def __repr__(self) -> str: n_providers = len(self._providers_registry) diff --git a/modern_di/dependency_graph.py b/modern_di/dependency_graph.py index 9d7853bc..81c11f7f 100644 --- a/modern_di/dependency_graph.py +++ b/modern_di/dependency_graph.py @@ -14,7 +14,6 @@ if typing.TYPE_CHECKING: - from modern_di import Container from modern_di.registries.providers_registry import ProvidersRegistry @@ -49,7 +48,7 @@ class DependenciesError(NamedTuple): def terminal_chain( - provider: "AbstractProvider[typing.Any]", container: "Container" + provider: "AbstractProvider[typing.Any]", registry: "ProvidersRegistry" ) -> "list[AbstractProvider[typing.Any]]": """Follow ``_redirect_target`` hops from ``provider``, ``provider`` first. @@ -58,7 +57,7 @@ def terminal_chain( """ chain = [provider] seen: set[int] = set() - while (nxt := provider._redirect_target(container)) is not None: # noqa: SLF001 + while (nxt := provider._redirect_target(registry)) is not None: # noqa: SLF001 if provider.provider_id in seen: return [provider] seen.add(provider.provider_id) @@ -67,25 +66,25 @@ def terminal_chain( return chain -def effective_scope(provider: "AbstractProvider[typing.Any]", container: "Container") -> enum.IntEnum: +def effective_scope(provider: "AbstractProvider[typing.Any]", registry: "ProvidersRegistry") -> enum.IntEnum: """Return the scope a provider actually resolves at: its terminal's, once redirects are followed.""" - return terminal_chain(provider, container)[-1].scope + return terminal_chain(provider, registry)[-1].scope def redirect_hops( - provider: "AbstractProvider[typing.Any]", container: "Container" + provider: "AbstractProvider[typing.Any]", registry: "ProvidersRegistry" ) -> "list[exceptions.ResolutionStep]": """Return the chain steps for the redirects between ``provider`` and its terminal, terminal excluded. A redirect owns no lifetime of its own, so each hop is drawn at the scope the terminal resolves at. """ - *hops, terminal = terminal_chain(provider, container) + *hops, terminal = terminal_chain(provider, registry) return [p._resolution_step(terminal.scope) for p in hops] # noqa: SLF001 def build_cycle_error( providers: "list[AbstractProvider[typing.Any]]", - container: "Container", + registry: "ProvidersRegistry", ) -> "exceptions.CircularDependencyError": """Build a ``CircularDependencyError`` from a cycle's providers (first node repeated last). @@ -97,27 +96,27 @@ def build_cycle_error( rotated = [*ring[lead:], *ring[:lead]] canonical = [*rotated, rotated[0]] return exceptions.CircularDependencyError( - steps=[p._resolution_step(effective_scope(p, container)) for p in canonical] # noqa: SLF001 + steps=[p._resolution_step(effective_scope(p, registry)) for p in canonical] # noqa: SLF001 ) def walk( roots: "typing.Iterable[AbstractProvider[typing.Any]]", - container: "Container", + registry: "ProvidersRegistry", ) -> "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) + yield from _walk_from(root, registry, visiting, visited) def find_cycle_from( start: "AbstractProvider[typing.Any]", - container: "Container", + registry: "ProvidersRegistry", ) -> "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): + for event in walk([start], registry): if isinstance(event, Cycle): return event.providers return None @@ -125,7 +124,7 @@ def find_cycle_from( def _walk_from( start: "AbstractProvider[typing.Any]", - container: "Container", + registry: "ProvidersRegistry", visiting: set[int], visited: set[int], ) -> "typing.Iterator[Event]": @@ -135,7 +134,7 @@ def _walk_from( path: list[AbstractProvider[typing.Any]] = [] stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = [] - yield from _enter(start, container, visiting, path, stack) + yield from _enter(start, registry, visiting, path, stack) while stack: try: @@ -154,12 +153,12 @@ def _walk_from( continue if dep.provider_id in visited: continue - yield from _enter(dep, container, visiting, path, stack) + yield from _enter(dep, registry, visiting, path, stack) def _enter( provider: "AbstractProvider[typing.Any]", - container: "Container", + registry: "ProvidersRegistry", visiting: set[int], path: "list[AbstractProvider[typing.Any]]", stack: "list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]]", @@ -169,26 +168,26 @@ def _enter( path.append(provider) yield NodeEntered(provider) try: - dependencies = provider._get_dependencies(container) # noqa: SLF001 + dependencies = provider._get_dependencies(registry) # 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]: +def collect_errors(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 walk(registry, container): + for event in walk(registry, registry): match event: case NodeEntered(provider): - errors.extend(provider._iter_validation_issues(container)) # noqa: SLF001 + errors.extend(provider._iter_validation_issues(registry)) # noqa: SLF001 case DependenciesError(_, error): errors.append(error) case Edge(parent, name, dep): - dependency_chain = terminal_chain(dep, container) + dependency_chain = terminal_chain(dep, registry) dependency_scope = dependency_chain[-1].scope - parent_scope = effective_scope(parent, container) + parent_scope = effective_scope(parent, registry) if dependency_scope > parent_scope: errors.append( exceptions.InvalidScopeDependencyError( @@ -206,7 +205,7 @@ def collect_errors(container: "Container", registry: "ProvidersRegistry") -> lis ) ) case Cycle(providers): - errors.append(build_cycle_error(providers, container)) + errors.append(build_cycle_error(providers, registry)) case _: typing.assert_never(event) return errors diff --git a/modern_di/providers/abstract.py b/modern_di/providers/abstract.py index 501b0dae..055b5a01 100644 --- a/modern_di/providers/abstract.py +++ b/modern_di/providers/abstract.py @@ -7,7 +7,7 @@ if typing.TYPE_CHECKING: - from modern_di import Container + from modern_di.registries.providers_registry import ProvidersRegistry _provider_id_counter = itertools.count() @@ -22,12 +22,14 @@ def __init__( self, *, scope: enum.IntEnum | types.UnsetType, - bound_type: type | None, + bound_type: type | types.UnsetType | None, + inferred_bound_type: type | None = None, ) -> None: + """Set the shared state; an unset ``bound_type`` falls back to ``inferred_bound_type``.""" self._explicit_scope: enum.IntEnum | None = scope if isinstance(scope, enum.IntEnum) else None self._group_claim: tuple[enum.IntEnum, str] | None = None self._registered = False - self._bound_type = bound_type + self._bound_type = inferred_bound_type if isinstance(bound_type, types.UnsetType) else bound_type self._provider_id = next(_provider_id_counter) @property @@ -96,13 +98,13 @@ def _resolution_step(self, scope: enum.IntEnum | None = None) -> exceptions.Reso scope=self.scope if scope is None else scope, name=self.display_name, location=self.definition_site ) - def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: # noqa: ARG002 + def _get_dependencies(self, registry: "ProvidersRegistry") -> dict[str, "AbstractProvider[typing.Any]"]: # noqa: ARG002 return {} - def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None": # noqa: ARG002 + def _redirect_target(self, registry: "ProvidersRegistry") -> "AbstractProvider[typing.Any] | None": # noqa: ARG002 """Return the provider this transparently forwards to, or None if resolution terminates here.""" return None - def _iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]: # noqa: ARG002 + def _iter_validation_issues(self, registry: "ProvidersRegistry") -> typing.Iterable[Exception]: # noqa: ARG002 """Yield validation-time issues for this provider. Default: no issues.""" return iter(()) diff --git a/modern_di/providers/alias.py b/modern_di/providers/alias.py index 5a0b9489..b61ee2a6 100644 --- a/modern_di/providers/alias.py +++ b/modern_di/providers/alias.py @@ -5,7 +5,7 @@ if typing.TYPE_CHECKING: - from modern_di import Container + from modern_di.registries.providers_registry import ProvidersRegistry class Alias(AbstractProvider[types.T_co]): @@ -19,19 +19,17 @@ def __init__( *, bound_type: type | types.UnsetType | None = types.UNSET, ) -> None: - super().__init__( - scope=types.UNSET, bound_type=source_type if isinstance(bound_type, types.UnsetType) else bound_type - ) + super().__init__(scope=types.UNSET, bound_type=bound_type, inferred_bound_type=source_type) self._source_type = source_type def __repr__(self) -> str: return f"Alias(source_type={self._source_type!r}, bound_type={self.bound_type!r}, scope={self.scope!r})" - def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: - source = self._redirect_target(container) + def _get_dependencies(self, registry: "ProvidersRegistry") -> dict[str, "AbstractProvider[typing.Any]"]: + source = self._redirect_target(registry) if source is None: raise exceptions.AliasSourceNotRegisteredError(source_type=self._source_type) return {"source": source} - def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None": - return container.find_provider(self._source_type) + def _redirect_target(self, registry: "ProvidersRegistry") -> "AbstractProvider[typing.Any] | None": + return registry.find_provider(self._source_type) diff --git a/modern_di/providers/context_provider.py b/modern_di/providers/context_provider.py index b0cf82c9..9eaffeb0 100644 --- a/modern_di/providers/context_provider.py +++ b/modern_di/providers/context_provider.py @@ -45,9 +45,7 @@ def __init__( bound_type: type | types.UnsetType | None = types.UNSET, default: typing.Any = types.UNSET, ) -> None: - super().__init__( - scope=scope, bound_type=context_type if isinstance(bound_type, types.UnsetType) else bound_type - ) + super().__init__(scope=scope, bound_type=bound_type, inferred_bound_type=context_type) self._context_type = context_type self._default = default diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index 33112cfc..5946f20d 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -11,7 +11,6 @@ if typing.TYPE_CHECKING: - from modern_di import Container from modern_di.registries.providers_registry import ProvidersRegistry @@ -66,10 +65,7 @@ def __init__( # noqa: PLR0913 ) self._params = parsed.params self._has_positional_only_gap = parsed.has_positional_only_gap - super().__init__( - scope=scope, - bound_type=parsed.return_type.arg_type if isinstance(bound_type, types.UnsetType) else bound_type, - ) + super().__init__(scope=scope, bound_type=bound_type, inferred_bound_type=parsed.return_type.arg_type) self._creator = creator self._cache_settings: CacheSettings[typing.Any] | None = CacheSettings._coerce(cache) # noqa: SLF001 self._kwargs = kwargs @@ -131,13 +127,17 @@ 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): + if ( + item.unresolvable_generic 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; " + f"parameterized generic annotation {item.unresolvable_generic!r} cannot be resolved by type; " "pass the value via the kwargs parameter or give the parameter a default" ), ) @@ -196,10 +196,6 @@ def _argument_resolution_error( 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) - def _can_call_positionally(self, plan: WiringPlan) -> bool: """Whether this creator can be called positionally under `plan`. @@ -215,13 +211,12 @@ def _can_call_positionally(self, plan: WiringPlan) -> bool: return False return not (names and self._has_positional_only_gap) - def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: + def _get_dependencies(self, registry: "ProvidersRegistry") -> 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).provider_kwargs # noqa: SLF001 + return registry.plan_for(self).provider_kwargs - def _iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]: + def _iter_validation_issues(self, registry: "ProvidersRegistry") -> typing.Iterable[Exception]: """Yield ArgumentResolutionError for parameters with no provider, no default, no static kwarg.""" - registry = container._providers_registry # noqa: SLF001 - plan = self._wiring_plan(registry) + plan = registry.plan_for(self) for name, item in plan.unwireable: yield self._argument_resolution_error(arg_name=name, item=item, registry=registry) diff --git a/modern_di/registries/context_registry.py b/modern_di/registries/context_registry.py deleted file mode 100644 index a2416fe5..00000000 --- a/modern_di/registries/context_registry.py +++ /dev/null @@ -1,13 +0,0 @@ -import typing - -from modern_di import types - - -class ContextRegistry: - __slots__ = ("context",) - - def __init__(self, context: dict[type[typing.Any], typing.Any]) -> None: - self.context = context - - def set_context(self, context_type: type[types.T], obj: types.T) -> None: - self.context[context_type] = obj diff --git a/modern_di/registries/providers_registry.py b/modern_di/registries/providers_registry.py index 3b8a52e1..b08e8cb4 100644 --- a/modern_di/registries/providers_registry.py +++ b/modern_di/registries/providers_registry.py @@ -132,12 +132,7 @@ def drop_resolvers(self) -> None: self._drop_resolvers() def register(self, provider_type: type, provider: AbstractProvider[typing.Any]) -> None: - with self._lock: - if provider_type in self._providers: - raise exceptions.DuplicateProviderTypeError(provider_type=provider_type) - self._providers[provider_type] = provider - provider._mark_registered() # noqa: SLF001 - self._invalidate() + self._add({provider_type: provider}, (provider,)) def add_providers(self, *args: AbstractProvider[typing.Any]) -> None: new_providers: dict[type, AbstractProvider[typing.Any]] = {} @@ -147,14 +142,21 @@ def add_providers(self, *args: AbstractProvider[typing.Any]) -> None: if provider.bound_type in new_providers: raise exceptions.DuplicateProviderTypeError(provider_type=provider.bound_type) new_providers[provider.bound_type] = provider - + self._add(new_providers, args) + + def _add( + self, + new_providers: dict[type, AbstractProvider[typing.Any]], + registered: tuple[AbstractProvider[typing.Any], ...], + ) -> None: + """Bind ``new_providers`` and latch every provider in ``registered``; a type already bound raises.""" with self._lock: for provider_type in new_providers: if provider_type in self._providers: raise exceptions.DuplicateProviderTypeError(provider_type=provider_type) self._providers.update(new_providers) - # Over `args`: a reference-only provider never enters `_providers` but is still compiled. - for provider in args: + # Over `registered`: a reference-only provider never enters `_providers` but is still compiled. + for provider in registered: provider._mark_registered() # noqa: SLF001 self._invalidate() diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 4e8db083..984bc89e 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -20,7 +20,7 @@ import typing from modern_di import exceptions, types -from modern_di.dependency_graph import redirect_hops +from modern_di.dependency_graph import redirect_hops, terminal_chain from modern_di.providers.abstract import AbstractProvider from modern_di.providers.alias import Alias from modern_di.providers.container_provider import container_provider @@ -38,7 +38,6 @@ Resolver: typing.TypeAlias = typing.Callable[[Container], typing.Any] _SCOPE_ERRORS = (exceptions.ScopeNotInitializedError, exceptions.ScopeSkippedError) -STEP_ERRORS = (exceptions.ResolutionError,) def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "ProvidersRegistry") -> "Resolver": @@ -76,11 +75,11 @@ def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "Provid name = [*edges][arg_lines[exc.__traceback__.tb_lineno]] if not exc.dependency_path: exc._name_parameter(name) - exc._prepend_step(resolution_step(), *redirect_hops(edges[name], target)) + exc._prepend_step(resolution_step(), *redirect_hops(edges[name], registry)) raise - except _STEP_ERRORS as exc: + except ResolutionError as exc: name = [*edges][arg_lines[exc.__traceback__.tb_lineno]] - exc._prepend_step(resolution_step(), *redirect_hops(edges[name], target)) + exc._prepend_step(resolution_step(), *redirect_hops(edges[name], registry)) raise """ @@ -92,7 +91,7 @@ def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "Provid if error is None: raise raise error from exc - except _STEP_ERRORS as exc: + except ResolutionError as exc: exc._prepend_step(resolution_step()) raise """ @@ -160,7 +159,7 @@ def _code(arity: int, names: tuple[str, ...] | None, static: bool, cached: bool) def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "Resolver": - plan = f._wiring_plan(registry) + plan = registry.plan_for(f) if plan.unwireable: return _compile_unwireable_factory(f, plan) positional = f._can_call_positionally(plan) @@ -182,11 +181,12 @@ def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "UNSET": types.UNSET, "partial": functools.partial, "_navigate": _navigate, - "_STEP_ERRORS": STEP_ERRORS, + "ResolutionError": exceptions.ResolutionError, "CreatorCallError": exceptions.CreatorCallError, "ContainerClosedError": exceptions.ContainerClosedError, "ContextValueNotSetError": exceptions.ContextValueNotSetError, "redirect_hops": redirect_hops, + "registry": registry, **{ f"r{i}": _argument_resolver(f, name, p, registry) for i, (name, p) in enumerate(plan.provider_kwargs.items()) @@ -211,21 +211,15 @@ def _argument_resolver( def _defaultless_context_terminal( - provider: "AbstractProvider[typing.Any] | None", registry: "ProvidersRegistry" + provider: "AbstractProvider[typing.Any]", registry: "ProvidersRegistry" ) -> "ContextProvider[typing.Any] | None": - """Follow un-overridden alias redirects to a `ContextProvider` with no `default=` and no override.""" - seen: set[int] = set() - while type(provider) is Alias and provider.provider_id not in seen: - if registry.overrides.fetch_override(provider.provider_id) is not types.UNSET: - return None - seen.add(provider.provider_id) - provider = registry.find_provider(provider._source_type) - if ( - type(provider) is ContextProvider - and provider.default is types.UNSET - and registry.overrides.fetch_override(provider.provider_id) is types.UNSET - ): - return provider + """Follow un-overridden redirects to a `ContextProvider` with no `default=` and no override.""" + chain = terminal_chain(provider, registry) + if any(registry.overrides.fetch_override(p.provider_id) is not types.UNSET for p in chain): + return None + terminal = chain[-1] + if type(terminal) is ContextProvider and terminal.default is types.UNSET: + return terminal return None @@ -290,7 +284,7 @@ def resolve(container: "Container") -> typing.Any: target = _navigate(container, scope, resolution_step) if target._closed: raise exceptions.ContainerClosedError(container_scope=target._scope) - context = target._context_registry.context + context = target._context # Not `.get(key, UNSET)`: that skips a dict subclass's `__contains__`/`__getitem__`. if context_type in context: return context[context_type] diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 18ee09d8..58e76b1f 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -18,7 +18,7 @@ class SignatureItem: member_types: list[type] = dataclasses.field(default_factory=list) is_nullable: bool = False default: object = UNSET - raw_annotation: object = None + unresolvable_generic: object = None is_keyword_only: bool = False @classmethod @@ -49,7 +49,7 @@ def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem": result["arg_type"] = non_none_members[0] elif origin is not None: - result["raw_annotation"] = type_ + result["unresolvable_generic"] = type_ elif isinstance(type_, (type, _NAMED_TYPE_FORMS)): result["arg_type"] = type_ @@ -135,7 +135,7 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: ) type_hints = {} - param_hints = {} + params = {} has_positional_only_gap = False accepts_any_kwargs = False for param_name, param in sig.parameters.items(): @@ -148,20 +148,20 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: if item is None: has_positional_only_gap = True continue - param_hints[param_name] = item + params[param_name] = item if is_class: return_sig = SignatureItem.from_type(creator) elif "return" in type_hints: return_sig = SignatureItem.from_type(type_hints["return"]) - if return_sig.raw_annotation is not None: - return_sig = SignatureItem(arg_type=typing.get_origin(return_sig.raw_annotation)) + if return_sig.unresolvable_generic is not None: + return_sig = SignatureItem(arg_type=typing.get_origin(return_sig.unresolvable_generic)) else: return_sig = SignatureItem() return ParsedCreator( return_type=return_sig, - params=param_hints, + params=params, has_positional_only_gap=has_positional_only_gap, accepts_any_kwargs=accepts_any_kwargs, ) diff --git a/tests/providers/test_alias.py b/tests/providers/test_alias.py index 626e83fb..713d845b 100644 --- a/tests/providers/test_alias.py +++ b/tests/providers/test_alias.py @@ -372,7 +372,7 @@ class _MutualAliasGroup(Group): def test_terminal_chain_handles_mutual_alias_cycle() -> None: # Mutual aliases: the walk must terminate via the `seen` guard and fall back to `a` itself. container = Container(scope=Scope.APP, groups=[_MutualAliasGroup]) - assert terminal_chain(_MutualAliasGroup.a, container) == [_MutualAliasGroup.a] + assert terminal_chain(_MutualAliasGroup.a, container._providers_registry) == [_MutualAliasGroup.a] # validate() also reports the cycle separately. with pytest.raises(exceptions.ValidationFailedError) as exc_info: container.validate() @@ -415,7 +415,7 @@ def test_alias_redirect_target_returns_source() -> None: container = Container(groups=[MyGroup]) source = container.find_provider(PostgresRepository) assert source is not None - target = MyGroup.abstract_repo._redirect_target(container) + target = MyGroup.abstract_repo._redirect_target(container._providers_registry) assert target is not None assert target.provider_id == source.provider_id @@ -425,7 +425,7 @@ class G(Group): abstract = providers.Alias(source_type=PostgresRepository, bound_type=AbstractRepository) container = Container(groups=[G]) - assert G.abstract._redirect_target(container) is None + assert G.abstract._redirect_target(container._providers_registry) is None # A genuine scope inversion through an alias must still raise, even measured against a custom diff --git a/tests/providers/test_context_provider.py b/tests/providers/test_context_provider.py index 17655778..f8d85015 100644 --- a/tests/providers/test_context_provider.py +++ b/tests/providers/test_context_provider.py @@ -1071,7 +1071,7 @@ class G(Group): out = providers.Factory(creator, bound_type=None, cache=cache) container = Container(groups=[G]) - assert G.out._can_call_positionally(G.out._wiring_plan(container._providers_registry)) is positional + assert G.out._can_call_positionally(container._providers_registry.plan_for(G.out)) is positional return container, G.out diff --git a/tests/registries/test_providers_registry.py b/tests/registries/test_providers_registry.py index f3a64f33..04d1ae00 100644 --- a/tests/registries/test_providers_registry.py +++ b/tests/registries/test_providers_registry.py @@ -288,7 +288,7 @@ def hold_open(owner: "providers.Factory[typing.Any]", *, registry: ProvidersRegi return plan monkeypatch.setattr(pr_mod.WiringPlan, "build", staticmethod(hold_open)) - worker = threading.Thread(target=lambda: svc._wiring_plan(registry)) + worker = threading.Thread(target=lambda: registry.plan_for(svc)) worker.start() try: assert built.wait(5), "plan build never reached the publication window" diff --git a/tests/test_container.py b/tests/test_container.py index ec56998a..bd38028f 100644 --- a/tests/test_container.py +++ b/tests/test_container.py @@ -27,6 +27,7 @@ ValidationFailedError, ) from modern_di.providers.abstract import AbstractProvider +from modern_di.registries.providers_registry import ProvidersRegistry from tests.helpers import cache_item @@ -223,10 +224,10 @@ class Top: class _CountingFactory(providers.Factory[Bottom]): __slots__ = () - def _get_dependencies(self, container: Container) -> dict[str, AbstractProvider[typing.Any]]: + def _get_dependencies(self, registry: ProvidersRegistry) -> dict[str, AbstractProvider[typing.Any]]: nonlocal call_count call_count += 1 - return super()._get_dependencies(container) + return super()._get_dependencies(registry) bottom_provider = _CountingFactory(creator=Bottom) @@ -367,7 +368,7 @@ class G(Group): svc = providers.Factory(creator=_NeedsMissing) container = Container(scope=Scope.APP, groups=[G]) - errors = collect_errors(container, container._providers_registry) + errors = collect_errors(container._providers_registry) # Root order is registration order (a, b, svc): the cycle closes while walking from root # `a`, so it is appended before `svc`'s missing dependency is reached. diff --git a/tests/test_dependency_graph.py b/tests/test_dependency_graph.py index 9e802c38..8b011d33 100644 --- a/tests/test_dependency_graph.py +++ b/tests/test_dependency_graph.py @@ -1,6 +1,8 @@ """Event-stream tests for ``dependency_graph.walk``: the module's test surface is the event SEQUENCE.""" -from modern_di import Container, Scope +import typing + +from modern_di import Scope from modern_di.dependency_graph import ( Cycle, DependenciesError, @@ -13,7 +15,14 @@ walk, ) from modern_di.group import Group -from modern_di.providers import Alias, Factory +from modern_di.providers import AbstractProvider, Alias, Factory +from modern_di.registries.providers_registry import ProvidersRegistry + + +def _registry(*providers: AbstractProvider[typing.Any]) -> ProvidersRegistry: + registry = ProvidersRegistry() + registry.add_providers(*providers) + return registry class Leaf: ... @@ -37,8 +46,8 @@ class G(Group): root = Factory(scope=Scope.APP, creator=Root) leaf = Factory(scope=Scope.APP, creator=Leaf) - c = Container(scope=Scope.APP, groups=[G]) - events = list(walk([G.root, G.leaf], c)) + registry = _registry(*G.get_providers()) + events = list(walk([G.root, G.leaf], registry)) kinds = [type(e).__name__ for e in events] assert kinds[0] == "NodeEntered" assert "Edge" in kinds @@ -51,8 +60,8 @@ class G(Group): root = Factory(scope=Scope.APP, creator=Root) leaf = Factory(scope=Scope.APP, creator=Leaf) - c = Container(scope=Scope.APP, groups=[G]) - events = list(walk([G.root], c)) + registry = _registry(*G.get_providers()) + events = list(walk([G.root], registry)) assert events == [ NodeEntered(G.root), Edge(G.root, "leaf", G.leaf), @@ -65,8 +74,8 @@ class G(Group): a = Factory(scope=Scope.APP, creator=CycA) b = Factory(scope=Scope.APP, creator=CycB) - c = Container(scope=Scope.APP, groups=[G]) - cycles = [e for e in walk([G.a], c) if isinstance(e, Cycle)] + registry = _registry(*G.get_providers()) + cycles = [e for e in walk([G.a], registry) if isinstance(e, Cycle)] assert cycles assert cycles[0].providers[0].provider_id == cycles[0].providers[-1].provider_id @@ -76,8 +85,8 @@ class G(Group): a = Factory(scope=Scope.APP, creator=CycA) b = Factory(scope=Scope.APP, creator=CycB) - c = Container(scope=Scope.APP, groups=[G]) - events = list(walk([G.a], c)) + registry = _registry(*G.get_providers()) + events = list(walk([G.a], registry)) assert events == [ NodeEntered(G.a), Edge(G.a, "b", G.b), @@ -101,8 +110,8 @@ class G(Group): right = Factory(scope=Scope.APP, creator=R) shared = Factory(scope=Scope.APP, creator=Shared) - c = Container(scope=Scope.APP, groups=[G]) - events = list(walk([G.left, G.right], c)) + registry = _registry(*G.get_providers()) + events = list(walk([G.left, G.right], registry)) # 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. @@ -115,9 +124,9 @@ class G(Group): root = Factory(scope=Scope.APP, creator=Root) leaf = Factory(scope=Scope.APP, creator=Leaf) - c = Container(scope=Scope.APP, groups=[G]) + registry = _registry(*G.get_providers()) # leaf appears as a dep of root (first root) AND as a later root; the later root is skipped. - events = list(walk([G.root, G.leaf], c)) + events = list(walk([G.root, G.leaf], registry)) assert sum(isinstance(e, NodeEntered) and e.provider is G.leaf for e in events) == 1 @@ -125,8 +134,8 @@ def test_find_cycle_from_returns_none_when_acyclic() -> None: class G(Group): leaf = Factory(scope=Scope.APP, creator=Leaf) - c = Container(scope=Scope.APP, groups=[G]) - assert find_cycle_from(G.leaf, c) is None + registry = _registry(*G.get_providers()) + assert find_cycle_from(G.leaf, registry) is None def test_find_cycle_from_returns_loop() -> None: @@ -134,8 +143,8 @@ class G(Group): a = Factory(scope=Scope.APP, creator=CycA) b = Factory(scope=Scope.APP, creator=CycB) - c = Container(scope=Scope.APP, groups=[G]) - cycle = find_cycle_from(G.a, c) + registry = _registry(*G.get_providers()) + cycle = find_cycle_from(G.a, registry) assert cycle == [G.a, G.b, G.a] @@ -151,9 +160,9 @@ class G(Group): mid = Alias(source_type=ChainTerminal, bound_type=ChainMid) top = Alias(source_type=ChainMid, bound_type=ChainTop) - c = Container(scope=Scope.APP, groups=[G]) - assert terminal_chain(G.top, c) == [G.top, G.mid, G.terminal] - assert effective_scope(G.top, c) == Scope.REQUEST + registry = _registry(*G.get_providers()) + assert terminal_chain(G.top, registry) == [G.top, G.mid, G.terminal] + assert effective_scope(G.top, registry) == Scope.REQUEST def test_terminal_chain_alias_cycle_falls_back_to_the_starting_provider() -> None: @@ -165,9 +174,9 @@ class G(Group): a = Alias(source_type=MutualY, bound_type=MutualX) b = Alias(source_type=MutualX, bound_type=MutualY) - c = Container(scope=Scope.APP, groups=[G]) - assert terminal_chain(G.a, c) == [G.a] - assert effective_scope(G.a, c) == G.a.scope + registry = _registry(*G.get_providers()) + assert terminal_chain(G.a, registry) == [G.a] + assert effective_scope(G.a, registry) == G.a.scope def test_walk_dangling_dep_emits_dependencies_error() -> None: @@ -179,8 +188,8 @@ class G(Group): # Bound under Marker, sourced from the unregistered Missing -> get_dependencies raises. alias = Alias(Missing, bound_type=Marker) - c = Container(scope=Scope.APP, groups=[G]) - events = list(walk([G.alias], c)) + registry = _registry(*G.get_providers()) + events = list(walk([G.alias], registry)) assert isinstance(events[0], NodeEntered) assert events[0].provider is G.alias assert isinstance(events[1], DependenciesError) @@ -205,10 +214,9 @@ def test_walk_emits_cycle_closed_through_kwargs_overlay() -> None: a = Factory(scope=Scope.APP, creator=KwCycA) b = Factory(scope=Scope.APP, creator=KwCycB, kwargs={"a": a}) # kwargs edge B -> A - c = Container(scope=Scope.APP) - c.add_providers(a, b) + registry = _registry(a, b) - events = list(walk([a], c)) + events = list(walk([a], registry)) 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"] @@ -230,7 +238,7 @@ def test_build_cycle_error_rotates_to_minimum_provider_id() -> None: # Seed the ring at the higher-id node, closing back to itself last -- the shape `Cycle.providers` # is in (first node repeated last), regardless of which provider the walk happened to start from. - error = build_cycle_error([second, first, second], Container(scope=Scope.APP)) + error = build_cycle_error([second, first, second], ProvidersRegistry()) # Rotated to the minimum-provider_id node (`first`), not left seeded at `second`. assert error.cycle_path == ["RingFirst", "RingSecond", "RingFirst"] diff --git a/tests/test_resolver_compiler.py b/tests/test_resolver_compiler.py index 7912b6a4..d2b6956d 100644 --- a/tests/test_resolver_compiler.py +++ b/tests/test_resolver_compiler.py @@ -17,11 +17,10 @@ import pytest from modern_di import Container, Group, Scope, exceptions, providers -from modern_di.dependency_graph import terminal_chain from modern_di.providers import ContextProvider from modern_di.providers.abstract import AbstractProvider from modern_di.registries.providers_registry import ProvidersRegistry -from modern_di.resolver_compiler import _defaultless_context_terminal, compile_resolver +from modern_di.resolver_compiler import compile_resolver from modern_di.wiring import WiringPlan @@ -80,7 +79,7 @@ def _make(a: _A, b: _B, c: _C) -> _Ordered: def _plan(registry: ProvidersRegistry, owner: "providers.Factory[object]") -> WiringPlan: """Build ``owner``'s wiring plan the way production does (via the registry memo).""" - return owner._wiring_plan(registry) + return registry.plan_for(owner) @dataclasses.dataclass(slots=True) @@ -663,41 +662,6 @@ def resolve_in_new_child(context: dict[type, object] | None) -> "typing.Callable assert unset_calls == set_calls -def _all_subclasses(cls: "type[AbstractProvider[typing.Any]]") -> "list[type[AbstractProvider[typing.Any]]]": - return [sub for direct in cls.__subclasses__() for sub in (direct, *_all_subclasses(direct))] - - -def test_context_fallback_walk_reaches_the_same_terminal_as_terminal_chain() -> None: - """INVARIANT: the compiler's context-fallback walk follows redirects exactly as `terminal_chain` does. - - `_defaultless_context_terminal` walks redirects through the registry instead of a container. A - new redirecting provider type must be taught to it, or a parameter behind that type stops - falling back. - """ - redirecting = { - cls - for cls in _all_subclasses(AbstractProvider) - if cls.__module__.startswith("modern_di.") and cls._redirect_target is not AbstractProvider._redirect_target - } - assert redirecting == {providers.Alias}, f"teach _defaultless_context_terminal to follow {redirecting}" - - class _Ctx: ... - - class _Mid: ... - - class _Top: ... - - class _G(Group): - ctx = providers.ContextProvider(_Ctx, scope=Scope.APP) - mid = providers.Alias(_Ctx, bound_type=_Mid) - top = providers.Alias(_Mid, bound_type=_Top) - - container = Container(scope=Scope.APP, groups=[_G]) - registry = container._providers_registry - for start in (_G.top, _G.mid, _G.ctx): - assert _defaultless_context_terminal(start, registry) is terminal_chain(start, container)[-1] is _G.ctx - - def test_overridden_alias_compiles_nothing_of_its_source() -> None: """INVARIANT: an override short-circuits before its provider's subtree is compiled. diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index b66297de..40b3b86f 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -21,14 +21,14 @@ class GenericClass(typing.Generic[types.T]): ... [ (int, SignatureItem(arg_type=int)), (typing.Annotated[int, None], SignatureItem(arg_type=int)), - (list[int], SignatureItem(raw_annotation=list[int])), - (dict[str, typing.Any], SignatureItem(raw_annotation=dict[str, typing.Any])), + (list[int], SignatureItem(unresolvable_generic=list[int])), + (dict[str, typing.Any], SignatureItem(unresolvable_generic=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(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], SignatureItem(unresolvable_generic=GenericClass[str])), (GenericClass[str] | None, SignatureItem(arg_type=GenericClass, is_nullable=True)), # `None` is the degenerate nullable: a union with zero non-None members. (type(None), SignatureItem(is_nullable=True)), @@ -147,7 +147,7 @@ def __init__(self, arg1: "WrongType", arg2: "int") -> None: ... # ty: ignore[un SignatureItem(is_nullable=True), { "arg1": SignatureItem(arg_type=str), - "arg2": SignatureItem(raw_annotation=tuple[int, ...], default=()), + "arg2": SignatureItem(unresolvable_generic=tuple[int, ...], default=()), }, ), ), From b2541eb7f7d7adb9e9a9aad64ee1c16c8cccfbb3 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 14:57:04 +0300 Subject: [PATCH 02/17] perf: keep cache items and creation order on the container (#605) `CacheRegistry` is gone. A container holds its cache items and creation order in two slots, and `fetch_cache_item`, `close_async` and `close_sync` are module functions in `cache_registry` over that dict and list. Both close functions now update the list in place, so the container no longer reaches into another object's private list to decide whether to close. Child construction allocates one object fewer and the cached template drops an attribute hop. Tests that counted cached items through the registry now check observable behaviour: the container repr, finalizer effects, or whether a resolve after reopening returns the same instance. A/B against the parent commit, 12 interleaved rounds, paired median: G6 -11.6%, G6b -9.9%, G7 -4.8%, G7b -3.3%, G13 -2.9%, C6 -6.2%; G1-G4, G8, G10, G11 within noise. --- modern_di/container.py | 19 ++--- modern_di/registries/cache_registry.py | 96 +++++++++++-------------- modern_di/resolver_compiler.py | 12 ++-- pyproject.toml | 3 - tests/helpers.py | 4 +- tests/providers/test_cached_factory.py | 34 +++++---- tests/providers/test_factory.py | 3 +- tests/registries/test_cache_registry.py | 20 +++--- tests/test_container.py | 4 +- tests/test_custom_scope.py | 4 +- tests/test_free_threading.py | 2 +- 11 files changed, 98 insertions(+), 103 deletions(-) diff --git a/modern_di/container.py b/modern_di/container.py index b82e5bc4..6e3c01ef 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -7,7 +7,8 @@ from modern_di.group import Group from modern_di.providers.abstract import AbstractProvider from modern_di.providers.container_provider import container_provider -from modern_di.registries.cache_registry import CacheRegistry +from modern_di.registries import cache_registry +from modern_di.registries.cache_registry import CacheItem from modern_di.registries.overrides_registry import OverrideHandle from modern_di.registries.providers_registry import ProvidersRegistry from modern_di.scope import Scope @@ -34,9 +35,10 @@ class Container: """ __slots__ = ( - "_cache_registry", + "_cache_items", "_closed", "_context", + "_creation_order", "_parent_container", "_providers_registry", "_scope", @@ -125,7 +127,8 @@ def _set_state( self._scope = scope self._parent_container = parent self._scope_map = scope_map - self._cache_registry = CacheRegistry() + self._cache_items: dict[int, CacheItem] = {} + self._creation_order: list[CacheItem] = [] self._context = copy.copy(context) if context is not None else {} self._providers_registry = providers_registry @@ -238,8 +241,8 @@ async def close_async(self) -> None: raises; the failures come back together as one :class:`~modern_di.exceptions.FinalizerError`. """ self._closed = True - if self._cache_registry._creation_order: # noqa: SLF001 - await self._cache_registry.close_async() + if self._creation_order: + await cache_registry.close_async(self._creation_order) def close_sync(self) -> None: """Mark this container closed, then run its sync finalizers, newest first. @@ -250,8 +253,8 @@ def close_sync(self) -> None: An async finalizer fails here and stays pending for a later :meth:`close_async`. """ self._closed = True - if self._cache_registry._creation_order: # noqa: SLF001 - self._cache_registry.close_sync() + if self._creation_order: + cache_registry.close_sync(self._creation_order) def override(self, provider: AbstractProvider[types.T], override_object: types.T) -> OverrideHandle[types.T]: """Apply an override immediately, tree-wide. @@ -287,7 +290,7 @@ def set_context(self, context_type: type[types.T], obj: types.T) -> None: def __repr__(self) -> str: n_providers = len(self._providers_registry) - n_cached = self._cache_registry.cached_count() + n_cached = sum(1 for item in self._cache_items.values() if item.cache is not types.UNSET) parent = self.parent_container.scope.name if self.parent_container else None return f"Container(scope={self.scope.name}, parent={parent}, providers={n_providers}, cached={n_cached})" diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index 926d2c20..da017822 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -77,56 +77,46 @@ def close_sync(self) -> None: self.clear() -class CacheRegistry: - __slots__ = ("_creation_order", "_items") - - def __init__(self) -> None: - self._items: dict[int, CacheItem] = {} - self._creation_order: list[CacheItem] = [] - - 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[typing.Any]) -> CacheItem: - """Return the cache item for a cached ``provider``, creating it on first use.""" - # Get before setdefault: a bare setdefault builds a throwaway CacheItem on every hit. - provider_id = provider._provider_id # noqa: SLF001 - item = self._items.get(provider_id) - if item is not None: - return item - settings = typing.cast("CacheSettings[typing.Any]", provider._cache_settings) # noqa: SLF001 - return self._items.setdefault(provider_id, CacheItem(settings=settings)) - - def mark_created(self, cache_item: CacheItem) -> None: - """Record creation completion; close finalizes in reverse of this order (LIFO).""" - self._creation_order.append(cache_item) - - async def close_async(self) -> None: - finalizer_errors: list[Exception] = [] - for cache_item in reversed(self._creation_order): - if cache_item.settings.finalizer is None: - cache_item.clear() - continue - try: - await cache_item.close_async() - except Exception as e: # noqa: BLE001 - finalizer_errors.append(e) - self._creation_order.clear() - if finalizer_errors: - raise exceptions.FinalizerError(finalizer_errors=finalizer_errors, is_async=True) - - def close_sync(self) -> None: - finalizer_errors: list[Exception] = [] - remaining: list[CacheItem] = [] - for cache_item in reversed(self._creation_order): - try: - cache_item.close_sync() - except exceptions.AsyncFinalizerInSyncCloseError as e: - finalizer_errors.append(e) - remaining.append(cache_item) - except Exception as e: # noqa: BLE001 - finalizer_errors.append(e) - remaining.reverse() - self._creation_order = remaining - if finalizer_errors: - raise exceptions.FinalizerError(finalizer_errors=finalizer_errors, is_async=False) +def fetch_cache_item(items: dict[int, CacheItem], provider: Factory[typing.Any]) -> CacheItem: + """Return the cache item in ``items`` for a cached ``provider``, creating it on first use.""" + # Get before setdefault: a bare setdefault builds a throwaway CacheItem on every hit. + provider_id = provider._provider_id # noqa: SLF001 + item = items.get(provider_id) + if item is not None: + return item + settings = typing.cast("CacheSettings[typing.Any]", provider._cache_settings) # noqa: SLF001 + return items.setdefault(provider_id, CacheItem(settings=settings)) + + +async def close_async(creation_order: list[CacheItem]) -> None: + """Close every item newest first and empty ``creation_order``; failures raise together at the end.""" + finalizer_errors: list[Exception] = [] + for cache_item in reversed(creation_order): + if cache_item.settings.finalizer is None: + cache_item.clear() + continue + try: + await cache_item.close_async() + except Exception as e: # noqa: BLE001 + finalizer_errors.append(e) + creation_order.clear() + if finalizer_errors: + raise exceptions.FinalizerError(finalizer_errors=finalizer_errors, is_async=True) + + +def close_sync(creation_order: list[CacheItem]) -> None: + """Close every item newest first; an async finalizer stays in ``creation_order`` for a later async close.""" + finalizer_errors: list[Exception] = [] + remaining: list[CacheItem] = [] + for cache_item in reversed(creation_order): + try: + cache_item.close_sync() + except exceptions.AsyncFinalizerInSyncCloseError as e: + finalizer_errors.append(e) + remaining.append(cache_item) + except Exception as e: # noqa: BLE001 + finalizer_errors.append(e) + remaining.reverse() + creation_order[:] = remaining + if finalizer_errors: + raise exceptions.FinalizerError(finalizer_errors=finalizer_errors, is_async=False) diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 984bc89e..338eea66 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -9,7 +9,7 @@ ``ProvidersRegistry.drop_resolvers``). Why a template and not shared helpers: see docs/adr/0001-resolver-hot-path-generated-source.md. -The template reaches into `Container._scope`, `Container._scope_map` and `CacheRegistry._items` to stay +The template reaches into the `Container` slots `_scope`, `_scope_map`, `_cache_items` and `_creation_order` to stay within that frame budget. No linter sees the template, so those reaches are outside every suppression here. """ @@ -26,6 +26,7 @@ from modern_di.providers.container_provider import container_provider from modern_di.providers.context_provider import ContextProvider from modern_di.providers.factory import Factory +from modern_di.registries.cache_registry import fetch_cache_item if typing.TYPE_CHECKING: @@ -106,16 +107,16 @@ def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "Provid + "\ndef resolve(container):\n" + _NAVIGATE + """\ - cache_registry = target._cache_registry - cache_item = cache_registry._items.get(pid) + cache_items = target._cache_items + cache_item = cache_items.get(pid) if cache_item is None: - cache_item = cache_registry.fetch_cache_item(provider) + cache_item = fetch_cache_item(cache_items, provider) cached = cache_item.cache if cached is not UNSET: return cached value, created = cache_item.get_or_create(partial(build, target), create) if created: - cache_registry.mark_created(cache_item) + target._creation_order.append(cache_item) return value """ ) @@ -179,6 +180,7 @@ def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "arg_lines": arg_lines, "static": plan.static_kwargs, "UNSET": types.UNSET, + "fetch_cache_item": fetch_cache_item, "partial": functools.partial, "_navigate": _navigate, "ResolutionError": exceptions.ResolutionError, diff --git a/pyproject.toml b/pyproject.toml index b760ba1c..e073cf98 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -82,15 +82,12 @@ isort.no-lines-before = ["standard-library", "local-folder"] # White-box tests that assert on container and registry internals with no public equivalent. "tests/helpers.py" = ["SLF001"] "tests/providers/test_alias.py" = ["SLF001"] -"tests/providers/test_cached_factory.py" = ["SLF001"] "tests/providers/test_context_provider.py" = ["SLF001"] "tests/providers/test_factory.py" = ["SLF001"] -"tests/registries/test_cache_registry.py" = ["SLF001"] "tests/registries/test_providers_registry.py" = ["SLF001"] "tests/test_container.py" = ["SLF001"] "tests/test_custom_scope.py" = ["SLF001"] "tests/test_exception_pickling.py" = ["SLF001"] -"tests/test_free_threading.py" = ["SLF001"] "tests/test_resolver_compiler.py" = ["SLF001"] "tests/test_wiring.py" = ["SLF001"] # Guard-tier benchmarks are pytest tests too: they assert on the measured ratios. diff --git a/tests/helpers.py b/tests/helpers.py index 81a58c4b..e8265af4 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -2,9 +2,9 @@ from modern_di import Container from modern_di.providers import Factory -from modern_di.registries.cache_registry import CacheItem +from modern_di.registries.cache_registry import CacheItem, fetch_cache_item def cache_item(container: Container, provider: Factory[typing.Any]) -> CacheItem: """Return `container`'s cache item for the cached `provider`, creating it if needed.""" - return container._cache_registry.fetch_cache_item(provider) + return fetch_cache_item(container._cache_items, provider) diff --git a/tests/providers/test_cached_factory.py b/tests/providers/test_cached_factory.py index 4040cf19..9fa1185f 100644 --- a/tests/providers/test_cached_factory.py +++ b/tests/providers/test_cached_factory.py @@ -114,7 +114,7 @@ async def test_request_cached_factory() -> None: assert instance3 is instance4 assert instance1 is not instance3 - cache_item = request_container._cache_registry.fetch_cache_item(MyGroup.request_cached) + item = cache_item(request_container, MyGroup.request_cached) with pytest.raises(FinalizerError) as exc_info: request_container.close_sync() @@ -122,10 +122,10 @@ async def test_request_cached_factory() -> None: assert len(exc_info.value.exceptions) == 1 assert isinstance(exc_info.value.exceptions[0], AsyncFinalizerInSyncCloseError) - assert cache_item.cache is not UNSET # preserved — user can still recover via close_async + assert item.cache is not UNSET # preserved — user can still recover via close_async await request_container.close_async() - assert cache_item.cache is UNSET + assert item.cache is UNSET def test_app_cached_factory_resolves_once_across_request_children() -> None: @@ -368,7 +368,7 @@ class NoneGroup(Group): app_container.resolve_provider(NoneGroup.none_resource) assert call_count == 1 # cached after first call, not re-created - assert app_container._cache_registry.cached_count() == 1 + assert "cached=1" in repr(app_container) app_container.close_sync() assert cleaned_up == [None] @@ -628,16 +628,18 @@ class G(Group): ) container = Container(scope=Scope.APP, groups=[G]) - container.resolve(_AwaitableFinSvc) + instance = container.resolve(_AwaitableFinSvc) with pytest.raises(FinalizerError) as exc: container.close_sync() (inner,) = exc.value.exceptions assert isinstance(inner, AsyncFinalizerInSyncCloseError) assert inner.instance_type is _AwaitableFinSvc - assert container._cache_registry.cached_count() == 1 + container.open() + assert container.resolve(_AwaitableFinSvc) is instance future.set_result(None) await container.close_async() - assert container._cache_registry.cached_count() == 0 + container.open() + assert container.resolve(_AwaitableFinSvc) is not instance class _First: ... @@ -676,8 +678,10 @@ class FinGroup(Group): assert len(errors) == 1 assert _Built.count == 1 - assert container._cache_registry.cached_count() == 0 assert container.closed is True + container.open() + container.resolve(_Built) + assert _Built.count == 1 + 1 # rebuilt: the close dropped the cached instance async def test_resolve_from_a_finalizer_during_close_async_raises_container_closed() -> None: @@ -703,8 +707,10 @@ class FinGroup(Group): assert len(errors) == 1 assert _Built.count == 1 - assert container._cache_registry.cached_count() == 0 assert container.closed is True + container.open() + container.resolve(_Built) + assert _Built.count == 1 + 1 # rebuilt: the close dropped the cached instance def _failing_finalizer(_: _First) -> None: @@ -728,7 +734,6 @@ class FailGroup(Group): container.close_sync() assert [type(e) for e in exc.value.exceptions] == [ValueError] - assert container._cache_registry.cached_count() == 0 container.open() assert container.resolve(_First) is not stale @@ -744,7 +749,6 @@ class FailGroup(Group): await container.close_async() assert [type(e) for e in exc.value.exceptions] == [ValueError] - assert container._cache_registry.cached_count() == 0 container.open() assert container.resolve(_First) is not stale @@ -767,8 +771,8 @@ class SlowGroup(Group): second = providers.Factory(creator=_Second, cache=providers.CacheSettings(finalizer=slow_finalizer)) container = Container(groups=[SlowGroup]) - container.resolve(_First) - container.resolve(_Second) + first = container.resolve(_First) + second = container.resolve(_Second) with pytest.raises(TimeoutError): await asyncio.wait_for(container.close_async(), timeout=0.05) @@ -781,7 +785,9 @@ class SlowGroup(Group): assert events == ["second", "first"] assert slow_calls == [0, 1] - assert container._cache_registry.cached_count() == 0 + container.open() + assert container.resolve(_First) is not first + assert container.resolve(_Second) is not second async def test_close_async_after_a_cancelled_close_does_not_refinalize_closed_items() -> None: diff --git a/tests/providers/test_factory.py b/tests/providers/test_factory.py index b8c4b27a..1e58483a 100644 --- a/tests/providers/test_factory.py +++ b/tests/providers/test_factory.py @@ -430,8 +430,7 @@ def test_creator_raising_mid_creation_caches_nothing_and_retry_succeeds() -> Non container = Container(scope=Scope.APP, groups=[_FlakyGroup]) with pytest.raises(RuntimeError, match="boom"): container.resolve(_FlakySvc) - expected_cached_after_failure = 1 # only the dep cached; failed svc not cached - assert container._cache_registry.cached_count() == expected_cached_after_failure + assert "cached=1" in repr(container) # only the dep cached; failed svc not cached retried = container.resolve(_FlakySvc) assert isinstance(retried, _FlakySvc) container.close_sync() diff --git a/tests/registries/test_cache_registry.py b/tests/registries/test_cache_registry.py index 99eca76b..d3d2cb33 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/registries/test_cache_registry.py @@ -5,7 +5,7 @@ import pytest from modern_di.providers import CacheSettings, Factory -from modern_di.registries.cache_registry import CacheItem, CacheRegistry +from modern_di.registries.cache_registry import CacheItem, close_async, fetch_cache_item from modern_di.types import UNSET @@ -108,21 +108,21 @@ def test_get_or_create_releases_the_item_lock() -> None: def test_each_cache_item_owns_its_lock() -> None: - registry = CacheRegistry() - first = registry.fetch_cache_item(Factory(creator=lambda: 1, bound_type=int, cache=True)) - second = registry.fetch_cache_item(Factory(creator=lambda: "", bound_type=str, cache=True)) + cache_items: dict[int, CacheItem] = {} + first = fetch_cache_item(cache_items, Factory(creator=lambda: 1, bound_type=int, cache=True)) + second = fetch_cache_item(cache_items, Factory(creator=lambda: "", bound_type=str, cache=True)) assert first.lock is not second.lock def test_concurrent_fetches_of_one_provider_share_one_item() -> None: n = 8 - registry = CacheRegistry() + cache_items: dict[int, CacheItem] = {} provider = Factory(creator=lambda: 1, bound_type=int, cache=True) barrier = threading.Barrier(n, timeout=5) def fetch() -> CacheItem: barrier.wait() - return registry.fetch_cache_item(provider) + return fetch_cache_item(cache_items, provider) with ThreadPoolExecutor(max_workers=n) as pool: items = [f.result(timeout=5) for f in [pool.submit(fetch) for _ in range(n)]] @@ -140,17 +140,15 @@ async def _recording(self: CacheItem) -> None: monkeypatch.setattr(CacheItem, "close_async", _recording) finalized: list[object] = [] - registry = CacheRegistry() plain = CacheItem(settings=CacheSettings(), cache="plain") persistent = CacheItem(settings=CacheSettings(clear_cache=False), cache="persistent") with_finalizer = CacheItem(settings=CacheSettings(finalizer=finalized.append), cache="finalized") - for item in (plain, persistent, with_finalizer): - registry.mark_created(item) + creation_order = [plain, persistent, with_finalizer] - await registry.close_async() + await close_async(creation_order) assert awaited == [with_finalizer] assert finalized == ["finalized"] assert plain.cache is UNSET assert persistent.cache == "persistent" - assert registry._creation_order == [] + assert creation_order == [] diff --git a/tests/test_container.py b/tests/test_container.py index bd38028f..ad0f6439 100644 --- a/tests/test_container.py +++ b/tests/test_container.py @@ -580,7 +580,7 @@ def test_resolving_through_closed_parent_via_open_child_raises() -> None: child.resolve(_PersistentBroker) assert exc.value.container_scope is Scope.APP assert app.closed is True - assert app._cache_registry.cached_count() == 0 + assert "cached=0" in repr(app) async def test_async_context_manager_reopens() -> None: @@ -614,7 +614,7 @@ class G(Group): with pytest.raises(ContainerClosedError): container.resolve(str) assert calls == [] - assert container._cache_registry.cached_count() == 0 + assert "cached=0" in repr(container) def test_reopen_rebuilds_a_value_that_close_cleared() -> None: diff --git a/tests/test_custom_scope.py b/tests/test_custom_scope.py index d8fc4242..86e37f87 100644 --- a/tests/test_custom_scope.py +++ b/tests/test_custom_scope.py @@ -244,7 +244,7 @@ class ConflictingGroup(Group): session.resolve(TenantService) assert exc.value.provider_scope is ConflictingScope.LOWER_THAN_REQUEST assert exc.value.dependency_path[0].scope is ConflictingScope.LOWER_THAN_REQUEST - assert session._cache_registry.cached_count() == 0 + assert "cached=0" in repr(session) @pytest.mark.parametrize("cache", [False, True]) @@ -258,7 +258,7 @@ class ConflictingGroup(Group): request.resolve(TenantService) assert exc.value.provider_scope is ConflictingScope.LOWER_THAN_REQUEST assert exc.value.dependency_path[0].scope is ConflictingScope.LOWER_THAN_REQUEST - assert session._cache_registry.cached_count() == 0 + assert "cached=0" in repr(session) @pytest.mark.parametrize("scope", [Scope.SESSION, Scope.REQUEST], ids=["SESSION", "REQUEST"]) diff --git a/tests/test_free_threading.py b/tests/test_free_threading.py index 6dd3c87c..24334d5f 100644 --- a/tests/test_free_threading.py +++ b/tests/test_free_threading.py @@ -101,4 +101,4 @@ def worker() -> None: assert len(raised) == n assert container.closed is True - assert container._cache_registry.cached_count() == 0 + assert "cached=0" in repr(container) From 8cb88b5b28c1bb01ea35f32950df9382b79c8678 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 14:58:38 +0300 Subject: [PATCH 03/17] perf: build a child's scope map with copy() and one assignment (#605) `{**parent_map, parent_scope: parent}` builds the dict through the generic unpacking path. A `copy()` of the parent's map plus one item assignment does the same work faster. A/B against the parent commit, 12 interleaved rounds, paired median: G6 -8.5%, G6b -7.5%, C6 -4.2%, G7 -1.1%. --- modern_di/container.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/modern_di/container.py b/modern_di/container.py index 6e3c01ef..aa2788b7 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -107,7 +107,9 @@ def build_child_container( cls = type(self) child: typing.Self = cls.__new__(cls) # Ancestors only: a `scope: self` entry is a reference cycle that refcounting never frees. - child._set_state(scope, self, {**self._scope_map, self._scope: self}, context, self._providers_registry) + scope_map = self._scope_map.copy() + scope_map[self._scope] = self + child._set_state(scope, self, scope_map, context, self._providers_registry) return child def _set_state( From a34b2ce553df8223429315de49877e13849f42a7 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 14:59:03 +0300 Subject: [PATCH 04/17] perf: copy a plain dict context with dict.copy (#605) A child built with `context={...}` copied it through `copy.copy`, which looks up the type's copier before it lands on `dict.copy`. A plain dict now calls `dict.copy` directly; any other mapping still goes through `copy.copy`, so a dict subclass keeps its type. A/B against the parent commit, 12 interleaved rounds, paired median: C6 (request cycle with a context value) -4.7%; G6 unchanged, as it passes no context. --- modern_di/container.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/modern_di/container.py b/modern_di/container.py index aa2788b7..50d717c8 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -131,7 +131,12 @@ def _set_state( self._scope_map = scope_map self._cache_items: dict[int, CacheItem] = {} self._creation_order: list[CacheItem] = [] - self._context = copy.copy(context) if context is not None else {} + if context is None: + self._context = {} + elif type(context) is dict: + self._context = context.copy() + else: + self._context = copy.copy(context) self._providers_registry = providers_registry def find_container(self, scope: enum.IntEnum) -> typing.Self: From c2c0a0d4a0cac037f170f9bca24266981bfe2aff Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:00:18 +0300 Subject: [PATCH 05/17] perf: check a child's scope only when the caller passes one (#605) `_set_state` checked the scope type and the parent ordering for every container, including an auto-scoped child whose scope came from `next_deeper` and so is always a deeper member of the same enum. The checks now sit at the two entry points: the root constructor checks the type, and `build_child_container` checks an explicit scope. `_set_state` stays the one place that sets a container's slots and does nothing else. A/B against the parent commit, 12 interleaved rounds, paired median: G6b -17.2%, G6 -4.4%, C6 -2.0%; G7 and G7b within noise. Tried and dropped: `object.__new__(cls)` instead of `cls.__new__(cls)` (G6 -0.2%, noise). Setting the child's slots inline in `build_child_container` is worth another ~9% on G6 and G6b, but it would duplicate `_set_state`, so it is not done. --- modern_di/container.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/modern_di/container.py b/modern_di/container.py index 50d717c8..2065b123 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -62,6 +62,8 @@ def __init__( Raises :class:`~modern_di.exceptions.InvalidScopeTypeError` when ``scope`` is not an ``IntEnum``. """ + if not isinstance(scope, enum.IntEnum): + raise exceptions.InvalidScopeTypeError(scope_value=scope) self._set_state(scope, None, {}, context, ProvidersRegistry()) self._providers_registry.register(Container, container_provider) if groups: @@ -104,6 +106,10 @@ def build_child_container( scope = next_deeper(self._scope) if scope is None: raise exceptions.MaxScopeReachedError(parent_scope=self._scope) + elif not isinstance(scope, enum.IntEnum): + raise exceptions.InvalidScopeTypeError(scope_value=scope) + elif scope <= self._scope: + raise exceptions.InvalidChildScopeError(parent_scope=self._scope, child_scope=scope) cls = type(self) child: typing.Self = cls.__new__(cls) # Ancestors only: a `scope: self` entry is a reference cycle that refcounting never frees. @@ -120,11 +126,7 @@ def _set_state( context: dict[type[typing.Any], typing.Any] | None, providers_registry: ProvidersRegistry, ) -> None: - """Check ``scope`` and set every slot of an open container; ``parent`` is ``None`` for a root.""" - if not isinstance(scope, enum.IntEnum): - raise exceptions.InvalidScopeTypeError(scope_value=scope) - if parent is not None and scope <= parent.scope: - raise exceptions.InvalidChildScopeError(parent_scope=parent.scope, child_scope=scope) + """Set every slot of an open container; ``parent`` is ``None`` for a root.""" self._closed = False self._scope = scope self._parent_container = parent From 0b3090a067571284868da3f4ba5b167f8ff663e1 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:00:39 +0300 Subject: [PATCH 06/17] perf: skip the awaitable check when a finalizer returns None (#605) `inspect.isawaitable` is a Python-level function that runs on every finalizer result, and a sync finalizer almost always returns `None`. A `None` check in front of it skips the call for that case. A/B against the parent commit, 12 interleaved rounds, paired median: G13 -14.7%, G7b -9.5%; G7 within noise (its finalizer is async). --- modern_di/registries/cache_registry.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index da017822..18101160 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -50,7 +50,7 @@ async def close_async(self) -> None: if (finalizer := self._pending_finalizer()) is not None: try: result = finalizer(self.cache) - if inspect.isawaitable(result): + if result is not None and inspect.isawaitable(result): await result except Exception: self.clear() @@ -68,7 +68,7 @@ def close_sync(self) -> None: except Exception: self.clear() raise - if inspect.isawaitable(result): + if result is not None and inspect.isawaitable(result): if inspect.iscoroutine(result): result.close() # suppress "never awaited" warning raise exceptions.AsyncFinalizerInSyncCloseError(instance_type=type(self.cache)) From d38570e1ed7c7b87a119f4870a99914a2fc08547 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:01:52 +0300 Subject: [PATCH 07/17] perf: run finalizers directly in the async close loop (#605) `close_async` awaited `CacheItem.close_async` for every item that had a finalizer, a coroutine per item, and short-circuited the items without one. The loop now calls the owed finalizer itself and awaits only an awaitable result, so a sync finalizer costs no coroutine, and one condition replaces both the short-circuit and `CacheItem.close_async`. The sync and async loops now have the same shape: finalize what is owed, collect failures, clear. A/B against the parent commit, 12 interleaved rounds, paired median: G7 -4.2%; G13b (no finalizers), G13 and G7b within noise. --- modern_di/registries/cache_registry.py | 31 +++++++++---------------- tests/registries/test_cache_registry.py | 17 ++++---------- 2 files changed, 15 insertions(+), 33 deletions(-) diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index 18101160..05bf9a1f 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -46,19 +46,6 @@ def _pending_finalizer(self) -> typing.Callable[[typing.Any], typing.Awaitable[N """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 (finalizer := self._pending_finalizer()) is not None: - try: - result = finalizer(self.cache) - if result is not None and inspect.isawaitable(result): - await result - except Exception: - self.clear() - raise - self.finalized = True - - self.clear() - def close_sync(self) -> None: if (finalizer := self._pending_finalizer()) is not None: if self.settings._is_async_finalizer: # noqa: SLF001 @@ -92,13 +79,17 @@ async def close_async(creation_order: list[CacheItem]) -> None: """Close every item newest first and empty ``creation_order``; failures raise together at the end.""" finalizer_errors: list[Exception] = [] for cache_item in reversed(creation_order): - if cache_item.settings.finalizer is None: - cache_item.clear() - continue - try: - await cache_item.close_async() - except Exception as e: # noqa: BLE001 - finalizer_errors.append(e) + finalizer = cache_item.settings.finalizer + if finalizer is not None and cache_item.cache is not types.UNSET and not cache_item.finalized: + try: + result = finalizer(cache_item.cache) + if result is not None and inspect.isawaitable(result): + await result + except Exception as e: # noqa: BLE001 + finalizer_errors.append(e) + else: + cache_item.finalized = True + cache_item.clear() creation_order.clear() if finalizer_errors: raise exceptions.FinalizerError(finalizer_errors=finalizer_errors, is_async=True) diff --git a/tests/registries/test_cache_registry.py b/tests/registries/test_cache_registry.py index d3d2cb33..c5dfefbe 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/registries/test_cache_registry.py @@ -2,8 +2,6 @@ import typing from concurrent.futures import ThreadPoolExecutor -import pytest - from modern_di.providers import CacheSettings, Factory from modern_di.registries.cache_registry import CacheItem, close_async, fetch_cache_item from modern_di.types import UNSET @@ -130,25 +128,18 @@ def fetch() -> CacheItem: assert all(item is items[0] for item in items) -async def test_close_async_awaits_only_items_with_a_finalizer(monkeypatch: pytest.MonkeyPatch) -> None: - awaited: list[CacheItem] = [] - original = CacheItem.close_async - - async def _recording(self: CacheItem) -> None: - awaited.append(self) - await original(self) - - monkeypatch.setattr(CacheItem, "close_async", _recording) +async def test_close_async_runs_only_owed_finalizers_and_empties_the_order() -> None: finalized: list[object] = [] plain = CacheItem(settings=CacheSettings(), cache="plain") persistent = CacheItem(settings=CacheSettings(clear_cache=False), cache="persistent") with_finalizer = CacheItem(settings=CacheSettings(finalizer=finalized.append), cache="finalized") - creation_order = [plain, persistent, with_finalizer] + never_built = CacheItem(settings=CacheSettings(finalizer=finalized.append)) + creation_order = [plain, persistent, with_finalizer, never_built] await close_async(creation_order) - assert awaited == [with_finalizer] assert finalized == ["finalized"] assert plain.cache is UNSET assert persistent.cache == "persistent" + assert with_finalizer.cache is UNSET assert creation_order == [] From ef41e8e85cff5cc7c0c7c49d78b27ac8b74c722f Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:02:32 +0300 Subject: [PATCH 08/17] perf: pass the target to get_or_create instead of a partial (#605) A cache miss built `functools.partial(build, target)` only to call it once under the item lock. `get_or_create(build, target, create)` takes the target as an argument, which allocates nothing and still keeps `target` out of a closure cell on the warm path. A/B against the parent commit, 12 interleaved rounds, paired median: G13b -10.6%, G13 -9.8%, G7b -6.3%, G7 -5.8%, G8b -4.1%. --- modern_di/registries/cache_registry.py | 10 +++++---- modern_di/resolver_compiler.py | 3 +-- tests/registries/test_cache_registry.py | 28 ++++++++++++------------- tests/test_resolver_compiler.py | 2 +- 4 files changed, 22 insertions(+), 21 deletions(-) diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index 05bf9a1f..b4e1d9d9 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -7,6 +7,7 @@ from modern_di.providers import CacheSettings, Factory +_T = typing.TypeVar("_T") _R = typing.TypeVar("_R") _V = typing.TypeVar("_V") @@ -25,12 +26,13 @@ def clear(self) -> None: def get_or_create( self, - resolve: typing.Callable[[], _R], + build: typing.Callable[[_T], _R], + target: _T, create: typing.Callable[[_R], _V], ) -> tuple[_V, bool]: - """Return the memoized singleton, or resolve-and-create it once under this item's lock. + """Return the memoized singleton, or ``create(build(target))`` it once under this item's lock. - A hit never takes the lock. A miss resolves and creates under it, so concurrent misses + A hit never takes the lock. A miss builds and creates under it, so concurrent misses build the value and its dependencies once. `created` is True only for the caller that built. """ if self.cache is not types.UNSET: @@ -38,7 +40,7 @@ def get_or_create( with self.lock: if self.cache is not types.UNSET: return self.cache, False - value = create(resolve()) + value = create(build(target)) self.cache = value return value, True diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 338eea66..6016aa5b 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -114,7 +114,7 @@ def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "Provid cached = cache_item.cache if cached is not UNSET: return cached - value, created = cache_item.get_or_create(partial(build, target), create) + value, created = cache_item.get_or_create(build, target, create) if created: target._creation_order.append(cache_item) return value @@ -181,7 +181,6 @@ def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "static": plan.static_kwargs, "UNSET": types.UNSET, "fetch_cache_item": fetch_cache_item, - "partial": functools.partial, "_navigate": _navigate, "ResolutionError": exceptions.ResolutionError, "CreatorCallError": exceptions.CreatorCallError, diff --git a/tests/registries/test_cache_registry.py b/tests/registries/test_cache_registry.py index c5dfefbe..bfe174c4 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/registries/test_cache_registry.py @@ -27,12 +27,12 @@ def try_acquire() -> None: return acquired == [True] -def test_get_or_create_miss_resolves_and_creates_once_under_the_item_lock() -> None: +def test_get_or_create_miss_builds_and_creates_once_under_the_item_lock() -> None: item = _item() - calls = {"resolve": 0, "create": 0} + calls = {"build": 0, "create": 0} - def resolve() -> dict[str, typing.Any]: - calls["resolve"] += 1 + def build(_: object) -> dict[str, typing.Any]: + calls["build"] += 1 assert not _acquirable_from_another_thread(item) return {"x": 1} @@ -40,27 +40,27 @@ def create(kwargs: dict[str, typing.Any]) -> tuple[str, dict[str, typing.Any]]: calls["create"] += 1 return ("made", kwargs) - value, created = item.get_or_create(resolve=resolve, create=create) + value, created = item.get_or_create(build=build, target=None, create=create) assert created is True assert value == ("made", {"x": 1}) assert item.cache == ("made", {"x": 1}) - assert calls == {"resolve": 1, "create": 1} + assert calls == {"build": 1, "create": 1} -def test_get_or_create_hit_returns_cache_without_resolving() -> None: +def test_get_or_create_hit_returns_cache_without_building() -> None: item = _item() item.cache = "cached" - def resolve() -> object: - msg = "resolve must not run on a cache hit" + def build(_: object) -> object: + msg = "build must not run on a cache hit" raise AssertionError(msg) def create(_: object) -> str: msg = "create must not run on a cache hit" raise AssertionError(msg) - value, created = item.get_or_create(resolve=resolve, create=create) + value, created = item.get_or_create(build=build, target=None, create=create) assert created is False assert value == "cached" @@ -83,15 +83,15 @@ def test_get_or_create_double_checks_under_the_lock() -> None: item = _item() item.lock = _LosingRaceLock(item) # ty: ignore[invalid-assignment] - def resolve() -> object: - msg = "resolve must not run when another thread already stored the value" + def build(_: object) -> object: + msg = "build must not run when another thread already stored the value" raise AssertionError(msg) def create(_: object) -> str: msg = "create must not run when another thread already stored the value" raise AssertionError(msg) - value, created = item.get_or_create(resolve=resolve, create=create) + value, created = item.get_or_create(build=build, target=None, create=create) assert created is False assert value == "won-the-race" @@ -100,7 +100,7 @@ def create(_: object) -> str: def test_get_or_create_releases_the_item_lock() -> None: item = _item() - value, created = item.get_or_create(resolve=lambda: 0, create=lambda _: "v") + value, created = item.get_or_create(build=lambda _: 0, target=None, create=lambda _: "v") assert (value, created) == ("v", True) assert _acquirable_from_another_thread(item) diff --git a/tests/test_resolver_compiler.py b/tests/test_resolver_compiler.py index d2b6956d..e2b2beed 100644 --- a/tests/test_resolver_compiler.py +++ b/tests/test_resolver_compiler.py @@ -853,7 +853,7 @@ class _G(Group): def test_cached_resolver_has_no_cell_on_the_warm_path() -> None: """INVARIANT: the cached-factory resolver has no cell variables. - The cold-miss thunk must stay a `functools.partial`, never a lambda closing over `target`: a + The cold miss must pass `target` to `get_or_create`, never a lambda closing over it: a closure promotes `target` to a cell, so MAKE_CELL runs in the prologue on every call -- including the warm hit that returns two lines later. Nothing else in the suite catches a revert. """ From b87166f30e4126c79bad8e2ccecbe94b0dfc0d84 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:03:09 +0300 Subject: [PATCH 09/17] perf: decide a factory's positional-call names once, at construction (#605) `_can_call_positionally` rebuilt the parameter-name tuple and rescanned for keyword-only parameters each time a resolver compiled, and a resolver compiles once per container tree, so a suite that builds a container per test paid it per test. The part that depends only on the signature is now computed in `Factory.__init__`; compiling compares the plan's names against it. A/B against the parent commit, 12 interleaved rounds, paired median: G8 -8.0%, G8b -6.3%. The work moves to construction: `Factory(...)` +1.2%, paid once per provider. --- modern_di/providers/factory.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index 5946f20d..4b06717c 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -45,9 +45,9 @@ class Factory(AbstractProvider[types.T_co]): "_cache_settings", "_cached_definition_site", "_creator", - "_has_positional_only_gap", "_kwargs", "_params", + "_positional_names", ) def __init__( # noqa: PLR0913 @@ -64,7 +64,13 @@ def __init__( # noqa: PLR0913 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 + names = tuple(parsed.params) + self._positional_names: tuple[str, ...] | None = ( + None + if (names and parsed.has_positional_only_gap) + or any(item.is_keyword_only for item in parsed.params.values()) + else names + ) super().__init__(scope=scope, bound_type=bound_type, inferred_bound_type=parsed.return_type.arg_type) self._creator = creator self._cache_settings: CacheSettings[typing.Any] | None = CacheSettings._coerce(cache) # noqa: SLF001 @@ -202,14 +208,9 @@ 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 plan.static_kwargs: - return False - names = tuple(self._params) - if tuple(plan.provider_kwargs) != names: - return False - if any(item.is_keyword_only for item in self._params.values()): + if plan.static_kwargs or self._positional_names is None: return False - return not (names and self._has_positional_only_gap) + return tuple(plan.provider_kwargs) == self._positional_names def _get_dependencies(self, registry: "ProvidersRegistry") -> dict[str, "AbstractProvider[typing.Any]"]: """Return parameter name → dependency provider: a pure registry lookup, no scope or cache touched.""" From 1e64d3a03a18e13ea653fb67faa29edcf6e171ad Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:03:34 +0300 Subject: [PATCH 10/17] perf: start each factory resolver's globals from a copied base dict (#605) Every compiled factory rebuilt the same eight module-level entries of its `exec` namespace from a dict literal, plus a merged comprehension for the argument resolvers. The constant entries now live in one module-level dict that each compile copies, then adds the factory's own names. A/B against the parent commit, 12 interleaved rounds, paired median: G8 -4.5%, G8b -3.7%. --- modern_di/resolver_compiler.py | 47 +++++++++++++++++----------------- 1 file changed, 24 insertions(+), 23 deletions(-) diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 6016aa5b..3c711783 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -170,29 +170,18 @@ def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> bool(plan.static_kwargs), f.cache_settings is not None, ) - namespace: dict[str, typing.Any] = { - "provider": f, - "pid": f.provider_id, - "scope": f.scope, - "creator": f._creator, - "resolution_step": f._resolution_step, - "edges": plan.provider_kwargs, - "arg_lines": arg_lines, - "static": plan.static_kwargs, - "UNSET": types.UNSET, - "fetch_cache_item": fetch_cache_item, - "_navigate": _navigate, - "ResolutionError": exceptions.ResolutionError, - "CreatorCallError": exceptions.CreatorCallError, - "ContainerClosedError": exceptions.ContainerClosedError, - "ContextValueNotSetError": exceptions.ContextValueNotSetError, - "redirect_hops": redirect_hops, - "registry": registry, - **{ - f"r{i}": _argument_resolver(f, name, p, registry) - for i, (name, p) in enumerate(plan.provider_kwargs.items()) - }, - } + namespace = _FACTORY_GLOBALS.copy() + namespace["provider"] = f + namespace["pid"] = f.provider_id + namespace["scope"] = f.scope + namespace["creator"] = f._creator + namespace["resolution_step"] = f._resolution_step + namespace["edges"] = plan.provider_kwargs + namespace["arg_lines"] = arg_lines + namespace["static"] = plan.static_kwargs + namespace["registry"] = registry + for i, (name, p) in enumerate(plan.provider_kwargs.items()): + namespace[f"r{i}"] = _argument_resolver(f, name, p, registry) exec(code, namespace) # noqa: S102 # the source is a fixed template; user data enters only via `namespace` resolve = namespace["resolve"] resolve.__qualname__ = f"resolve[{f.display_name}]" @@ -310,3 +299,15 @@ def _navigate( except _SCOPE_ERRORS as exc: exc._prepend_step(resolution_step()) raise + + +_FACTORY_GLOBALS: dict[str, typing.Any] = { + "UNSET": types.UNSET, + "fetch_cache_item": fetch_cache_item, + "_navigate": _navigate, + "ResolutionError": exceptions.ResolutionError, + "CreatorCallError": exceptions.CreatorCallError, + "ContainerClosedError": exceptions.ContainerClosedError, + "ContextValueNotSetError": exceptions.ContextValueNotSetError, + "redirect_hops": redirect_hops, +} From 4541309a1409b74893542555036ac68f726ad811 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:05:00 +0300 Subject: [PATCH 11/17] perf: dispatch validation events on their exact type (#605) `collect_errors` matched each event with class patterns, which run an `isinstance` and a `__match_args__` unpack per case. It now compares `type(event)` against the four event classes, which are marked `typing.final` so `ty` narrows each branch and the last one is known to be a `Cycle`. `walk` skips a root it already visited before starting a generator for it, and `_walk_from` drops its own, now unreachable, visited check. A/B against the parent commit, 12 interleaved rounds, paired median: G10 -8.6%, G11 -8.6%. Tried and dropped: memoizing each parent's effective scope across its edges. It helps the wide graph (G11 -1.6%) and costs the chain, where every parent has one edge (G10 +2.9%). --- modern_di/dependency_graph.py | 64 +++++++++++++++++------------------ pyproject.toml | 2 +- 2 files changed, 33 insertions(+), 33 deletions(-) diff --git a/modern_di/dependency_graph.py b/modern_di/dependency_graph.py index 81c11f7f..955c1bca 100644 --- a/modern_di/dependency_graph.py +++ b/modern_di/dependency_graph.py @@ -17,12 +17,14 @@ from modern_di.registries.providers_registry import ProvidersRegistry +@typing.final class NodeEntered(NamedTuple): """A provider was reached for the first time, before its dependencies are read.""" provider: "AbstractProvider[typing.Any]" +@typing.final class Edge(NamedTuple): """A dependency edge from ``parent`` to ``dep`` via parameter ``name``.""" @@ -31,12 +33,14 @@ class Edge(NamedTuple): dep: "AbstractProvider[typing.Any]" +@typing.final class Cycle(NamedTuple): """A cycle closing on the active path; ``providers`` repeats the first node last.""" providers: "list[AbstractProvider[typing.Any]]" +@typing.final class DependenciesError(NamedTuple): """Reading ``provider``'s dependencies raised; it is then treated as having none.""" @@ -108,7 +112,8 @@ def walk( visiting: set[int] = set() visited: set[int] = set() for root in roots: - yield from _walk_from(root, registry, visiting, visited) + if root.provider_id not in visited: + yield from _walk_from(root, registry, visiting, visited) def find_cycle_from( @@ -128,10 +133,7 @@ def _walk_from( 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 - + """Explicit-stack DFS from an unvisited ``start``.""" path: list[AbstractProvider[typing.Any]] = [] stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = [] yield from _enter(start, registry, visiting, path, stack) @@ -179,33 +181,31 @@ def collect_errors(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 walk(registry, registry): - match event: - case NodeEntered(provider): - errors.extend(provider._iter_validation_issues(registry)) # noqa: SLF001 - case DependenciesError(_, error): - errors.append(error) - case Edge(parent, name, dep): - dependency_chain = terminal_chain(dep, registry) - dependency_scope = dependency_chain[-1].scope - parent_scope = effective_scope(parent, registry) - if dependency_scope > parent_scope: - errors.append( - exceptions.InvalidScopeDependencyError( - provider=parent, - parameter_name=name, - dependency_chain=dependency_chain, - ) + if type(event) is Edge: + parent, name, dep = event + dependency_chain = terminal_chain(dep, registry) + dependency_scope = dependency_chain[-1].scope + parent_scope = effective_scope(parent, registry) + if dependency_scope > parent_scope: + errors.append( + exceptions.InvalidScopeDependencyError( + provider=parent, + parameter_name=name, + dependency_chain=dependency_chain, ) - elif dependency_scope == parent_scope and dependency_scope is not parent_scope: - errors.append( - exceptions.ScopeEnumMismatchError( - provider=parent, - parameter_name=name, - dependency_chain=dependency_chain, - ) + ) + elif dependency_scope == parent_scope and dependency_scope is not parent_scope: + errors.append( + exceptions.ScopeEnumMismatchError( + provider=parent, + parameter_name=name, + dependency_chain=dependency_chain, ) - case Cycle(providers): - errors.append(build_cycle_error(providers, registry)) - case _: - typing.assert_never(event) + ) + elif type(event) is NodeEntered: + errors.extend(event.provider._iter_validation_issues(registry)) # noqa: SLF001 + elif type(event) is DependenciesError: + errors.append(event.error) + else: + errors.append(build_cycle_error(event.providers, registry)) return errors diff --git a/pyproject.toml b/pyproject.toml index e073cf98..c28aa1df 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -104,5 +104,5 @@ asyncio_default_fixture_loop_scope = "function" [tool.coverage] run.source = ["modern_di"] run.branch = true -report.exclude_also = ["if typing.TYPE_CHECKING:", "typing.assert_never\\("] +report.exclude_also = ["if typing.TYPE_CHECKING:"] report.fail_under = 100 From 6d1df2f4e2d019a6bd94a2a248636af073128a45 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:05:23 +0300 Subject: [PATCH 12/17] perf: read a plain class annotation without get_origin (#605) Most creator parameters are annotated with a plain class. `from_type` ran every such annotation through `typing.get_origin` and the union and generic checks before landing on the plain-class branch. An annotation whose type is exactly `type` (and is not `NoneType`) now returns there directly; a class with a custom metaclass, a generic alias, a union or a `NewType` takes the full path as before. A/B against the parent commit, 12 interleaved rounds, paired median: `Factory(...)` on a two-parameter dataclass -5.3%, on a three-parameter plain class -5.9%. --- modern_di/types_parser.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 58e76b1f..1826d343 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -23,6 +23,8 @@ class SignatureItem: @classmethod def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem": + if type(type_) is type and type_ is not types.NoneType: + return cls(arg_type=type_, default=default) if type_ is types.NoneType: # The degenerate nullable: the union branch below would take it for a plain type and # try to resolve `NoneType` from the registry. From 0e7b522c65ac9e439aa3e3772f304d33e21c4e7f Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:09:20 +0300 Subject: [PATCH 13/17] perf: read a plain class's signature straight off its __init__ (#605) `inspect.signature(SomeClass)` checks the metaclass for `__call__`, looks `__new__` and `__init__` up statically, unwraps both and walks the MRO before it reads `__init__` as a bound method. `_signature` takes that last step directly when the earlier ones cannot change the answer: the metaclass is `type`, the class has no `__signature__` or `__wrapped__`, and the first MRO entry that defines `__new__` or `__init__` defines a plain-function `__init__` with no `__wrapped__` and no `__new__` beside it. Anything else goes to `inspect.signature` as before. `test_class_signature_matches_inspect_signature` checks the result, or the `ValueError`, against `inspect.signature` for plain, inherited, dataclass, generated, `*args`-first, self-less, wrapped, `__signature__`, NamedTuple, `__new__`-only, `__init__`-over-`__new__`, metaclass, Generic, exception, staticmethod and no-`__init__` classes, plus non-class creators. It passes on 3.11, 3.12, 3.13, 3.14 and 3.15. A/B against the parent commit, 12 interleaved rounds, paired median: `Factory(...)` on a dataclass -21.8%, on a plain class -19.1%; G8 and G8b within noise, since their providers are built outside the timed call. --- modern_di/types_parser.py | 22 +++++- tests/test_types_parser.py | 135 ++++++++++++++++++++++++++++++++++++- 2 files changed, 155 insertions(+), 2 deletions(-) diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 1826d343..7084d920 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -114,9 +114,29 @@ def _class_type_hints(creator: type) -> dict[str, typing.Any]: return typing.get_type_hints(creator.__init__) +def _signature(creator: typing.Callable[..., typing.Any]) -> inspect.Signature: + """Return ``inspect.signature(creator)``, read straight off ``__init__`` when that is where it comes from. + + Only for a class of metaclass ``type`` with no ``__signature__`` or ``__wrapped__``, whose first + MRO entry defining ``__new__`` or ``__init__`` defines a plain-function ``__init__`` alone. + """ + if ( + type(creator) is type + and getattr(creator, "__signature__", None) is None + and not hasattr(creator, "__wrapped__") + ): + owner = next( + base.__dict__ for base in creator.__mro__ if "__new__" in base.__dict__ or "__init__" in base.__dict__ + ) + init = owner.get("__init__") + if "__new__" not in owner and type(init) is types.FunctionType and not hasattr(init, "__wrapped__"): + return inspect.signature(types.MethodType(init, creator)) + return inspect.signature(creator) + + def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator: try: - sig = inspect.signature(creator) + sig = _signature(creator) except (ValueError, TypeError): return ParsedCreator( return_type=SignatureItem.from_type(typing.cast(type, creator)), diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index 40b3b86f..b88f5a37 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -1,12 +1,13 @@ import dataclasses import functools +import inspect import sys import typing import pytest from modern_di import Container, Group, Scope, exceptions, providers, types -from modern_di.types_parser import SignatureItem, parse_creator +from modern_di.types_parser import SignatureItem, _signature, parse_creator class GenericClass(typing.Generic[types.T]): ... @@ -455,3 +456,135 @@ def test_union_return_type_is_silent_when_bound_type_is_known( creator: typing.Callable[..., typing.Any], bound_type: type | None ) -> None: providers.Factory(creator, bound_type=bound_type) + + +class _PlainInit: + def __init__(self, dep: _Dep, label: str = "x") -> None: + """Never called: these classes exist for their signatures.""" + + +class _InheritsInit(_PlainInit): + """Inherits its ``__init__``.""" + + +@dataclasses.dataclass(slots=True) +class _SlotsDataclass: + dep: _Dep + + +class _GeneratedInit: + __init__ = eval("lambda self, dep: None") # noqa: S307 # an attrs-style generated __init__ + + +class _StarArgsInit: + def __init__(*args: object, dep: _Dep) -> None: + """Never called.""" + + +class _NoSelfInit: + def __init__() -> None: + """Never called.""" + + +class _WrappedInit: + @functools.wraps(_PlainInit.__init__) + def __init__(self, *args: object, **kwargs: object) -> None: + """Never called.""" + + +class _SignatureOverride: + __signature__ = inspect.Signature([inspect.Parameter("dep", inspect.Parameter.KEYWORD_ONLY, annotation=_Dep)]) + + def __init__(self, **kwargs: object) -> None: + """Never called.""" + + +class _InitOverNew(_NewOnlyCreator): + def __init__(self, dep: _Dep, label: str) -> None: + """Never called.""" + + +class _InitBesideNew: + def __new__(cls, *args: object, **kwargs: object) -> typing.Self: # noqa: ARG004 + return super().__new__(cls) + + def __init__(self, dep: _Dep) -> None: + """Never called.""" + + +class _CallingMeta(type): + def __call__(cls, other: _OtherDep) -> object: + return other + + +class _MetaCall(metaclass=_CallingMeta): + def __init__(self, dep: _Dep) -> None: + """Never called.""" + + +class _PlainMeta(type): + """A metaclass with no ``__call__``.""" + + +class _MetaNoCall(metaclass=_PlainMeta): + def __init__(self, dep: _Dep) -> None: + """Never called.""" + + +class _GenericInit(typing.Generic[types.T]): + def __init__(self, dep: _Dep) -> None: + """Never called.""" + + +class _InitError(Exception): + def __init__(self, dep: _Dep) -> None: + """Never called.""" + + +class _StaticInit: + __init__ = staticmethod(lambda dep: None) # noqa: ARG005 + + +class _NoInit: + """Defines neither ``__new__`` nor ``__init__``.""" + + +@pytest.mark.parametrize( + "creator", + [ + _PlainInit, + _InheritsInit, + SomeDataClass, + _SlotsDataclass, + _GeneratedInit, + _StarArgsInit, + _NoSelfInit, + _WrappedInit, + _SignatureOverride, + _NamedTupleCreator, + _NewOnlyCreator, + _NewOnlySubclass, + _InitOverNew, + _InitBesideNew, + _MetaCall, + _MetaNoCall, + _GenericInit, + _InitError, + _StaticInit, + _NoInit, + _UserId, + _greeter_of(int), + functools.partial(_PlainInit, label="y"), + ], +) +def test_class_signature_matches_inspect_signature(creator: typing.Callable[..., object]) -> None: + """INVARIANT: `_signature` reads the same signature `inspect.signature` does, error included.""" + try: + expected: object = inspect.signature(creator) + except ValueError as exc: + expected = (ValueError, str(exc)) + try: + actual: object = _signature(creator) + except ValueError as exc: + actual = (ValueError, str(exc)) + assert actual == expected From 04ba8d98c44bc0754fa8407895f152a76afade62 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 15:16:01 +0300 Subject: [PATCH 14/17] docs: regenerate the comparative tables after #605 (#605) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Re-ran `just bench-report` (5 runs, Apple M2, CPython 3.14.7, the same rival versions) at the last perf commit in this series. C4 fell from 2.28 to 1.85 µs and now beats dishka (0.88, was 1.12); C6 fell from 1.09 µs to 806 ns. C1-C3 did not move. "What moved in this publication" gains a column for the previous 4.0 run, and the performance history gains a paragraph for --- AGENTS.md | 4 +- docs/introduction/performance.md | 113 +++++++++++++++++++------------ 2 files changed, 71 insertions(+), 46 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 6f86afa5..9d4b7e13 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -46,7 +46,9 @@ Every module under `modern_di/` is named for what it does; read it. What a singl `docs/troubleshooting/`, enforced by `tests/test_docs_slug_census.py`). Add a message, a glyph, or a 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. + `overrides_registry` are shared tree-wide. Cache and context are per-container and live in the + container's own slots: `cache_registry` holds `CacheItem` and the module functions that fetch an + item and close a container's items, and the context is a plain dict. - `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. diff --git a/docs/introduction/performance.md b/docs/introduction/performance.md index 90901954..77580494 100644 --- a/docs/introduction/performance.md +++ b/docs/introduction/performance.md @@ -59,7 +59,7 @@ for why that matters and which cells it moved. ## Results -Measured 2026-10-05 with modern-di `main` at `f300c2e` (the 4.0 code, unreleased when measured) +Measured 2026-10-06 with modern-di at the end of the #605 perf series (the 4.0 code, unreleased when measured) on an Apple M2 (macOS 26.6.2), CPython 3.14.7 with the GIL, median over 5 runs (ratios paired within each run); the footnote under each table bounds the across-run dispersion of each side's own median. Rival versions: dishka 1.10.1, dependency-injector 4.49.1, that-depends 4.1.0, @@ -86,67 +86,67 @@ as a verdict. | Scenario | modern-di | vs dependency-injector | vs that-depends | |---|---|---|---| -| C1 transient | 242 ns ±0.4% | **0.41** ±0.5% | **0.56** ±0.4% | -| C2 warm singleton | 146 ns ±0.6% | 2.20 ±0.6% | 1.70 ±1.6% | -| C3 deep chain (6) | 760 ns ±1.4% | **0.35** ±1.3% | **0.51** ±1.3% | +| C1 transient | 242 ns ±0.4% | **0.40** ±0.2% | **0.56** ±0.4% | +| C2 warm singleton | 143 ns ±0.3% | 2.20 ±0.6% | 1.69 ±0.3% | +| C3 deep chain (6) | 756 ns ±1.7% | **0.35** ±2.1% | **0.51** ±1.5% | -_Across-run IQR of each side's own median (5 runs): modern-di ≤1.4%, rivals ≤0.9%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ +_Across-run IQR of each side's own median (5 runs): modern-di ≤1.7%, rivals ≤1.3%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ ### By-type resolution | Scenario | modern-di | vs dishka | vs wireup | |---|---|---|---| -| C1 transient | 239 ns ±0.1% | **0.70** ±0.5% | **0.80** ±0.1% | -| C2 warm singleton | 142 ns ±0.3% | **0.60** ±1.2% | 1.41 ±1.1% | -| C3 deep chain (6) | 764 ns ±1.4% | 1.21 ±1.9% | **0.83** ±4.3% | +| C1 transient | 239 ns ±0.6% | **0.70** ±1.1% | **0.80** ±0.5% | +| C2 warm singleton | 140 ns ±0.6% | **0.59** ±1.0% | 1.40 ±0.6% | +| C3 deep chain (6) | 772 ns ±2.7% | 1.22 ±2.4% | **0.83** ±3.4% | -_Across-run IQR of each side's own median (5 runs): modern-di ≤1.4%, rivals ≤2.8%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ +_Across-run IQR of each side's own median (5 runs): modern-di ≤2.7%, rivals ≤1.1%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ ### Request lifecycle (batched, published per request) | Scenario | modern-di | vs dependency-injector | vs that-depends | vs dishka | vs wireup | |---|---|---|---|---|---| -| C4 request lifecycle | 2.28 µs ±0.4% | **0.02** ±0.9% | **0.18** ±0.4% | 1.12 ±0.8% | **0.14** ±0.6% | +| C4 request lifecycle | 1.85 µs ±0.5% | **0.02** ±2.1% | **0.15** ±0.7% | **0.88** ±0.1% | **0.12** ±0.9% | -_Across-run IQR of each side's own median (5 runs): modern-di ≤0.4%, rivals ≤1.5%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ +_Across-run IQR of each side's own median (5 runs): modern-di ≤0.5%, rivals ≤1.8%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ ### Per-request context | Scenario | modern-di | vs dependency-injector | vs that-depends | vs dishka | vs wireup | |---|---|---|---|---|---| -| C6 context | 1.09 µs ±0.1% | **0.28** ±0.4% | **0.36** ±1.2% | **0.88** ±1.3% | **0.78** ±2.3% | +| C6 context | 806 ns ±0.5% | **0.21** ±0.7% | **0.27** ±0.9% | **0.65** ±2.2% | **0.57** ±0.7% | -_Across-run IQR of each side's own median (5 runs): modern-di ≤0.1%, rivals ≤1.6%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ +_Across-run IQR of each side's own median (5 runs): modern-di ≤0.5%, rivals ≤1.8%. The ± on each ratio cell is a different quantity: the spread of the paired per-run ratios._ ## What the numbers show -- Against `dependency-injector`, modern-di is faster by reference on C1 (**0.41**) and +- Against `dependency-injector`, modern-di is faster by reference on C1 (**0.40**) and C3 (**0.35**), and far faster on the batched C4 request lifecycle (**0.02**). dependency-injector's C4 body calls `init_resources()`/`shutdown_resources()` every cycle in addition to resolving; the suite doesn't decompose how much of its per-request cost is that lifecycle work versus the resolve itself, so C4 should be read as a whole-lifecycle comparison, not an isolated resolve (see the caveat below). dependency-injector is still - faster on C2 warm-singleton (2.20, an implied ~66 ns cache hit against modern-di's 146 ns): + faster on C2 warm-singleton (2.20, an implied ~65 ns cache hit against modern-di's 143 ns): its hit is a C-level slot read on a Cython-compiled core, where modern-di's is a Python dict lookup plus a slot read on the cache item. Pure Python does not reach ~66 ns, so this cell is expected to stay above 1.0 however much of modern-di's own overhead is removed. - Against `that-depends` (4.1.0), modern-di leads by reference on C1 (**0.56**) and C3 (**0.51**). The C1 series across publications is 1.08, 1.12, 0.98, 0.98, 0.97, 0.89, 0.91, - 0.65, 0.58, 0.58, now 0.56. that-depends remains faster on C2 warm-singleton (1.70); the suite + 0.65, 0.58, 0.58, 0.56, now 0.56. that-depends remains faster on C2 warm-singleton (1.69); the suite does not decompose its `resolve_sync` cache-hit path, so no mechanism is asserted for the remaining gap. - Against the two `exec`-codegen frameworks, modern-di leads most of the by-type table. - modern-di is faster than `dishka` on C1 (**0.70**) and C2 (**0.60**), and faster than `wireup` - on C1 (**0.80**) and C3 (**0.83**). dishka keeps its lead on C3 (1.21), the deepest graph, and - wireup keeps C2 (1.41). Since #470 modern-di also generates its resolvers from a source + modern-di is faster than `dishka` on C1 (**0.70**) and C2 (**0.59**), and faster than `wireup` + on C1 (**0.80**) and C3 (**0.83**). dishka keeps its lead on C3 (1.22), the deepest graph, and + wireup keeps C2 (1.40). Since #470 modern-di also generates its resolvers from a source template, so both sides run one generated frame per node, and the suite does not decompose what dishka does differently on a six-node chain. No mechanism is asserted for that cell. - By-type resolution carries no surcharge. `Container.resolve` memoizes type → resolver directly, - so the by-type and by-reference cells differ by 3-4 ns, with by-type the faster of the two on C1 + so the by-type and by-reference cells differ by 3-16 ns, with by-type the faster of the two on C1 and C2 and the slower on C3. The two tables measure the same resolve; only the rival set differs. - On C6 (per-request context) modern-di is faster than all four rivals: `dependency-injector` - (**0.28**), `that-depends` (**0.36**), `dishka` (**0.88**) and `wireup` (**0.78**). In the 3.x + (**0.21**), `that-depends` (**0.27**), `dishka` (**0.65**) and `wireup` (**0.57**). In the 3.x publication it was slower than dishka (1.15) and level with wireup (1.02). The cells are not on one basis: each framework supplies the request value through its own idiom, and two of those are structural analogs rather than equivalents (see the caveat below). No mechanism is asserted @@ -154,32 +154,38 @@ _Across-run IQR of each side's own median (5 runs): modern-di ≤0.1%, rivals - On C4 (request lifecycle), the corrected batching does not *remove* the ~35 µs asyncio floor, it amortizes it. The guard tier's `test_g7c_event_loop_floor_control` times the same batch shape with an empty body and puts the residual at ~0.3 µs per request still inside every - C4 cell (~14% of modern-di's C4 figure), shared identically by all five frameworks. With the - floor amortized, dishka is measurably faster than modern-di here (1.12); modern-di remains - far faster than that-depends, dependency-injector, and wireup on this scenario. + C4 cell (~16% of modern-di's C4 figure), shared identically by all five frameworks. With the + floor amortized, modern-di is faster than dishka here (**0.88**), which led this cell (1.12) + before #605, and far faster than that-depends, dependency-injector, and wireup. ### What moved in this publication -4.0 moved C6 the most. C6 builds a REQUEST child seeded with a context value, resolves a handler -that needs that value and an APP dependency, and closes the child. It takes #557's context -resolver, #585's cheaper child build and close, and #596's cross-scope check. C4 builds a child, -first-resolves one request-scoped cached factory with an async finalizer and closes the child, so -it gets #585's cheaper child build and pays for #597's lock allocation. - -| | 3.x (`630de77`) | 4.0 (`f300c2e`) | -|---|---|---| -| C1 / C2 / C3 by reference, modern-di | 253 / 151 / 789 ns | 242 / 146 / 760 ns | -| C4 request lifecycle, modern-di | 2.29 µs | 2.28 µs | -| C4 vs that-depends / dishka / wireup | **0.18** / 1.08 / **0.14** | **0.18** / 1.12 / **0.14** | -| C6 context, modern-di | 1.45 µs | 1.09 µs | -| C6 vs dependency-injector / that-depends | **0.38** / **0.48** | **0.28** / **0.36** | -| C6 vs dishka / wireup | 1.15 / 1.02 | **0.88** / **0.78** | - -C6 fell 25%, close to what the guard tier predicts: G9, C6's guard twin, fell 28% with #557 and another 5.9% -with #585. That moved modern-di ahead of dishka and wireup on this scenario. C4 did not move. On -G7, C4's guard twin, #585 measured −2.9% and #597 +4.6%, so the two roughly cancel. The dishka C4 -ratio went from 1.08 to 1.12 because dishka's own cell is about 3% lower than the 3.x ratio -implies, while modern-di's stayed put. +4.0 moved C4 and C6, the two per-request scenarios. C6 builds a REQUEST child seeded with a +context value, resolves a handler that needs that value and an APP dependency, and closes the +child. It takes #557's context resolver, #585's and #605's cheaper child build and close, and +#596's cross-scope check. C4 builds a child, first-resolves one request-scoped cached factory with +an async finalizer and closes the child, so it gets the cheaper child build and close from #585 +and #605 and pays for #597's lock allocation. The middle column is the previous 4.0 publication, +before #605. + +| | 3.x (`630de77`) | 4.0 before #605 (`f300c2e`) | 4.0 with #605 | +|---|---|---|---| +| C1 / C2 / C3 by reference, modern-di | 253 / 151 / 789 ns | 242 / 146 / 760 ns | 242 / 143 / 756 ns | +| C4 request lifecycle, modern-di | 2.29 µs | 2.28 µs | 1.85 µs | +| C4 vs that-depends / dishka / wireup | **0.18** / 1.08 / **0.14** | **0.18** / 1.12 / **0.14** | **0.15** / **0.88** / **0.12** | +| C6 context, modern-di | 1.45 µs | 1.09 µs | 806 ns | +| C6 vs dependency-injector / that-depends | **0.38** / **0.48** | **0.28** / **0.36** | **0.21** / **0.27** | +| C6 vs dishka / wireup | 1.15 / 1.02 | **0.88** / **0.78** | **0.65** / **0.57** | + +Before #605, C6 fell 25% from 3.x, close to what the guard tier predicts: G9, C6's guard twin, +fell 28% with #557 and another 5.9% with #585. That moved modern-di ahead of dishka and wireup on +this scenario. C4 did not move then. On G7, C4's guard twin, #585 measured −2.9% and #597 +4.6%, +so the two roughly cancel. + +#605 then took C4 down 19% and C6 down 26%. Its guard-tier figures, measured against the commit +before it, are in the history below: −16.6% on G7 and −21.3% on a C6-shaped request cycle. That +moved modern-di ahead of dishka on C4 (1.12 to **0.88**). dishka's own C4 cell, as the ratios +imply it, stayed within about 4% across the three runs, so the change in that ratio is modern-di's. The C1-C3 cells fell 3-4% on both the by-reference and by-type sides, while the rival cells implied by the 3.x ratios stayed within about 2% of this run. Every generated `Factory` resolver @@ -312,7 +318,7 @@ coroutine that only did that (−8% on a request cycle closing ten such items). publication at `630de77`, C6 fell 10% and C4 5%, the sum of the first two; the third has no cell on this page. -Several 4.0 changes moved these numbers. The figures in the next three paragraphs are guard-tier +Several 4.0 changes moved these numbers. The figures in the next four paragraphs are guard-tier numbers from the PR that made each change, measured against the commit before it. [What moved in this publication](#what-moved-in-this-publication) reads the comparative cells that moved. @@ -343,6 +349,23 @@ to the first resolve of a cached factory in each container. A request that resol request-scoped cached factory pays about 120 ns to allocate its lock: +7.7% on G7b, one request cycle with a sync close, and +4.6% on G7. +#605 went after the per-request cycle and the cold path again, as a series of small changes, each +measured against the commit before it. A container now holds its cache items, its creation order +and its context in its own slots instead of in two registry objects, so a child build allocates +two objects fewer. The child's ancestor map is copied instead of rebuilt by unpacking, a plain +dict context is copied with `dict.copy`, and an auto-scoped child skips the scope checks its +scope cannot fail. On the close path, a finalizer that returns `None` skips the +`inspect.isawaitable` call, the async close loop calls finalizers itself instead of awaiting a +coroutine per item, and a cache miss passes its target container instead of building a +`functools.partial`. On the cold path, a factory decides its positional-call names once, each +resolver's globals start from a copied dict, `validate()` dispatches its events on their exact +type, and building a `Factory` reads a plain class's signature off its `__init__` and a plain-class +annotation without `typing.get_origin`. Against the commit before the series: −31% on a child +build (G6), −37% with an automatic scope (G6b), −24% on G7b, −17% on G7, −27% on ten sync +finalizers (G13), −12% on a cold first resolve (G8), −9 to −11% on `validate()` (G10, G11) and +about −23% on building a `Factory`. A child container takes 504 bytes instead of 584. Warm +resolves (G1-G4) moved less than 2% either way. + ## Reproduce it yourself ```bash From 93fc00aa9fa6c8b0d0d0237332f78206c43828e9 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 20:09:51 +0300 Subject: [PATCH 15/17] fix: keep a bare typing.Generic annotation unresolvable at declaration (#605) The plain-class fast path in `SignatureItem.from_type` took any annotation whose type is exactly `type`. `typing.Generic` is one, but `typing.get_origin(Generic)` is `Generic`, so on main it is an unresolvable generic and `Factory(...)` rejects a parameter annotated with it. The fast path took it for a plain class and moved the failure to resolve time. It now excludes `Generic`, and the `NoneType` branch runs first so the fast path no longer needs its own `NoneType` check. A/B against the parent commit, 12 interleaved rounds: `Factory(...)` +0.5% and +0.8%, within noise. --- modern_di/types_parser.py | 4 ++-- tests/test_types_parser.py | 11 +++++++++++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 7084d920..a37f552c 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -23,12 +23,12 @@ class SignatureItem: @classmethod def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem": - if type(type_) is type and type_ is not types.NoneType: - return cls(arg_type=type_, default=default) if type_ is types.NoneType: # The degenerate nullable: the union branch below would take it for a plain type and # try to resolve `NoneType` from the registry. return cls(default=default, is_nullable=True) + if type(type_) is type and type_ is not typing.Generic: # `get_origin(Generic)` is `Generic` + return cls(arg_type=type_, default=default) origin = typing.get_origin(type_) if origin is typing.Annotated: diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index b88f5a37..5359b86e 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -31,6 +31,7 @@ class GenericClass(typing.Generic[types.T]): ... (list[str] | None, SignatureItem(arg_type=list, is_nullable=True)), (GenericClass[str], SignatureItem(unresolvable_generic=GenericClass[str])), (GenericClass[str] | None, SignatureItem(arg_type=GenericClass, is_nullable=True)), + (typing.Generic, SignatureItem(unresolvable_generic=typing.Generic)), # `None` is the degenerate nullable: a union with zero non-None members. (type(None), SignatureItem(is_nullable=True)), ], @@ -279,6 +280,16 @@ def test_parameterized_generic_param_without_default_raises_at_declaration() -> assert "skip_creator_parsing" not in str(exc_info.value) +def _bare_generic_param_creator(x: typing.Generic) -> str: # ty: ignore[invalid-type-form] + return str(x) + + +def test_bare_generic_param_without_default_raises_at_declaration() -> None: + assert _bare_generic_param_creator(1) == "1" + with pytest.raises(exceptions.UnsupportedCreatorParameterError, match=r"typing\.Generic"): + providers.Factory(creator=_bare_generic_param_creator) + + def test_parameterized_generic_param_supplied_via_kwargs_is_allowed() -> None: sentinel = [_GenericDep()] provider = providers.Factory(creator=_generic_param_creator, kwargs={"x": sentinel}) From ce52a166502da86134ad917d7672a312ea855e1a Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 20:10:20 +0300 Subject: [PATCH 16/17] refactor: move the cache module out of registries and name it cache (#605) `registries/cache_registry.py` stopped holding a registry when the cache items moved onto the container. It is now `modern_di/cache.py`: `CacheItem` plus the functions over a container's own items and creation order. A new `cached_count(items)` there replaces the inline count in `Container.__repr__`. Its tests move to `tests/test_cache.py`. --- AGENTS.md | 7 ++++--- .../{registries/cache_registry.py => cache.py} | 5 +++++ modern_di/container.py | 14 ++++++-------- modern_di/resolver_compiler.py | 2 +- tests/helpers.py | 2 +- tests/providers/test_context_provider.py | 2 +- .../test_cache_registry.py => test_cache.py} | 2 +- 7 files changed, 19 insertions(+), 15 deletions(-) rename modern_di/{registries/cache_registry.py => cache.py} (96%) rename tests/{registries/test_cache_registry.py => test_cache.py} (98%) diff --git a/AGENTS.md b/AGENTS.md index 9d4b7e13..5b1ddf1f 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -46,9 +46,10 @@ Every module under `modern_di/` is named for what it does; read it. What a singl `docs/troubleshooting/`, enforced by `tests/test_docs_slug_census.py`). Add a message, a glyph, or a 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 and context are per-container and live in the - container's own slots: `cache_registry` holds `CacheItem` and the module functions that fetch an - item and close a container's items, and the context is a plain dict. + `overrides_registry` are shared tree-wide. +- `cache.py` — `CacheItem` and the functions over a container's own cache items and creation order + (fetch, count, close). Cache and context are per-container and live in the container's slots; the + context is a plain dict. - `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. diff --git a/modern_di/registries/cache_registry.py b/modern_di/cache.py similarity index 96% rename from modern_di/registries/cache_registry.py rename to modern_di/cache.py index b4e1d9d9..934cc500 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/cache.py @@ -77,6 +77,11 @@ def fetch_cache_item(items: dict[int, CacheItem], provider: Factory[typing.Any]) return items.setdefault(provider_id, CacheItem(settings=settings)) +def cached_count(items: dict[int, CacheItem]) -> int: + """Return how many of ``items`` hold a value.""" + return sum(1 for item in items.values() if item.cache is not types.UNSET) + + async def close_async(creation_order: list[CacheItem]) -> None: """Close every item newest first and empty ``creation_order``; failures raise together at the end.""" finalizer_errors: list[Exception] = [] diff --git a/modern_di/container.py b/modern_di/container.py index 2065b123..b285ed10 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -2,13 +2,11 @@ import enum import typing -from modern_di import dependency_graph, exceptions, types +from modern_di import cache, dependency_graph, exceptions, types from modern_di._scope_algebra import next_deeper from modern_di.group import Group from modern_di.providers.abstract import AbstractProvider from modern_di.providers.container_provider import container_provider -from modern_di.registries import cache_registry -from modern_di.registries.cache_registry import CacheItem from modern_di.registries.overrides_registry import OverrideHandle from modern_di.registries.providers_registry import ProvidersRegistry from modern_di.scope import Scope @@ -131,8 +129,8 @@ def _set_state( self._scope = scope self._parent_container = parent self._scope_map = scope_map - self._cache_items: dict[int, CacheItem] = {} - self._creation_order: list[CacheItem] = [] + self._cache_items: dict[int, cache.CacheItem] = {} + self._creation_order: list[cache.CacheItem] = [] if context is None: self._context = {} elif type(context) is dict: @@ -251,7 +249,7 @@ async def close_async(self) -> None: """ self._closed = True if self._creation_order: - await cache_registry.close_async(self._creation_order) + await cache.close_async(self._creation_order) def close_sync(self) -> None: """Mark this container closed, then run its sync finalizers, newest first. @@ -263,7 +261,7 @@ def close_sync(self) -> None: """ self._closed = True if self._creation_order: - cache_registry.close_sync(self._creation_order) + cache.close_sync(self._creation_order) def override(self, provider: AbstractProvider[types.T], override_object: types.T) -> OverrideHandle[types.T]: """Apply an override immediately, tree-wide. @@ -299,7 +297,7 @@ def set_context(self, context_type: type[types.T], obj: types.T) -> None: def __repr__(self) -> str: n_providers = len(self._providers_registry) - n_cached = sum(1 for item in self._cache_items.values() if item.cache is not types.UNSET) + n_cached = cache.cached_count(self._cache_items) parent = self.parent_container.scope.name if self.parent_container else None return f"Container(scope={self.scope.name}, parent={parent}, providers={n_providers}, cached={n_cached})" diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 3c711783..6c3c92a4 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -20,13 +20,13 @@ import typing from modern_di import exceptions, types +from modern_di.cache import fetch_cache_item from modern_di.dependency_graph import redirect_hops, terminal_chain from modern_di.providers.abstract import AbstractProvider from modern_di.providers.alias import Alias from modern_di.providers.container_provider import container_provider from modern_di.providers.context_provider import ContextProvider from modern_di.providers.factory import Factory -from modern_di.registries.cache_registry import fetch_cache_item if typing.TYPE_CHECKING: diff --git a/tests/helpers.py b/tests/helpers.py index e8265af4..3de5a138 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -1,8 +1,8 @@ import typing from modern_di import Container +from modern_di.cache import CacheItem, fetch_cache_item from modern_di.providers import Factory -from modern_di.registries.cache_registry import CacheItem, fetch_cache_item def cache_item(container: Container, provider: Factory[typing.Any]) -> CacheItem: diff --git a/tests/providers/test_context_provider.py b/tests/providers/test_context_provider.py index f8d85015..9af30e7f 100644 --- a/tests/providers/test_context_provider.py +++ b/tests/providers/test_context_provider.py @@ -210,7 +210,7 @@ def test_context_provider_reads_registry_at_its_own_scope_not_resolving_containe # --- set_context cross-scope staleness (2026-06-14 deep audit) --- # # An APP-scoped ContextProvider consumed by a deeper (REQUEST) scoped Factory: -# the factory's compiled kwargs live in the request child's cache_registry, so a +# the factory's compiled kwargs live in the request child's cache, so a # late app.set_context must still be picked up by subsequent resolves from that # child. Context values are resolved live, not baked in at first resolve. diff --git a/tests/registries/test_cache_registry.py b/tests/test_cache.py similarity index 98% rename from tests/registries/test_cache_registry.py rename to tests/test_cache.py index bfe174c4..2eb0e727 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/test_cache.py @@ -2,8 +2,8 @@ import typing from concurrent.futures import ThreadPoolExecutor +from modern_di.cache import CacheItem, close_async, fetch_cache_item from modern_di.providers import CacheSettings, Factory -from modern_di.registries.cache_registry import CacheItem, close_async, fetch_cache_item from modern_di.types import UNSET From 4668331da3c0f1cf6272e5a28e956ee384ad8fbd Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 6 Oct 2026 20:10:51 +0300 Subject: [PATCH 17/17] refactor: keep the resolver globals beside their reader and tidy three spots (#605) - `_navigate` and `_FACTORY_GLOBALS` sit just above `_compile_factory`, the one function that reads them, instead of at the end of the module. - `Factory._positional_names` is set with a plain `if` instead of a multi-line conditional expression. - `collect_errors` calls `walk(roots=registry, registry=registry)`, so it is clear the registry is both the root set and the lookup. - The `resolver_compiler` docstring lists `_closed` among the container slots the generated code reads. No behaviour change. --- modern_di/dependency_graph.py | 2 +- modern_di/providers/factory.py | 9 ++--- modern_di/resolver_compiler.py | 61 +++++++++++++++++----------------- 3 files changed, 35 insertions(+), 37 deletions(-) diff --git a/modern_di/dependency_graph.py b/modern_di/dependency_graph.py index 955c1bca..ddf5ad00 100644 --- a/modern_di/dependency_graph.py +++ b/modern_di/dependency_graph.py @@ -180,7 +180,7 @@ def _enter( def collect_errors(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 walk(registry, registry): + for event in walk(roots=registry, registry=registry): if type(event) is Edge: parent, name, dep = event dependency_chain = terminal_chain(dep, registry) diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index 4b06717c..6e340a50 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -65,12 +65,9 @@ def __init__( # noqa: PLR0913 ) self._params = parsed.params names = tuple(parsed.params) - self._positional_names: tuple[str, ...] | None = ( - None - if (names and parsed.has_positional_only_gap) - or any(item.is_keyword_only for item in parsed.params.values()) - else names - ) + self._positional_names: tuple[str, ...] | None = names + if (names and parsed.has_positional_only_gap) or any(item.is_keyword_only for item in parsed.params.values()): + self._positional_names = None super().__init__(scope=scope, bound_type=bound_type, inferred_bound_type=parsed.return_type.arg_type) self._creator = creator self._cache_settings: CacheSettings[typing.Any] | None = CacheSettings._coerce(cache) # noqa: SLF001 diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index 6c3c92a4..cf0f0d2e 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -9,8 +9,9 @@ ``ProvidersRegistry.drop_resolvers``). Why a template and not shared helpers: see docs/adr/0001-resolver-hot-path-generated-source.md. -The template reaches into the `Container` slots `_scope`, `_scope_map`, `_cache_items` and `_creation_order` to stay -within that frame budget. No linter sees the template, so those reaches are outside every suppression here. +The template reaches into the `Container` slots `_scope`, `_scope_map`, `_closed`, `_cache_items` and +`_creation_order` to stay within that frame budget. No linter sees the template, so those reaches are +outside every suppression here. """ import enum @@ -159,6 +160,34 @@ def _code(arity: int, names: tuple[str, ...] | None, static: bool, cached: bool) return compile(source, filename, "exec"), arg_lines +def _navigate( + container: "Container", + scope: enum.IntEnum, + resolution_step: "typing.Callable[[], exceptions.ResolutionStep]", +) -> "Container": + """Miss path for a scope absent from `_scope_map`, or held there by another enum's same-valued member. + + The scope error carries this provider's resolution step. + """ + try: + return container.find_container(scope) + except _SCOPE_ERRORS as exc: + exc._prepend_step(resolution_step()) + raise + + +_FACTORY_GLOBALS: dict[str, typing.Any] = { + "UNSET": types.UNSET, + "fetch_cache_item": fetch_cache_item, + "_navigate": _navigate, + "ResolutionError": exceptions.ResolutionError, + "CreatorCallError": exceptions.CreatorCallError, + "ContainerClosedError": exceptions.ContainerClosedError, + "ContextValueNotSetError": exceptions.ContextValueNotSetError, + "redirect_hops": redirect_hops, +} + + def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "Resolver": plan = registry.plan_for(f) if plan.unwireable: @@ -283,31 +312,3 @@ def resolve(container: "Container") -> typing.Any: raise exceptions.ContextValueNotSetError(context_type=context_type, provider_scope=scope) return resolve - - -def _navigate( - container: "Container", - scope: enum.IntEnum, - resolution_step: "typing.Callable[[], exceptions.ResolutionStep]", -) -> "Container": - """Miss path for a scope absent from `_scope_map`, or held there by another enum's same-valued member. - - The scope error carries this provider's resolution step. - """ - try: - return container.find_container(scope) - except _SCOPE_ERRORS as exc: - exc._prepend_step(resolution_step()) - raise - - -_FACTORY_GLOBALS: dict[str, typing.Any] = { - "UNSET": types.UNSET, - "fetch_cache_item": fetch_cache_item, - "_navigate": _navigate, - "ResolutionError": exceptions.ResolutionError, - "CreatorCallError": exceptions.CreatorCallError, - "ContainerClosedError": exceptions.ContainerClosedError, - "ContextValueNotSetError": exceptions.ContextValueNotSetError, - "redirect_hops": redirect_hops, -}