Skip to content

Commit 9d3814c

Browse files
author
remimd
committed
feat: Bidirectional link possible between modules
1 parent 8ad8760 commit 9d3814c

4 files changed

Lines changed: 76 additions & 20 deletions

File tree

injection/_core/module.py

Lines changed: 43 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -120,15 +120,24 @@ def __str__(self) -> str:
120120
return f"`{self.module}` has propagated an event: {self.origin}"
121121

122122
@property
123-
def history(self) -> Iterator[Event]:
124-
if isinstance(self.event, ModuleEventProxy):
125-
yield from self.event.history
126-
127-
yield self.event
123+
def is_duplicate(self) -> bool:
124+
return any(
125+
self.module is event.module and self.origin is event.origin
126+
for event in self.proxy_history
127+
)
128128

129129
@property
130130
def origin(self) -> Event:
131-
return next(self.history)
131+
reversed_proxy_history = reversed(tuple(self.proxy_history))
132+
return next(reversed_proxy_history, self).event
133+
134+
@property
135+
def proxy_history(self) -> Iterator[ModuleEventProxy]:
136+
event = self.event
137+
138+
if isinstance(event, ModuleEventProxy):
139+
yield event
140+
yield from event.proxy_history
132141

133142

134143
@dataclass(frozen=True, slots=True)
@@ -162,8 +171,10 @@ def __str__(self) -> str:
162171

163172
@dataclass(frozen=True, slots=True)
164173
class UnlockCalled(Event):
174+
module: Module
175+
165176
def __str__(self) -> str:
166-
return "An `unlock` method has been called."
177+
return f"`{self.module}.unlock` has been called."
167178

168179

169180
"""
@@ -420,23 +431,18 @@ def __post_init__(self) -> None:
420431
self.__locator.add_listener(self)
421432

422433
def __getitem__[T](self, cls: InputType[T], /) -> Injectable[T]:
423-
for broker in self.__brokers:
434+
for broker in self._iter_brokers():
424435
with suppress(KeyError):
425436
return broker[cls]
426437

427438
raise NoInjectable(cls)
428439

429440
def __contains__(self, cls: InputType[Any], /) -> bool:
430-
return any(cls in broker for broker in self.__brokers)
441+
return any(cls in broker for broker in self._iter_brokers())
431442

432443
@property
433444
def is_locked(self) -> bool:
434-
return any(broker.is_locked for broker in self.__brokers)
435-
436-
@property
437-
def __brokers(self) -> Iterator[Broker]:
438-
yield from self.__modules
439-
yield self.__locator
445+
return any(broker.is_locked for broker in self._iter_brokers())
440446

441447
def injectable[**P, T](
442448
self,
@@ -857,19 +863,19 @@ def change_priority(self, module: Module, priority: Priority | PriorityStr) -> S
857863
return self
858864

859865
def unlock(self) -> Self:
860-
event = UnlockCalled()
866+
event = UnlockCalled(self)
861867

862868
with self.dispatch(event, lock_bypass=True):
863869
self.unsafe_unlocking()
864870

865871
return self
866872

867873
def unsafe_unlocking(self) -> None:
868-
for broker in self.__brokers:
874+
for broker in self._iter_brokers():
869875
broker.unsafe_unlocking()
870876

871877
async def all_ready(self) -> None:
872-
for broker in self.__brokers:
878+
for broker in self._iter_brokers():
873879
await broker.all_ready()
874880

875881
def add_logger(self, logger: Logger) -> Self:
@@ -884,8 +890,12 @@ def remove_listener(self, listener: EventListener) -> Self:
884890
self.__channel.remove_listener(listener)
885891
return self
886892

887-
def on_event(self, event: Event, /) -> ContextManager[None]:
893+
def on_event(self, event: Event, /) -> ContextManager[None] | None:
888894
self_event = ModuleEventProxy(self, event)
895+
896+
if self_event.is_duplicate:
897+
return None
898+
889899
return self.dispatch(self_event)
890900

891901
@contextmanager
@@ -899,6 +909,20 @@ def dispatch(self, event: Event, *, lock_bypass: bool = False) -> Iterator[None]
899909
finally:
900910
self.__debug(event)
901911

912+
def _iter_brokers(self, visited: set[Module] | None = None, /) -> Iterator[Broker]:
913+
if visited is None:
914+
visited = set()
915+
916+
if self in visited:
917+
return
918+
919+
visited.add(self)
920+
921+
for module in self.__modules:
922+
yield from module._iter_brokers(visited)
923+
924+
yield self.__locator
925+
902926
def __debug(self, message: object) -> None:
903927
for logger in self.__loggers:
904928
logger.debug(message)

injection/loaders.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -176,12 +176,12 @@ def _unload(self, name: str, /) -> None:
176176

177177
def __init_subsets_for(self, module: Module) -> Module:
178178
if not self.__is_empty and not self.__is_initialized(module):
179+
self.__mark_initialized(module)
179180
target_modules = tuple(
180181
self.__init_subsets_for(mod(name))
181182
for name in self.module_subsets.get(module.name, ())
182183
)
183184
module.init_modules(*target_modules)
184-
self.__mark_initialized(module)
185185

186186
return module
187187

tests/core/test_module.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -292,6 +292,16 @@ def test_use_with_module_already_in_use_raise_module_error(
292292

293293
event_history.assert_length(1)
294294

295+
def test_use_with_bidirectional_use(self, module, event_history):
296+
second_module = Module()
297+
third_module = Module()
298+
299+
module.use(second_module)
300+
second_module.use(module)
301+
module.use(third_module)
302+
303+
event_history.assert_length(4)
304+
295305
"""
296306
stop_using
297307
"""

tests/loaders/test_profile_loader.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,28 @@ class B(A): ...
108108

109109
assert type(find_instance(A)) is A
110110

111+
def test_load_with_bidirectional_link(self):
112+
profile_name_1 = uuid4().hex
113+
profile_name_2 = uuid4().hex
114+
115+
@mod(profile_name_1).injectable
116+
class BaseConfig: ...
117+
118+
@mod(profile_name_2).injectable
119+
@dataclass
120+
class Dependency:
121+
config: BaseConfig
122+
123+
loader = ProfileLoader(
124+
{
125+
profile_name_1: [profile_name_2],
126+
profile_name_2: [profile_name_1],
127+
}
128+
)
129+
130+
with loader.load(profile_name_1):
131+
assert find_instance(Dependency)
132+
111133
def test_load_with_default_profile_do_nothing(self):
112134
default_profile_name = mod().name
113135
global_profile_name = uuid4().hex

0 commit comments

Comments
 (0)