diff --git a/CHANGES b/CHANGES index ae5df99..eaccdf9 100644 --- a/CHANGES +++ b/CHANGES @@ -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 ------ diff --git a/injector/__init__.py b/injector/__init__.py index 55f7ab2..c27e73d 100644 --- a/injector/__init__.py +++ b/injector/__init__.py @@ -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( @@ -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) @@ -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: @@ -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: @@ -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: diff --git a/injector_test.py b/injector_test.py index 917f34e..9e02507 100644 --- a/injector_test.py +++ b/injector_test.py @@ -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