From 3f3dd2603414805db4c6b9532579ca99fb66c07f Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Sun, 6 Sep 2026 11:29:40 +0900 Subject: [PATCH] Fix replacement of named provider injections --- docs/providers/factory.rst | 5 + src/dependency_injector/providers.pyx | 26 +++-- .../async/test_named_injection_replacement.py | 101 ++++++++++++++++++ 3 files changed, 126 insertions(+), 6 deletions(-) create mode 100644 tests/unit/providers/async/test_named_injection_replacement.py diff --git a/docs/providers/factory.rst b/docs/providers/factory.rst index f9377baf..600d8101 100644 --- a/docs/providers/factory.rst +++ b/docs/providers/factory.rst @@ -40,6 +40,11 @@ injected following these rules: :language: python :lines: 3- +Calling ``.add_kwargs()`` or ``.add_attributes()`` with an existing name replaces +the previous injection. The replaced provider is not called, including when it +is asynchronous. Call-time keyword arguments still take precedence over +registered keyword injections. + ``Factory`` provider can inject attributes. Use ``.add_attributes()`` method to specify attribute injections. diff --git a/src/dependency_injector/providers.pyx b/src/dependency_injector/providers.pyx index 1935c012..d3f5418c 100644 --- a/src/dependency_injector/providers.pyx +++ b/src/dependency_injector/providers.pyx @@ -1311,7 +1311,7 @@ cdef class Callable(Provider): :return: Reference ``self`` """ - self._kwargs += parse_named_injections(kwargs) + self._kwargs = _merge_named_injections(self._kwargs, kwargs) self._kwargs_len = len(self._kwargs) return self @@ -2622,7 +2622,7 @@ cdef class Factory(Provider): :return: Reference ``self`` """ - self._attributes += parse_named_injections(kwargs) + self._attributes = _merge_named_injections(self._attributes, kwargs) self._attributes_len = len(self._attributes) return self @@ -3534,8 +3534,8 @@ cdef class Dict(Provider): if dict_ is None: dict_ = {} - self._kwargs += parse_named_injections(dict_) - self._kwargs += parse_named_injections(kwargs) + self._kwargs = _merge_named_injections(self._kwargs, dict_) + self._kwargs = _merge_named_injections(self._kwargs, kwargs) self._kwargs_len = len(self._kwargs) return self @@ -3551,7 +3551,7 @@ cdef class Dict(Provider): dict_ = {} self._kwargs = parse_named_injections(dict_) - self._kwargs += parse_named_injections(kwargs) + self._kwargs = _merge_named_injections(self._kwargs, kwargs) self._kwargs_len = len(self._kwargs) return self @@ -3815,7 +3815,7 @@ cdef class BaseResource(Provider): :return: Reference ``self`` """ - self._kwargs += parse_named_injections(kwargs) + self._kwargs = _merge_named_injections(self._kwargs, kwargs) self._kwargs_len = len(self._kwargs) return self @@ -4697,6 +4697,20 @@ cpdef tuple parse_named_injections(dict kwargs): return tuple(injections) +cdef tuple _merge_named_injections(tuple injections, dict kwargs): + """Replace named injections before any superseded provider is evaluated.""" + cdef dict merged = {} + cdef NamedInjection injection + cdef object name + cdef object value + + for injection in injections: + merged[injection._name] = injection + for name, value in kwargs.items(): + merged[name] = NamedInjection(name, value) + return tuple(merged.values()) + + cdef class OverridingContext: """Provider overriding context. diff --git a/tests/unit/providers/async/test_named_injection_replacement.py b/tests/unit/providers/async/test_named_injection_replacement.py new file mode 100644 index 00000000..d8f34e74 --- /dev/null +++ b/tests/unit/providers/async/test_named_injection_replacement.py @@ -0,0 +1,101 @@ +"""Named injection replacement must discard previously registered providers.""" + +import inspect +from types import SimpleNamespace + +from dependency_injector import providers +from pytest import mark + + +@mark.asyncio +@mark.parametrize("provider_type", [providers.Callable, providers.Factory, providers.Singleton, + providers.Resource, providers.Dict]) +@mark.parametrize("asynchronous", [False, True]) +@mark.parametrize("context", [{}, {"extra": 42}]) +async def test_add_kwargs_replaces_existing_injection(provider_type, asynchronous, context): + calls = [] + + def original(): + calls.append(True) + return "original" + + async def original_async(): + return original() + + dependency = providers.Callable(original_async if asynchronous else original) + if provider_type is providers.Dict: + provider = provider_type(value=dependency) + else: + provider = provider_type(dict, value=dependency) + provider.add_kwargs(value="intermediate").add_kwargs(value="replacement") + + result = provider(**context) + if inspect.isawaitable(result): + result = await result + + assert result == dict(value="replacement", **context) + assert calls == [] + assert provider.kwargs == {"value": "replacement"} + + +@mark.asyncio +@mark.parametrize("provider_type", [providers.Factory, providers.Singleton]) +async def test_add_attributes_replaces_async_injection(provider_type): + calls = [] + + async def original(): + calls.append(True) + return "original" + + provider = provider_type(SimpleNamespace) + provider.add_attributes(value=providers.Callable(original)) + provider.add_attributes(value="replacement") + result = provider() + if inspect.isawaitable(result): + result = await result + + assert result.value == "replacement" + assert calls == [] + + +@mark.asyncio +@mark.parametrize("method", ["__init__", "add_kwargs", "set_kwargs"]) +async def test_dict_keyword_overrides_mapping_injection(method): + calls = [] + + async def original(): + calls.append(True) + return "original" + + provider = providers.Dict() + getattr(provider, method)({"value": providers.Callable(original)}, value="replacement") + result = provider() + if inspect.isawaitable(result): + result = await result + + assert result == {"value": "replacement"} + assert calls == [] + + +@mark.asyncio +@mark.parametrize("provider_type", [providers.Callable, providers.Factory, providers.Singleton, + providers.Resource, providers.Dict]) +@mark.parametrize("override", [False, True]) +async def test_async_replacement_and_call_time_precedence(provider_type, override): + calls = [] + + async def replacement(): + calls.append(True) + return "replacement" + + if provider_type is providers.Dict: + provider = provider_type(value="original") + else: + provider = provider_type(dict, value="original") + provider.add_kwargs(value=providers.Callable(replacement)) + result = provider(**({"value": "explicit"} if override else {})) + if inspect.isawaitable(result): + result = await result + + assert result == {"value": "explicit" if override else "replacement"} + assert calls == ([] if override else [True])