diff --git a/src/braket/tracking/tracking_context.py b/src/braket/tracking/tracking_context.py index 65b8ea42b..f3875dd63 100644 --- a/src/braket/tracking/tracking_context.py +++ b/src/braket/tracking/tracking_context.py @@ -13,10 +13,13 @@ from __future__ import annotations +import threading + class TrackingContext: def __init__(self): self._trackers = set() + self._lock = threading.Lock() def register_tracker(self, tracker: Tracker) -> None: # ruff:ignore[undefined-name] """Registers a tracker. @@ -24,7 +27,8 @@ def register_tracker(self, tracker: Tracker) -> None: # ruff:ignore[undefined-n Args: tracker (Tracker): The tracker. """ - self._trackers.add(tracker) + with self._lock: + self._trackers.add(tracker) def deregister_tracker(self, tracker: Tracker) -> None: # ruff:ignore[undefined-name] """Deregisters a tracker. @@ -32,7 +36,8 @@ def deregister_tracker(self, tracker: Tracker) -> None: # ruff:ignore[undefined Args: tracker (Tracker): The tracker. """ - self._trackers.remove(tracker) + with self._lock: + self._trackers.remove(tracker) def broadcast_event(self, event: _TrackingEvent) -> None: # ruff:ignore[undefined-name] """Broadcasts an event to all trackers. @@ -40,7 +45,11 @@ def broadcast_event(self, event: _TrackingEvent) -> None: # ruff:ignore[undefin Args: event (_TrackingEvent): The event to broadcast. """ - for tracker in self._trackers: + # Iterate over a snapshot so that trackers registering or deregistering + # concurrently (or from receive_event) cannot mutate the set mid-iteration. + with self._lock: + trackers = list(self._trackers) + for tracker in trackers: tracker.receive_event(event) def active_trackers(self) -> set: diff --git a/test/unit_tests/braket/tracking/test_tracking_context.py b/test/unit_tests/braket/tracking/test_tracking_context.py index 2963bbd0b..3a11cd9c6 100644 --- a/test/unit_tests/braket/tracking/test_tracking_context.py +++ b/test/unit_tests/braket/tracking/test_tracking_context.py @@ -42,3 +42,20 @@ def test_broadcast_event(): broadcast_event("EVENT") tracker.receive_event.assert_called_with("EVENT") deregister_tracker(tracker) + + +def test_broadcast_event_tracker_deregisters_during_broadcast(): + class SelfDeregisteringTracker: + def __init__(self): + self.events = [] + + def receive_event(self, event): + self.events.append(event) + deregister_tracker(self) + + trackers = [SelfDeregisteringTracker() for _ in range(2)] + for tracker in trackers: + register_tracker(tracker) + broadcast_event("EVENT") + assert all(tracker.events == ["EVENT"] for tracker in trackers) + assert active_trackers() == set()