diff --git a/dapr/actor/runtime/state_manager.py b/dapr/actor/runtime/state_manager.py index c2882debb..03d10998b 100644 --- a/dapr/actor/runtime/state_manager.py +++ b/dapr/actor/runtime/state_manager.py @@ -267,6 +267,17 @@ async def save_state(self) -> None: ) for state_name in states_to_remove: state_change_tracker.pop(state_name, None) + if state_change_tracker is not self._default_state_change_tracker: + self._invalidate_default_tracker(state_changes) + + def _invalidate_default_tracker(self, state_changes: List[ActorStateChange]) -> None: + # Writes made through a reentrancy-scoped tracker are invisible to the default + # tracker, which activation, reminders and timers read from. Drop its clean copies + # of the written keys so the next read reloads them instead of serving stale data. + for change in state_changes: + metadata = self._default_state_change_tracker.get(change.state_name) + if metadata is not None and metadata.change_kind == StateChangeKind.none: + self._default_state_change_tracker.pop(change.state_name) def is_state_marked_for_remove(self, state_name: str) -> bool: state_change_tracker = self._get_contextual_state_tracker() diff --git a/tests/actor/test_state_manager.py b/tests/actor/test_state_manager.py index 5ed4e24e6..566e80457 100644 --- a/tests/actor/test_state_manager.py +++ b/tests/actor/test_state_manager.py @@ -20,6 +20,7 @@ from dapr.actor.id import ActorId from dapr.actor.runtime._type_information import ActorTypeInformation from dapr.actor.runtime.context import ActorRuntimeContext +from dapr.actor.runtime.reentrancy_context import reentrancy_ctx from dapr.actor.runtime.state_change import StateChangeKind from dapr.actor.runtime.state_manager import ActorStateManager, StateMetadata from dapr.serializers import DefaultJSONSerializer @@ -118,6 +119,35 @@ def test_get_state_for_removed_value(self): self.assertFalse(has_value) self.assertIsNone(val) + @mock.patch( + 'tests.actor.fake_client.FakeDaprActorClient.get_state', + new=_async_mock(return_value=b'"value1"'), + ) + @mock.patch( + 'tests.actor.fake_client.FakeDaprActorClient.save_state_transactionally', new=_async_mock() + ) + def test_reentrant_save_invalidates_default_tracker(self): + state_manager = ActorStateManager(self._fake_actor) + + # A read outside any reentrancy context caches the value in the default tracker. + has_value, val = _run(state_manager.try_get_state('state1')) + self.assertTrue(has_value) + self.assertEqual('value1', val) + + # A reentrancy-scoped call writes the same key through its own tracker. + reentrancy_ctx.set('reentrancy-id') + state_manager.set_state_context('ctx1') + _run(state_manager.set_state('state1', 'value2')) + _run(state_manager.save_state()) + state_manager.set_state_context(None) + reentrancy_ctx.set(None) + + # The default tracker must reload the key instead of serving its stale copy. + self._fake_client.get_state.mock.return_value = b'"value2"' + has_value, val = _run(state_manager.try_get_state('state1')) + self.assertTrue(has_value) + self.assertEqual('value2', val) + @mock.patch('tests.actor.fake_client.FakeDaprActorClient.get_state', new=_async_mock()) def test_set_state_for_new_state(self): state_manager = ActorStateManager(self._fake_actor)