Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions dapr/actor/runtime/_state_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
26 changes: 26 additions & 0 deletions dapr/actor/runtime/state_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
134 changes: 134 additions & 0 deletions tests/actor/test_state_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading