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
11 changes: 11 additions & 0 deletions dapr/actor/runtime/state_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Comment thread
JoshVanL marked this conversation as resolved.

def is_state_marked_for_remove(self, state_name: str) -> bool:
state_change_tracker = self._get_contextual_state_tracker()
Expand Down
30 changes: 30 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,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)
Expand Down
Loading