diff --git a/AGENTS.md b/AGENTS.md index 6f86afa5..5b1ddf1f 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -46,7 +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_registry` and `context_registry` are per-container. + `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/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 diff --git a/modern_di/cache.py b/modern_di/cache.py new file mode 100644 index 00000000..934cc500 --- /dev/null +++ b/modern_di/cache.py @@ -0,0 +1,120 @@ +import dataclasses +import inspect +import threading +import typing + +from modern_di import exceptions, types +from modern_di.providers import CacheSettings, Factory + + +_T = typing.TypeVar("_T") +_R = typing.TypeVar("_R") +_V = typing.TypeVar("_V") + + +@dataclasses.dataclass(kw_only=True, slots=True) +class CacheItem: + settings: CacheSettings[typing.Any] + cache: typing.Any = types.UNSET + finalized: bool = False + lock: threading.RLock = dataclasses.field(default_factory=threading.RLock, repr=False, compare=False) + + def clear(self) -> None: + if self.settings.clear_cache: + self.cache = types.UNSET + self.finalized = False + + def get_or_create( + self, + build: typing.Callable[[_T], _R], + target: _T, + create: typing.Callable[[_R], _V], + ) -> tuple[_V, bool]: + """Return the memoized singleton, or ``create(build(target))`` it once under this item's lock. + + 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: + return self.cache, False + with self.lock: + if self.cache is not types.UNSET: + return self.cache, False + value = create(build(target)) + self.cache = value + return value, True + + def _pending_finalizer(self) -> typing.Callable[[typing.Any], typing.Awaitable[None] | None] | None: + """Return the finalizer still owed to the cached value, or None when nothing is owed.""" + return None if self.cache is types.UNSET or self.finalized else self.settings.finalizer + + def close_sync(self) -> None: + if (finalizer := self._pending_finalizer()) is not None: + if self.settings._is_async_finalizer: # noqa: SLF001 + raise exceptions.AsyncFinalizerInSyncCloseError(instance_type=type(self.cache)) + try: + result = finalizer(self.cache) + except Exception: + self.clear() + raise + 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)) + self.finalized = True + + self.clear() + + +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)) + + +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] = [] + for cache_item in reversed(creation_order): + 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) + + +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/container.py b/modern_di/container.py index 94519093..b285ed10 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -2,29 +2,26 @@ 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.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: @@ -36,9 +33,10 @@ class Container: """ __slots__ = ( - "_cache_registry", + "_cache_items", "_closed", - "_context_registry", + "_context", + "_creation_order", "_parent_container", "_providers_registry", "_scope", @@ -62,6 +60,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,10 +104,16 @@ 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. - 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( @@ -118,17 +124,19 @@ 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 self._scope_map = scope_map - self._cache_registry = CacheRegistry() - self._context_registry = ContextRegistry(copy.copy(context) if context is not None else {}) + self._cache_items: dict[int, cache.CacheItem] = {} + self._creation_order: list[cache.CacheItem] = [] + 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: @@ -163,11 +171,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 +199,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 +216,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() @@ -240,8 +248,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.close_async(self._creation_order) def close_sync(self) -> None: """Mark this container closed, then run its sync finalizers, newest first. @@ -252,8 +260,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.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. @@ -285,11 +293,11 @@ 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) - n_cached = self._cache_registry.cached_count() + 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/dependency_graph.py b/modern_di/dependency_graph.py index 9d7853bc..ddf5ad00 100644 --- a/modern_di/dependency_graph.py +++ b/modern_di/dependency_graph.py @@ -14,16 +14,17 @@ if typing.TYPE_CHECKING: - from modern_di import Container 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``.""" @@ -32,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.""" @@ -49,7 +52,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 +61,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 +70,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 +100,28 @@ 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) + if root.provider_id not in 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,17 +129,14 @@ def find_cycle_from( def _walk_from( start: "AbstractProvider[typing.Any]", - container: "Container", + registry: "ProvidersRegistry", 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, container, visiting, path, stack) + yield from _enter(start, registry, visiting, path, stack) while stack: try: @@ -154,12 +155,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,44 +170,42 @@ 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): - match event: - case NodeEntered(provider): - errors.extend(provider._iter_validation_issues(container)) # noqa: SLF001 - case DependenciesError(_, error): - errors.append(error) - case Edge(parent, name, dep): - dependency_chain = terminal_chain(dep, container) - dependency_scope = dependency_chain[-1].scope - parent_scope = effective_scope(parent, container) - if dependency_scope > parent_scope: - errors.append( - exceptions.InvalidScopeDependencyError( - provider=parent, - parameter_name=name, - dependency_chain=dependency_chain, - ) + for event in walk(roots=registry, registry=registry): + 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, container)) - 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/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..6e340a50 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 @@ -46,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 @@ -65,11 +64,11 @@ 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 - super().__init__( - scope=scope, - bound_type=parsed.return_type.arg_type if isinstance(bound_type, types.UnsetType) else bound_type, - ) + names = tuple(parsed.params) + 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 self._kwargs = kwargs @@ -131,13 +130,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,32 +199,22 @@ 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`. 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, 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/cache_registry.py b/modern_di/registries/cache_registry.py deleted file mode 100644 index 926d2c20..00000000 --- a/modern_di/registries/cache_registry.py +++ /dev/null @@ -1,132 +0,0 @@ -import dataclasses -import inspect -import threading -import typing - -from modern_di import exceptions, types -from modern_di.providers import CacheSettings, Factory - - -_R = typing.TypeVar("_R") -_V = typing.TypeVar("_V") - - -@dataclasses.dataclass(kw_only=True, slots=True) -class CacheItem: - settings: CacheSettings[typing.Any] - cache: typing.Any = types.UNSET - finalized: bool = False - lock: threading.RLock = dataclasses.field(default_factory=threading.RLock, repr=False, compare=False) - - def clear(self) -> None: - if self.settings.clear_cache: - self.cache = types.UNSET - self.finalized = False - - def get_or_create( - self, - resolve: typing.Callable[[], _R], - create: typing.Callable[[_R], _V], - ) -> tuple[_V, bool]: - """Return the memoized singleton, or resolve-and-create it once under this item's lock. - - A hit never takes the lock. A miss resolves 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: - return self.cache, False - with self.lock: - if self.cache is not types.UNSET: - return self.cache, False - value = create(resolve()) - self.cache = value - return value, True - - def _pending_finalizer(self) -> typing.Callable[[typing.Any], typing.Awaitable[None] | None] | None: - """Return the finalizer still owed to the cached value, or None when nothing is owed.""" - return None if self.cache is types.UNSET or self.finalized else self.settings.finalizer - - async def close_async(self) -> None: - if (finalizer := self._pending_finalizer()) is not None: - try: - result = finalizer(self.cache) - if 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 - raise exceptions.AsyncFinalizerInSyncCloseError(instance_type=type(self.cache)) - try: - result = finalizer(self.cache) - except Exception: - self.clear() - raise - if inspect.isawaitable(result): - if inspect.iscoroutine(result): - result.close() # suppress "never awaited" warning - raise exceptions.AsyncFinalizerInSyncCloseError(instance_type=type(self.cache)) - self.finalized = True - - 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) 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..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 `Container._scope`, `Container._scope_map` and `CacheRegistry._items` 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 @@ -20,7 +21,8 @@ import typing from modern_di import exceptions, types -from modern_di.dependency_graph import redirect_hops +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 @@ -38,7 +40,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 +77,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 +93,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 """ @@ -107,16 +108,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) + value, created = cache_item.get_or_create(build, target, create) if created: - cache_registry.mark_created(cache_item) + target._creation_order.append(cache_item) return value """ ) @@ -159,8 +160,36 @@ 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 = f._wiring_plan(registry) + plan = registry.plan_for(f) if plan.unwireable: return _compile_unwireable_factory(f, plan) positional = f._can_call_positionally(plan) @@ -170,28 +199,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, - "partial": functools.partial, - "_navigate": _navigate, - "_STEP_ERRORS": STEP_ERRORS, - "CreatorCallError": exceptions.CreatorCallError, - "ContainerClosedError": exceptions.ContainerClosedError, - "ContextValueNotSetError": exceptions.ContextValueNotSetError, - "redirect_hops": redirect_hops, - **{ - 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}]" @@ -211,21 +230,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 +303,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] @@ -299,19 +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 diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 18ee09d8..a37f552c 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 @@ -27,6 +27,8 @@ def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem": # 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: @@ -49,7 +51,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_ @@ -112,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)), @@ -135,7 +157,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 +170,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/pyproject.toml b/pyproject.toml index b760ba1c..c28aa1df 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. @@ -107,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 diff --git a/tests/helpers.py b/tests/helpers.py index 81a58c4b..3de5a138 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -1,10 +1,10 @@ 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 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_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_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_context_provider.py b/tests/providers/test_context_provider.py index 17655778..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. @@ -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/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_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/registries/test_cache_registry.py b/tests/test_cache.py similarity index 64% rename from tests/registries/test_cache_registry.py rename to tests/test_cache.py index 99eca76b..2eb0e727 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/test_cache.py @@ -2,10 +2,8 @@ import typing from concurrent.futures import ThreadPoolExecutor -import pytest - +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, CacheRegistry from modern_di.types import UNSET @@ -29,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} @@ -42,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" @@ -85,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" @@ -102,27 +100,27 @@ 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) 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)]] @@ -130,27 +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] = [] - 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) + never_built = CacheItem(settings=CacheSettings(finalizer=finalized.append)) + creation_order = [plain, persistent, with_finalizer, never_built] - 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 with_finalizer.cache is UNSET + assert creation_order == [] diff --git a/tests/test_container.py b/tests/test_container.py index ec56998a..ad0f6439 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. @@ -579,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: @@ -613,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_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_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) diff --git a/tests/test_resolver_compiler.py b/tests/test_resolver_compiler.py index 7912b6a4..e2b2beed 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. @@ -889,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. """ diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index b66297de..5359b86e 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]): ... @@ -21,15 +22,16 @@ 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)), + (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)), ], @@ -147,7 +149,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=()), }, ), ), @@ -278,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}) @@ -455,3 +467,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