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
1 change: 1 addition & 0 deletions CHANGES
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ Injector Change Log
------

- Dropped support for Python 3.8 and 3.9. Python 3.10+ is now required.
- Child injectors now respect and extend parent injector multibindings.

0.24.0
------
Expand Down
18 changes: 15 additions & 3 deletions injector/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -330,9 +330,11 @@ class MultiBinder(Provider, Generic[T]):

_multi_bindings: List['Binding']

def __init__(self, parent: 'Binder') -> None:
def __init__(self, parent: 'Binder', interface: Optional[type] = None) -> None:
self._multi_bindings = []
self._binder = Binder(parent.injector, auto_bind=False, parent=parent)
self._owner_binder = parent
self._interface = interface

@abstractmethod
def multibind(
Expand All @@ -348,6 +350,14 @@ def append(self, provider: Provider[T], scope: Type['Scope']) -> None:
self._multi_bindings.append(Binding(pseudo_type, provider, scope))

def get_scoped_providers(self, injector: 'Injector') -> Generator[Provider[T], None, None]:
if self._owner_binder.parent and self._interface is not None:
try:
parent_binding, _ = self._owner_binder.parent._get_binding(self._interface)
if isinstance(parent_binding.provider, MultiBinder):
yield from parent_binding.provider.get_scoped_providers(injector)
except KeyError:
pass

for binding in self._multi_bindings:
scope_binding, _ = self._binder.get_binding(binding.scope)
scope_instance: Scope = scope_binding.provider.get(injector)
Expand All @@ -365,6 +375,7 @@ class MultiBindProvider(MultiBinder[List[T]]):
def multibind(
self, interface: type, to: Any, scope: Union['ScopeDecorator', Type['Scope'], None]
) -> None:
self._interface = interface
try:
element_type = get_args(_punch_through_alias(interface))[0]
except IndexError:
Expand All @@ -391,6 +402,7 @@ class MapBindProvider(MultiBinder[Dict[str, T]]):
def multibind(
self, interface: type, to: Any, scope: Union['ScopeDecorator', Type['Scope'], None]
) -> None:
self._interface = interface
try:
value_type = get_args(_punch_through_alias(interface))[1]
except IndexError:
Expand Down Expand Up @@ -582,9 +594,9 @@ def _get_multi_binder(self, interface: type) -> MultiBinder:
and issubclass(interface, dict)
or _get_origin(_punch_through_alias(interface)) is dict
):
multi_binder = MapBindProvider(self)
multi_binder = MapBindProvider(self, interface)
else:
multi_binder = MultiBindProvider(self)
multi_binder = MultiBindProvider(self, interface)
binding = self.create_binding(interface, multi_binder)
self._bindings[interface] = binding
else:
Expand Down
51 changes: 51 additions & 0 deletions injector_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -923,6 +923,57 @@ def configure(binder: Binder) -> None:
assert injector.get(PluginA) is not injector.get(PluginA)


def test_multibinds_are_extended_by_child_injectors() -> None:
parent_injector = Injector()
parent_injector.binder.multibind(List[str], to=['parent name'])

child_injector = parent_injector.create_child_injector()
child_injector.binder.multibind(List[str], to=['child name'])

assert parent_injector.get(List[str]) == ['parent name']
assert child_injector.get(List[str]) == ['parent name', 'child name']


def test_multibind_dict_is_extended_by_child_injectors() -> None:
parent_injector = Injector()
parent_injector.binder.multibind(Dict[str, str], to={'parent': 'p', 'shared': 'from_parent'})

child_injector = parent_injector.create_child_injector()
child_injector.binder.multibind(Dict[str, str], to={'child': 'c', 'shared': 'from_child'})

assert parent_injector.get(Dict[str, str]) == {'parent': 'p', 'shared': 'from_parent'}
assert child_injector.get(Dict[str, str]) == {'parent': 'p', 'shared': 'from_child', 'child': 'c'}


def test_multibind_multi_level_hierarchy_extended_by_child_injectors() -> None:
parent = Injector()
parent.binder.multibind(List[str], to=['parent'])

child = parent.create_child_injector()
child.binder.multibind(List[str], to=['child'])

grandchild = child.create_child_injector()
grandchild.binder.multibind(List[str], to=['grandchild'])

assert parent.get(List[str]) == ['parent']
assert child.get(List[str]) == ['parent', 'child']
assert grandchild.get(List[str]) == ['parent', 'child', 'grandchild']


def test_multibind_skipped_level_child_injector() -> None:
parent = Injector()
parent.binder.multibind(List[str], to=['parent'])

child = parent.create_child_injector()

grandchild = child.create_child_injector()
grandchild.binder.multibind(List[str], to=['grandchild'])

assert parent.get(List[str]) == ['parent']
assert child.get(List[str]) == ['parent']
assert grandchild.get(List[str]) == ['parent', 'grandchild']


def test_regular_bind_and_provider_dont_work_with_multibind():
# We only want multibind and multiprovider to work to avoid confusion

Expand Down
Loading