diff --git a/benchmarks/README.md b/benchmarks/README.md index 03bd496c..7b73ae85 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -17,6 +17,7 @@ cost. Runs in CI (informational, non-gating) and locally via `just bench`. | G5 | Cross-scope resolve, REQUEST -> APP dep | `find_container` traversal | | G6 | `build_child_container(REQUEST)` | per-request setup | | G7 | Full lifecycle batch: K=100 x (build REQUEST -> sync-init cached resolve -> `await close_async()`) | real per-request cost incl. async teardown | +| G7b | One request cycle: build REQUEST -> first-resolve one cached REQUEST provider -> `close_sync()` | per-request cost of a cached item, no event loop | | G7c | Control: K=100 empty awaits in one loop entry | residual event-loop floor inside G7 | | G8 | Cold first-resolve: build root container + compile + resolve, depth 6 | construction + first-compile cost | | G8b | G8 with every provider `cache=True` | the cached template's cold-miss `build`/`create`, read against G8 | @@ -27,7 +28,7 @@ cost. Runs in CI (informational, non-gating) and locally via `just bench`. | G13 | Per-request cycle finalizing 10 cached resources (`close_sync`) | LIFO teardown at scale | | G13b | Batch of K=100 request cycles, 10 finalizer-less cached REQUEST providers, `await close_async()` | the async close loop when there is nothing to finalize | | G14 | Concurrent cached-hit throughput, N threads (lock-free read) | free-threaded read scaling | -| G15 | Concurrent first-resolve, N threads (double-checked creation lock) | free-threaded creation-lock contention | +| G15 | Concurrent first-resolve, N threads (per-item double-checked creation lock) | free-threaded creation-lock contention | | G16 | Warm by-type `resolve(SomeType)`, small graph | `find_provider` lookup on the integration/`@inject` path | | G17 | Warm by-type `resolve(SomeType)`, 200-provider registry | lookup cost at realistic registry scale | | G18 | Warm resolve through an `Alias` to a cached source | the alias hop, read against G2 | diff --git a/benchmarks/test_guard_concurrency.py b/benchmarks/test_guard_concurrency.py index 950d3f19..c616f0e2 100644 --- a/benchmarks/test_guard_concurrency.py +++ b/benchmarks/test_guard_concurrency.py @@ -10,8 +10,8 @@ across N threads. The cached-hit path is lock-free, so on a free-threaded build (PEP 703) the batch time should *drop* as N rises (throughput scales); under the GIL it stays flat. - G15 concurrent first-resolve: N threads each race to resolve the *same* K cold singletons, so - they contend on the double-checked creation lock (`CacheItem.get_or_create`). Singleton - creation is serialized by design, so this is expected *not* to scale even free-threaded — the + they contend on each item's double-checked creation lock (`CacheItem.get_or_create`). Creating + one singleton is serialized by design, so this is expected *not* to scale even free-threaded: the measured cost is the contention itself (the known trade-off vs lock-free-slot rivals). Read the batch-time-vs-thread-count trend, not the absolutes. The GIL vs free-threaded comparison @@ -124,10 +124,12 @@ def _worker() -> None: @pytest.mark.parametrize("n_threads", _THREAD_COUNTS) def test_g15b_concurrent_first_resolve_sibling_children(benchmark, n_threads): - # Each thread builds its OWN REQUEST child and first-resolves K cached providers in it, so - # every creation is a cold miss in a container no other thread touches. With a lock per - # container these creations never contend; with one lock per tree they serialize. This is the - # scenario G15 does not cover -- G15 races on one root, whose lock is shared either way. + """Each thread builds its own REQUEST child and first-resolves K cached providers in it. + + Every creation is a cold miss in a container no other thread touches. Each cache item has its + own lock, so these creations never contend; under a lock shared by the tree they would + serialize. G15 does not cover this: it races on one root's items, whose locks are shared. + """ check = Container(scope=Scope.APP, groups=[_REQUEST_GROUP]) check.open() probe = check.build_child_container(scope=Scope.REQUEST) diff --git a/benchmarks/test_guard_lifecycle.py b/benchmarks/test_guard_lifecycle.py index 457dbf16..3be9ba03 100644 --- a/benchmarks/test_guard_lifecycle.py +++ b/benchmarks/test_guard_lifecycle.py @@ -7,13 +7,14 @@ K times per loop entry to isolate DI work from the ~27us event-loop entry cost. Divide G7's number by 100 for per-request cost. G7c is the control (K=100 empty awaits on the same shape) so the residual loop overhead is visible (~15% at K=100). +G7b is one request cycle with a synchronous close, so no event loop is in the number. See benchmarks/README.md. """ import asyncio import dataclasses -from benchmarks._pinned import ITER_UNDER_1US, ROUNDS +from benchmarks._pinned import ITER_UNDER_1US, ITER_UNDER_2US, ROUNDS from modern_di import Container, Group, Scope, providers @@ -117,6 +118,34 @@ def _run_batch() -> None: loop.close() +# --- G7b: one request cycle, cached REQUEST provider without a finalizer, sync close --- +@dataclasses.dataclass(slots=True) +class RequestService: + pass + + +class RequestCycleGroup(Group): + svc = providers.Factory(creator=RequestService, scope=Scope.REQUEST, cache=True) + + +def test_g7b_request_cycle_sync(benchmark): + """Build a REQUEST child, first-resolve one cached REQUEST provider, close it synchronously. + + Every cycle creates a fresh cache item and its lock, so this is where a per-item cost lands. + """ + app = Container(scope=Scope.APP, groups=[RequestCycleGroup]) + app.open() + + def _one_request() -> RequestService: + req = app.build_child_container(scope=Scope.REQUEST) + svc = req.resolve_provider(RequestCycleGroup.svc) + req.close_sync() + return svc + + result = benchmark.pedantic(_one_request, rounds=ROUNDS, iterations=ITER_UNDER_2US) + assert isinstance(result, RequestService) + + # --- G13: teardown at scale -- 10 cached REQUEST resources, sync finalizers --- # G7 finalizes one resource; a real request closes several. G13 measures the per-request cycle # with 10 cached REQUEST providers (each a sync finalizer) so the LIFO close loop is exercised. diff --git a/docs/introduction/design-decisions.md b/docs/introduction/design-decisions.md index e16f2544..36f4178b 100644 --- a/docs/introduction/design-decisions.md +++ b/docs/introduction/design-decisions.md @@ -10,11 +10,12 @@ Async resolution will not be added. ## 2. Cached factories are thread-safe -Cached `Factory` providers use one reentrant lock (`threading.RLock`) per container tree, created by the root and shared by every child, so concurrent resolves in multiple threads still produce exactly one instance per cache. The lock is taken only on a cache miss; a resolve that finds the instance already cached never touches it. +Each cache item, the cached instance of one `Factory` in one container, has its own reentrant lock (`threading.RLock`), so concurrent resolves in multiple threads still produce exactly one instance per cache. The lock is taken only on a cache miss; a resolve that finds the instance already cached never touches it. On a miss, the dependencies are resolved and the creator is called while the lock is held, so threads that miss together build the dependencies once, and the others wait for that instance. ### The thread-safety boundary -- Cached / singleton creation is locked. The tree-wide reentrant lock guards the create-and-store step, so two threads racing to resolve the same cached provider get the same single instance. +- Cached / singleton creation is locked per cached provider. Two threads racing to resolve the same cached provider get the same single instance, and transient dependencies of that provider are built once for it. Creations of different cached providers do not wait for each other, so a creator can hand a resolve of another cached type to a worker thread and wait for the result. A creator that waits on another thread resolving the provider it is creating, directly or through its dependencies, still deadlocks: that is a cycle. +- Call `validate()` at startup if the graph might contain a cycle. Without it, a single thread resolving a cyclic graph gets `CircularDependencyError` from the runtime guard. Two threads that cold-resolve different providers of the same cycle at the same time can each hold one cache item's lock while waiting for the other's, and block forever. - Provider registration is safe. `ProvidersRegistry` mutations (`register`, `add_providers`) are guarded by the registry's own lock, and iteration snapshots the provider dict (`iter(list(...))`), so registering providers concurrently, or while another thread iterates, will not corrupt the registry or raise "dict changed size during iteration". - Registration belongs to the setup phase. The registry is lock-guarded against corruption, but the supported model is to register every provider diff --git a/docs/introduction/performance.md b/docs/introduction/performance.md index ce0d30c7..413bbe42 100644 --- a/docs/introduction/performance.md +++ b/docs/introduction/performance.md @@ -322,6 +322,12 @@ resolver, so a hop through an alias runs no frame of its own (−33% on an alias now matches a plain cached resolve). An error that crosses an alias still shows the alias in its chain: the parent puts the hop back when it adds its own step. +4.0 also replaced the tree lock from #541 with one lock per cache item (#569), so creating one +cached factory no longer waits for another. A child build still allocates no lock, because the +lock comes with the cache item. The cost moved to the first resolve of a cached factory in each +container. A request that resolves one 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. + ## Reproduce it yourself ```bash diff --git a/docs/migration/to-4.x.md b/docs/migration/to-4.x.md index 32f0d98a..427547de 100644 --- a/docs/migration/to-4.x.md +++ b/docs/migration/to-4.x.md @@ -11,10 +11,29 @@ Every `Container` argument after `scope` is keyword-only: `Container(Scope.APP, ### `use_lock` is removed -`Container(use_lock=...)` raises `TypeError`; drop the argument. Every container tree is now +`Container(use_lock=...)` raises `TypeError`; drop the argument. Cached factories are always locked, and the lock is taken only on a cache miss (see [Design decisions](../introduction/design-decisions.md#2-cached-factories-are-thread-safe)). +### Each cached factory has its own lock + +In 3.x, with the default `use_lock=True`, every cached creator ran under a shared lock: the +container's lock, or from 3.6 the lock of the whole container tree. Two cached factories that +shared the lock were never created at the same time. In 4.0 each cached factory has its own lock +in each container, so creators of different cached factories can run at the same time on +different threads. If two creators share state that is not thread-safe, guard that state with +your own lock. + +On a cache miss the lock is now held while the dependencies are resolved as well. Threads that +miss together build the dependencies once, where 3.x built them once per thread and discarded all +but one. A cached creator can also wait on another thread that resolves a different cached type, +which deadlocked in 3.x. + +One case that raised in 3.x can now block. If the graph has a cycle and `validate()` was never +called, two threads that cold-resolve different providers of that cycle at the same time can each +hold one lock and wait for the other forever. In 3.x both got `CircularDependencyError`. A single +thread still gets that error. Call `validate()` at startup to catch cycles before serving. + ### Resolving on a closed container raises In 3.x, resolving from a closed container, or through a child whose resolve reached a closed diff --git a/docs/providers/advanced-api.md b/docs/providers/advanced-api.md index 647d6733..41af94c2 100644 --- a/docs/providers/advanced-api.md +++ b/docs/providers/advanced-api.md @@ -38,5 +38,7 @@ inspect or iterate all providers declared on a group hierarchy. `Container` subclass does not redirect navigation. `resolve` and `resolve_provider` are entry points, not hooks either: a compiled resolver calls its dependencies' resolvers directly, so an override of either sees only the top-level call. -- The root creates one `threading.RLock`, and every child shares it. A cached `Factory` holds - that lock while it builds on a cold cache miss, so one instance is created per cache key. +- Each cached `Factory` gets its own `threading.RLock` in each container that caches it, created + with the cache item on the first resolve there. Building a child allocates no lock. On a cold + cache miss the factory holds its lock while it resolves its dependencies and calls the creator, + so one instance is created per cache key. A warm resolve does not take the lock. diff --git a/docs/providers/errors-and-exceptions.md b/docs/providers/errors-and-exceptions.md index ca237fbe..35ef6cb8 100644 --- a/docs/providers/errors-and-exceptions.md +++ b/docs/providers/errors-and-exceptions.md @@ -121,7 +121,8 @@ to render the chain programmatically. [Troubleshooting: ArgumentResolutionError](../troubleshooting/argument-resolution-error.md). - `CircularDependencyError` is raised when the provider graph contains a cycle (A → B → A); the message shows the cycle path. Raised eagerly by `validate()`, and also by a bare `resolve()` on an - unvalidated cyclic graph via a runtime guard; see + unvalidated cyclic graph via a runtime guard. The guard covers one thread resolving: concurrent + first resolves of an unvalidated cyclic graph can block instead, so call `validate()` at startup. See [Troubleshooting: Circular dependency](../troubleshooting/circular-dependency.md#the-runtime-cycle-guard-without-validate). - `CreatorCallError` is raised when a creator's dependencies all resolved but argument binding failed while calling it (the assembled arguments don't match the signature, typically a `kwargs` / diff --git a/docs/providers/factories.md b/docs/providers/factories.md index 798fdc4c..91fe6726 100644 --- a/docs/providers/factories.md +++ b/docs/providers/factories.md @@ -51,7 +51,7 @@ This is modern-di's Singleton. There is no separate `Singleton` provider class: [Where is Singleton?](../introduction/comparison.md#where-is-singleton-cross-framework-vocabulary) for the full cross-framework mapping. -The caching mechanism is thread-safe: when multiple threads resolve the same cached factory simultaneously, only one instance is created. +The caching mechanism is thread-safe: when multiple threads resolve the same cached factory simultaneously, only one instance is created, and its dependencies are resolved once for it. Other threads wait for that instance; resolves of other cached factories do not. ```python import random diff --git a/docs/troubleshooting/circular-dependency.md b/docs/troubleshooting/circular-dependency.md index 78deaa69..1b621951 100644 --- a/docs/troubleshooting/circular-dependency.md +++ b/docs/troubleshooting/circular-dependency.md @@ -31,6 +31,11 @@ re-walks the static graph from the failing provider, and, since a cycle is reach `RecursionError`. A creator that merely recurses on its own, with no actual cycle in the provider graph, still raises the original `RecursionError` unchanged. This guard runs on every resolve, whether or not `validate()` was ever called. +The guard covers resolution on one thread. Each cached factory locks its cache item while it is +created, so two threads that cold-resolve different providers of the same cycle at the same time +can each hold one of those locks and wait for the other forever, and neither reaches the guard. +If the graph might have a cycle, call `validate()` at startup, before any thread resolves. + ### Cycle detection with `validate()` Calling `validate()` up front finds the *same* cycle earlier, and finds *every* issue in the graph diff --git a/modern_di/container.py b/modern_di/container.py index 0f0e9909..717b77df 100644 --- a/modern_di/container.py +++ b/modern_di/container.py @@ -1,6 +1,5 @@ import copy import enum -import threading import typing from modern_di import dependency_graph, exceptions, types @@ -40,7 +39,6 @@ class Container: "_cache_registry", "_closed", "_context_registry", - "_lock", "_providers_registry", "_scope_map", "parent_container", @@ -87,10 +85,8 @@ def __init__( self._context_registry = ContextRegistry(copy.copy(context) if context is not None else {}) self._providers_registry: ProvidersRegistry if parent_container: - self._lock = parent_container._lock # noqa: SLF001 self._providers_registry = parent_container._providers_registry # noqa: SLF001 else: - self._lock = threading.RLock() self._providers_registry = ProvidersRegistry() self._providers_registry.register(Container, container_provider) if groups: diff --git a/modern_di/registries/cache_registry.py b/modern_di/registries/cache_registry.py index 9ba91a95..135061d5 100644 --- a/modern_di/registries/cache_registry.py +++ b/modern_di/registries/cache_registry.py @@ -1,15 +1,12 @@ import dataclasses import inspect +import threading import typing from modern_di import exceptions, types from modern_di.providers import CacheSettings, Factory -if typing.TYPE_CHECKING: - import threading - - _R = typing.TypeVar("_R") _V = typing.TypeVar("_V") @@ -19,6 +16,7 @@ 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: @@ -27,22 +25,20 @@ def clear(self) -> None: def get_or_create( self, - lock: "threading.RLock", resolve: typing.Callable[[], _R], create: typing.Callable[[_R], _V], ) -> tuple[_V, bool]: - """Return the memoized singleton, or resolve-and-create it once under `lock`. + """Return the memoized singleton, or resolve-and-create it once under this item's lock. - `resolve()` runs unlocked — recursive resolution must not hold the lock; creation and - the store are double-checked under it. `created` is True only for the caller that built. + 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 - resolved = resolve() - with lock: + with self.lock: if self.cache is not types.UNSET: return self.cache, False - value = create(resolved) + value = create(resolve()) self.cache = value return value, True @@ -92,7 +88,7 @@ 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 slot for a cached ``provider``, creating it on first use.""" + """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. item = self._items.get(provider.provider_id) if item is not None: diff --git a/modern_di/resolver_compiler.py b/modern_di/resolver_compiler.py index cabfbc9e..406c18e9 100644 --- a/modern_di/resolver_compiler.py +++ b/modern_di/resolver_compiler.py @@ -9,7 +9,7 @@ ``ProvidersRegistry.drop_resolvers``). Why a template and not shared helpers: see docs/adr/0001-resolver-hot-path-generated-source.md. -The template reaches into `Container._lock`/`_scope_map` and `CacheRegistry._items` to stay +The template reaches into `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. """ @@ -114,7 +114,7 @@ def compile_resolver(provider: "AbstractProvider[typing.Any]", registry: "Provid cached = cache_item.cache if cached is not UNSET: return cached - value, created = cache_item.get_or_create(target._lock, resolve=partial(build, target), create=create) + value, created = cache_item.get_or_create(partial(build, target), create) if created: cache_registry.mark_created(cache_item) return value diff --git a/tests/helpers.py b/tests/helpers.py new file mode 100644 index 00000000..81a58c4b --- /dev/null +++ b/tests/helpers.py @@ -0,0 +1,10 @@ +import typing + +from modern_di import Container +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) diff --git a/tests/providers/test_cached_factory.py b/tests/providers/test_cached_factory.py index aacca6b0..672c53b6 100644 --- a/tests/providers/test_cached_factory.py +++ b/tests/providers/test_cached_factory.py @@ -8,6 +8,7 @@ from modern_di import Container, Group, Scope, providers from modern_di.exceptions import AsyncFinalizerInSyncCloseError, ContainerClosedError, FinalizerError, ModernDIError from modern_di.types import UNSET +from tests.helpers import cache_item @dataclasses.dataclass(kw_only=True, slots=True, frozen=True) @@ -373,44 +374,136 @@ class NoneGroup(Group): assert cleaned_up == [None] -class _Gate: ... +class _Conn: ... -class _GatedValue: - def __init__(self, gate: _Gate) -> None: - self.gate = gate +class _Svc: + def __init__(self, conn: _Conn) -> None: + self.conn = conn -def test_concurrent_cache_misses_create_once() -> None: - """Threads that all miss the cache together still run the creator once. +class _CountingRLock: + """An RLock that counts the threads blocked in `acquire`, so a test can wait for contention.""" - The barrier sits in a dependency, and dependencies resolve outside the lock, so every thread is - past the unlocked cache check before any of them creates. If the lock covered dependency - resolution, the barrier would time out and the test would fail instead of hanging. + def __init__(self) -> None: + self._lock = threading.RLock() + self._state = threading.Condition() + self.waiting = 0 + + def __enter__(self) -> None: + with self._state: + self.waiting += 1 + self._state.notify_all() + self._lock.acquire() + with self._state: + self.waiting -= 1 + + def __exit__(self, *_: object) -> None: + self._lock.release() + + def wait_for_waiters(self, count: int) -> bool: + with self._state: + return self._state.wait_for(lambda: self.waiting >= count, timeout=5) + + +def test_concurrent_cache_misses_build_the_value_and_its_dependencies_once() -> None: + """Threads that miss a cached `Svc(conn: Conn)` together build one Svc and one transient Conn. + + The first Conn build waits until every other thread is blocked on the item's lock, so a + dependency resolved outside that lock would be built once per thread and fail the count. """ - n = 4 + n = 8 + conns: list[_Conn] = [] + lock = _CountingRLock() + + def make_conn() -> _Conn: + conn = _Conn() + conns.append(conn) + if len(conns) == 1: + lock.wait_for_waiters(n - 1) + return conn + + class G(Group): + conn = providers.Factory(creator=make_conn) + svc = providers.Factory(creator=_Svc, cache=True) + + container = Container(groups=[G]) + cache_item(container, G.svc).lock = lock # ty: ignore[invalid-assignment] barrier = threading.Barrier(n, timeout=5) - created: list[_GatedValue] = [] - def make_gate() -> _Gate: + def worker() -> _Svc: barrier.wait() - return _Gate() + return container.resolve(_Svc) + + with ThreadPoolExecutor(max_workers=n) as pool: + results = [f.result(timeout=10) for f in [pool.submit(worker) for _ in range(n)]] + + assert len(conns) == 1 + assert all(result is results[0] for result in results) + assert results[0].conn is conns[0] + + +class _Blocked: ... - def make_value(gate: _Gate) -> _GatedValue: - value = _GatedValue(gate) - created.append(value) - return value + +class _Free: ... + + +def test_unrelated_cached_items_are_created_concurrently() -> None: + """One cached creator blocked mid-creation does not block another item's creation.""" + started = threading.Event() + release = threading.Event() + + def make_blocked() -> _Blocked: + started.set() + release.wait(timeout=5) + return _Blocked() class G(Group): - gate = providers.Factory(creator=make_gate) - value = providers.Factory(creator=make_value, cache=True) + blocked = providers.Factory(creator=make_blocked, cache=True) + free = providers.Factory(creator=_Free, cache=True) container = Container(groups=[G]) - with ThreadPoolExecutor(max_workers=n) as pool: - results = list(pool.map(lambda _: container.resolve(_GatedValue), range(n))) + blocked = threading.Thread(target=container.resolve, args=(_Blocked,), daemon=True) + blocked.start() + try: + assert started.wait(timeout=5) + free: list[_Free] = [] + worker = threading.Thread(target=lambda: free.append(container.resolve(_Free)), daemon=True) + worker.start() + worker.join(timeout=5) + assert not worker.is_alive(), "creating one cached item waited for another item's creation" + finally: + release.set() + blocked.join(timeout=5) + assert isinstance(free[0], _Free) + + +class _OnWorker: ... + - assert len(created) == 1 - assert all(result is created[0] for result in results) +class _ViaWorker: + def __init__(self, inner: _OnWorker) -> None: + self.inner = inner + + +def test_cached_creator_resolving_a_cached_type_on_another_thread_completes() -> None: + """A cached creator may hand a resolve of another cached type to a worker thread and wait for it.""" + pool = ThreadPoolExecutor(max_workers=1) + + def make_via_worker(container: Container) -> _ViaWorker: + return _ViaWorker(pool.submit(lambda: container.resolve(_OnWorker)).result(timeout=5)) + + class G(Group): + on_worker = providers.Factory(creator=_OnWorker, cache=True) + via_worker = providers.Factory(creator=make_via_worker, cache=True) + + container = Container(groups=[G]) + try: + result = container.resolve(_ViaWorker) + finally: + pool.shutdown(wait=False) + assert result.inner is container.resolve(_OnWorker) _lifo_events: list[str] = [] diff --git a/tests/registries/test_cache_registry.py b/tests/registries/test_cache_registry.py index 33d56caa..2f249be0 100644 --- a/tests/registries/test_cache_registry.py +++ b/tests/registries/test_cache_registry.py @@ -1,9 +1,10 @@ import threading import typing +from concurrent.futures import ThreadPoolExecutor import pytest -from modern_di.providers import CacheSettings +from modern_di.providers import CacheSettings, Factory from modern_di.registries.cache_registry import CacheItem, CacheRegistry from modern_di.types import UNSET @@ -12,19 +13,36 @@ def _item() -> CacheItem: return CacheItem(settings=CacheSettings()) -def test_get_or_create_miss_calls_resolve_and_create_once_and_caches() -> None: +def _acquirable_from_another_thread(item: CacheItem) -> bool: + """Whether another thread can take `item`'s lock right now; releases it again if so.""" + acquired: list[bool] = [] + + def try_acquire() -> None: + got = item.lock.acquire(blocking=False) + if got: + item.lock.release() + acquired.append(got) + + thread = threading.Thread(target=try_acquire) + thread.start() + thread.join(timeout=5) + return acquired == [True] + + +def test_get_or_create_miss_resolves_and_creates_once_under_the_item_lock() -> None: item = _item() calls = {"resolve": 0, "create": 0} def resolve() -> dict[str, typing.Any]: calls["resolve"] += 1 + assert not _acquirable_from_another_thread(item) return {"x": 1} 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(threading.RLock(), resolve=resolve, create=create) + value, created = item.get_or_create(resolve=resolve, create=create) assert created is True assert value == ("made", {"x": 1}) @@ -44,55 +62,72 @@ def create(_: object) -> str: # pragma: no cover - a cache hit must not create msg = "create must not run on a cache hit" raise AssertionError(msg) - value, created = item.get_or_create(threading.RLock(), resolve=resolve, create=create) + value, created = item.get_or_create(resolve=resolve, create=create) assert created is False assert value == "cached" -def test_get_or_create_double_checks_after_lock() -> None: - # The inner re-check fires when the cache is UNSET at the fast read but SET - # by the time the lock is held (a losing thread in production). Simulate it - # deterministically: resolve() sets the cache as a side effect, so the - # post-lock re-check must return it and skip create. +class _LosingRaceLock: + """Stores a value in the item while the caller waits for the lock, as a winning thread would.""" + + def __init__(self, item: CacheItem) -> None: + self._item = item + + def __enter__(self) -> None: + self._item.cache = "won-the-race" + + def __exit__(self, *_: object) -> None: + pass + + +def test_get_or_create_double_checks_under_the_lock() -> None: item = _item() - created_calls: list[object] = [] + item.lock = _LosingRaceLock(item) # ty: ignore[invalid-assignment] - def resolve() -> dict[str, typing.Any]: - item.cache = "won-the-race" - return {} + def resolve() -> object: # pragma: no cover - the re-check under the lock must skip resolve + msg = "resolve must not run when another thread already stored the value" + raise AssertionError(msg) - def create(kwargs: dict[str, typing.Any]) -> str: # pragma: no cover - the post-lock re-check must skip create - created_calls.append(kwargs) - return "should-not-be-used" + def create(_: object) -> str: # pragma: no cover - the re-check under the lock must skip create + msg = "create must not run when another thread already stored the value" + raise AssertionError(msg) - value, created = item.get_or_create(threading.RLock(), resolve=resolve, create=create) + value, created = item.get_or_create(resolve=resolve, create=create) assert created is False assert value == "won-the-race" - assert created_calls == [] -def test_get_or_create_releases_lock_and_fast_path_on_second_call() -> None: +def test_get_or_create_releases_the_item_lock() -> None: item = _item() - lock = threading.RLock() - value, created = item.get_or_create(lock, resolve=lambda: 0, create=lambda _: "v") + value, created = item.get_or_create(resolve=lambda: 0, create=lambda _: "v") assert (value, created) == ("v", True) + assert _acquirable_from_another_thread(item) - # Second call hits the fast path (cache set) — returns before touching the lock. - value2, created2 = item.get_or_create(lock, resolve=lambda: 0, create=lambda _: "v2") - assert (value2, created2) == ("v", False) - acquired: list[bool] = [] +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)) + assert first.lock is not second.lock - def try_acquire() -> None: - acquired.append(lock.acquire(blocking=False)) - thread = threading.Thread(target=try_acquire) - thread.start() - thread.join() - assert acquired == [True] +def test_concurrent_fetches_of_one_provider_share_one_item() -> None: + n = 8 + registry = CacheRegistry() + 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) + + with ThreadPoolExecutor(max_workers=n) as pool: + items = [f.result(timeout=5) for f in [pool.submit(fetch) for _ in range(n)]] + + 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: diff --git a/tests/test_container.py b/tests/test_container.py index aedfdccc..a7be45c7 100644 --- a/tests/test_container.py +++ b/tests/test_container.py @@ -25,6 +25,7 @@ ValidationFailedError, ) from modern_di.providers.abstract import AbstractProvider +from tests.helpers import cache_item def test_container_prevent_copy() -> None: @@ -385,14 +386,6 @@ def test_constructor_rejects_use_lock() -> None: Container(use_lock=False) # ty: ignore[unknown-argument] -def test_child_shares_the_root_lock() -> None: - root = Container() - child = root.build_child_container(scope=Scope.REQUEST) - grandchild = Container(scope=Scope.ACTION, parent_container=child) - assert child._lock is root._lock - assert grandchild._lock is root._lock - - def test_container_provider_resolves_on_subclasses() -> None: class MyContainer(Container): pass @@ -616,8 +609,8 @@ def test_fresh_container_builds_child_and_child_resolves_without_open() -> None: assert app.closed is False # building a child does not close the parent -def test_warm_cached_resolve_does_not_wait_for_the_lock() -> None: - """A cached value resolves while another thread holds the container lock. +def test_warm_cached_resolve_does_not_wait_for_the_item_lock() -> None: + """A cached value resolves while another thread holds its cache item's lock. The lock is taken only on a cache miss. If a warm resolve took it too, the worker would block until the join timeout and the test would fail instead of hanging. @@ -627,7 +620,7 @@ def test_warm_cached_resolve_does_not_wait_for_the_lock() -> None: child = root.build_child_container(scope=Scope.REQUEST) results: list[_PersistentBroker] = [] worker = threading.Thread(target=lambda: results.append(child.resolve(_PersistentBroker)), daemon=True) - with root._lock: + with cache_item(root, _AppBrokerGroup.broker).lock: worker.start() worker.join(timeout=5) finished_while_held = not worker.is_alive() diff --git a/tests/test_runtime_cycle_guard.py b/tests/test_runtime_cycle_guard.py index 4d636492..7b30b7da 100644 --- a/tests/test_runtime_cycle_guard.py +++ b/tests/test_runtime_cycle_guard.py @@ -9,6 +9,7 @@ import dataclasses import inspect import sys +import threading import pytest @@ -108,6 +109,46 @@ def test_unvalidated_cycle_raises_circular_dependency_error() -> None: sys.setrecursionlimit(original_limit) +class CachedCycleGroup(Group): + common = providers.Factory(creator=Common, cache=True) + a = providers.Factory(creator=NodeA, cache=True) + b = providers.Factory(creator=NodeB, cache=True) + + +def _resolve_in_daemon_thread(container: Container) -> BaseException | None: + """Resolve `NodeA` on a daemon thread under the shallow limit; return what it raised, or None if it hung.""" + raised: list[BaseException] = [] + + def worker() -> None: + try: + container.resolve(NodeA) + except Exception as exc: # noqa: BLE001 + raised.append(exc) + + original_limit = sys.getrecursionlimit() + sys.setrecursionlimit(_SHALLOW_RECURSION_LIMIT) + try: + thread = threading.Thread(target=worker, daemon=True) + thread.start() + thread.join(timeout=5) + finally: + sys.setrecursionlimit(original_limit) + return raised[0] if raised else None + + +def test_cached_cycle_reenters_the_item_lock_and_raises_circular_dependency_error() -> None: + """A same-thread cycle through cached factories re-enters each item's lock and still raises. + + The second resolve runs on another thread, so it hangs instead of raising if the first left + an item's lock held. + """ + container = Container(groups=[CachedCycleGroup]) + for _ in range(2): + exc = _resolve_in_daemon_thread(container) + assert isinstance(exc, exceptions.CircularDependencyError) + _assert_simple_cycle(exc) + + def _assert_deep_chain_cycle_is_self_contained(exc: exceptions.CircularDependencyError) -> None: # Reached via the Root -> Middle -> DeepNodeA approach path, but the cycle itself is only # DeepNodeA <-> DeepNodeB: CircularDependencyError.prepend_step is a no-op (ERR-1 canonicalization),