diff --git a/tests/test_meta.py b/tests/test_meta.py index 6f1a9353..bd5bf112 100644 --- a/tests/test_meta.py +++ b/tests/test_meta.py @@ -1,4 +1,9 @@ -from that_depends import ContextScopes +import typing + +import pytest + +from that_depends import BaseContainer, ContextScopes, providers +from that_depends.exceptions import TypeNotBoundError from that_depends.meta import BaseContainerMeta @@ -21,3 +26,18 @@ class _Test(metaclass=BaseContainerMeta): pass assert _Test.name() == "_Test" + + +def test_type_provider_cache_invalidates_after_rebinding() -> None: + provider = providers.Object(1).bind(int) + + class Container(BaseContainer): + value = provider + + assert Container.get_provider_for_type(int) is provider + + provider.bind(str) + + with pytest.raises(TypeNotBoundError): + Container.get_provider_for_type(int) + assert Container.get_provider_for_type(str) is typing.cast(providers.Object[str], provider) diff --git a/that_depends/meta.py b/that_depends/meta.py index 614479d0..f6dff5b2 100644 --- a/that_depends/meta.py +++ b/that_depends/meta.py @@ -62,6 +62,7 @@ def get_scope(cls) -> ContextScope | None: "alias", "default_scope", "type_provider_cache", + "type_provider_cache_revision", ) _lock: Lock = Lock() @@ -133,8 +134,10 @@ def get_provider_for_type(cls, t: type[T]) -> AbstractProvider[T]: Provider for the given type. """ - if not hasattr(cls, "type_provider_cache"): + current_revision = AbstractProvider._get_binding_revision() # noqa: SLF001 + if getattr(cls, "type_provider_cache_revision", -1) != current_revision: cls.type_provider_cache: dict[type[typing.Any], AbstractProvider[typing.Any]] = {} + cls.type_provider_cache_revision = current_revision if provider := cls.type_provider_cache.get(t): return typing.cast(AbstractProvider[T], provider) for provider in cls.get_providers().values(): diff --git a/that_depends/providers/base.py b/that_depends/providers/base.py index 31f59820..114cab9b 100644 --- a/that_depends/providers/base.py +++ b/that_depends/providers/base.py @@ -89,6 +89,9 @@ def _resolve_keyword_arguments_sync( class AbstractProvider(abc.ABC, typing.Generic[T_co]): """Base class for all providers.""" + _binding_revision: typing.ClassVar[int] = 0 + _binding_revision_lock: typing.ClassVar[threading.Lock] = threading.Lock() + def __init__(self) -> None: """Create a new provider.""" super().__init__() @@ -115,10 +118,17 @@ def bind(self, *types: type, contravariant: bool = False) -> typing_extensions.S The current provider instance. """ - self._bindings = set(types) - self._has_contravariant_bindings = contravariant + with AbstractProvider._binding_revision_lock: + self._bindings = set(types) + self._has_contravariant_bindings = contravariant + AbstractProvider._binding_revision += 1 return self + @classmethod + def _get_binding_revision(cls) -> int: + with AbstractProvider._binding_revision_lock: + return AbstractProvider._binding_revision + def _register(self, candidates: typing.Iterable[typing.Any]) -> None: """Register current provider as child.