diff --git a/dapr/actor/runtime/_state_provider.py b/dapr/actor/runtime/_state_provider.py index eeb1e4995..86752d150 100644 --- a/dapr/actor/runtime/_state_provider.py +++ b/dapr/actor/runtime/_state_provider.py @@ -51,6 +51,10 @@ async def try_load_state( result = self._state_serializer.deserialize(raw_state_value, state_type) return (True, result) + def round_trip_state_value(self, value: Any) -> Any: + """Returns value as try_load_state would decode it after a save, without a store read.""" + return self._state_serializer.deserialize(self._state_serializer.serialize(value), object) + async def contains_state(self, actor_type: str, actor_id: str, state_name: str) -> bool: raw_state_value = await self._state_client.get_state(actor_type, actor_id, state_name) return (raw_state_value is not None) and len(raw_state_value) > 0 diff --git a/dapr/actor/runtime/state_manager.py b/dapr/actor/runtime/state_manager.py index c2882debb..37b1a6d3e 100644 --- a/dapr/actor/runtime/state_manager.py +++ b/dapr/actor/runtime/state_manager.py @@ -267,6 +267,32 @@ 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._refresh_default_tracker(state_changes) + + def _refresh_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. Refresh its clean + # copies of the written keys in place, in the shape a fresh read would return, + # and drop removed keys. Entries with pending changes are left alone. + state_provider = self._actor.runtime_ctx.state_provider + for change in state_changes: + metadata = self._default_state_change_tracker.get(change.state_name) + if metadata is None or metadata.change_kind != StateChangeKind.none: + continue + # A None value is not written to the store, so let the next read reload it. + if change.change_kind == StateChangeKind.remove or change.value is None: + self._default_state_change_tracker.pop(change.state_name) + continue + try: + value = state_provider.round_trip_state_value(change.value) + except Exception: + # The save has already committed; fall back to reloading on the next read. + self._default_state_change_tracker.pop(change.state_name) + continue + self._default_state_change_tracker[change.state_name] = StateMetadata( + value, StateChangeKind.none, change.ttl_in_seconds + ) 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..658959b6a 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,139 @@ def test_get_state_for_removed_value(self): self.assertFalse(has_value) self.assertIsNone(val) + def _run_reentrant(self, state_manager, coro_fn): + # Runs coro_fn inside a reentrancy-scoped call, then saves its tracker. + token = reentrancy_ctx.set('reentrancy-id') + try: + state_manager.set_state_context('ctx1') + try: + _run(coro_fn()) + _run(state_manager.save_state()) + finally: + state_manager.set_state_context(None) + finally: + reentrancy_ctx.reset(token) + + @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_update_refreshes_default_tracker(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.try_get_state('state1')) + + self._run_reentrant(state_manager, lambda: state_manager.set_state_ttl('state1', 'v2', 60)) + + # The default read is served from the refreshed entry, not the state store. + calls = self._fake_client.get_state.mock.call_count + has_value, val = _run(state_manager.try_get_state('state1')) + self.assertTrue(has_value) + self.assertEqual('v2', val) + self.assertEqual(calls, self._fake_client.get_state.mock.call_count) + state = state_manager._default_state_change_tracker['state1'] + self.assertEqual(StateChangeKind.none, state.change_kind) + self.assertEqual(60, state.ttl_in_seconds) + + @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_remove_evicts_default_tracker(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.try_get_state('state1')) + + self._run_reentrant(state_manager, lambda: state_manager.remove_state('state1')) + + self.assertNotIn('state1', state_manager._default_state_change_tracker) + self._fake_client.get_state.mock.return_value = None + has_value, val = _run(state_manager.try_get_state('state1')) + 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_keeps_dirty_default_entry(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.set_state('state1', 'pending')) + + self._run_reentrant(state_manager, lambda: state_manager.set_state('state1', 'v2')) + + state = state_manager._default_state_change_tracker['state1'] + self.assertEqual('pending', state.value) + self.assertEqual(StateChangeKind.update, state.change_kind) + + @mock.patch( + 'tests.actor.fake_client.FakeDaprActorClient.get_state', + new=_async_mock(return_value=b'[1, 2]'), + ) + @mock.patch( + 'tests.actor.fake_client.FakeDaprActorClient.save_state_transactionally', new=_async_mock() + ) + def test_reentrant_refresh_matches_fresh_read(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.try_get_state('state1')) + + self._run_reentrant(state_manager, lambda: state_manager.set_state('state1', (3, 4))) + + _, refreshed = _run(state_manager.try_get_state('state1')) + self._fake_client.get_state.mock.return_value = self._serializer.serialize((3, 4)) + _, fresh = _run(ActorStateManager(self._fake_actor).try_get_state('state1')) + self.assertEqual([3, 4], fresh) + self.assertEqual(fresh, refreshed) + + @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_of_none_evicts_default_tracker(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.try_get_state('state1')) + + self._run_reentrant(state_manager, lambda: state_manager.set_state('state1', None)) + + self.assertNotIn('state1', state_manager._default_state_change_tracker) + + @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_refresh_failure_evicts_default_tracker(self): + state_manager = ActorStateManager(self._fake_actor) + _run(state_manager.try_get_state('state1')) + _run(state_manager.try_get_state('state2')) + + async def set_both(): + await state_manager.set_state('state1', 'v2') + await state_manager.set_state('state2', 'v2') + + # The save must not raise after the write has committed, and every key is handled. + with mock.patch.object( + self._runtime_ctx.state_provider, + 'round_trip_state_value', + side_effect=ValueError('cannot decode'), + ): + self._run_reentrant(state_manager, set_both) + + self.assertNotIn('state1', state_manager._default_state_change_tracker) + self.assertNotIn('state2', state_manager._default_state_change_tracker) + @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)