Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions docs/providers/factory.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
26 changes: 20 additions & 6 deletions src/dependency_injector/providers.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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.

Expand Down
101 changes: 101 additions & 0 deletions tests/unit/providers/async/test_named_injection_replacement.py
Original file line number Diff line number Diff line change
@@ -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])