Skip to content
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
83 changes: 64 additions & 19 deletions dapr/actor/runtime/state_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,21 @@ def ttl_in_seconds(self, new_ttl_in_seconds: Optional[int]) -> None:
self._ttl_in_seconds = new_ttl_in_seconds


class _NotFoundStateMetadata(StateMetadata[None]):
# A cached miss: the key is known to be absent from the state store. It reports the
# public 'none' change kind, so it is never saved, and is told apart by its type.
def __init__(self) -> None:
super().__init__(None, StateChangeKind.none)


def _is_not_found(state_metadata: StateMetadata) -> bool:
return isinstance(state_metadata, _NotFoundStateMetadata)


def _is_absent(state_metadata: StateMetadata) -> bool:
return _is_not_found(state_metadata) or state_metadata.change_kind == StateChangeKind.remove


class ActorStateManager(Generic[T]):
def __init__(self, actor: 'Actor'):
self._actor = actor
Expand All @@ -77,6 +92,9 @@ async def try_add_state(self, state_name: str, value: T) -> bool:
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
if _is_not_found(state_metadata):
state_change_tracker[state_name] = StateMetadata(value, StateChangeKind.add)
return True
if state_metadata.change_kind == StateChangeKind.remove:
state_change_tracker[state_name] = StateMetadata(value, StateChangeKind.update)
return True
Expand All @@ -102,14 +120,18 @@ async def try_get_state(self, state_name: str) -> Tuple[bool, Optional[T]]:
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
if state_metadata.change_kind == StateChangeKind.remove:
if _is_absent(state_metadata):
return False, None
return True, state_metadata.value
has_value, val = await self._actor.runtime_ctx.state_provider.try_load_state(
self._type_name, self._actor.id.id, state_name
)
if has_value:
state_change_tracker[state_name] = StateMetadata(val, StateChangeKind.none)
elif state_change_tracker is self._default_state_change_tracker:
# Only the default tracker is refreshed by reentrant saves, so a miss cached
# in a reentrant tracker could hide a nested call's write and overwrite it.
state_change_tracker[state_name] = _NotFoundStateMetadata()
return has_value, val

async def set_state(self, state_name: str, value: T) -> None:
Expand All @@ -122,6 +144,11 @@ async def set_state_ttl(self, state_name: str, value: T, ttl_in_seconds: Optiona
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
if _is_not_found(state_metadata):
state_change_tracker[state_name] = StateMetadata(
value, StateChangeKind.add, ttl_in_seconds
)
return
state_metadata.value = value
state_metadata.ttl_in_seconds = ttl_in_seconds

Expand All @@ -133,17 +160,11 @@ async def set_state_ttl(self, state_name: str, value: T, ttl_in_seconds: Optiona
state_change_tracker[state_name] = state_metadata
return

existed = await self._actor.runtime_ctx.state_provider.contains_state(
self._type_name, self._actor.id.id, state_name
# Add and update are both saved as an upsert, so skip the store read. Update also
# makes a later remove send a delete, since the key may already be in the store.
state_change_tracker[state_name] = StateMetadata(
value, StateChangeKind.update, ttl_in_seconds
)
if existed:
state_change_tracker[state_name] = StateMetadata(
value, StateChangeKind.update, ttl_in_seconds
)
else:
state_change_tracker[state_name] = StateMetadata(
value, StateChangeKind.add, ttl_in_seconds
)

async def remove_state(self, state_name: str) -> None:
if not await self.try_remove_state(state_name):
Expand All @@ -153,7 +174,7 @@ async def try_remove_state(self, state_name: str) -> bool:
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
if state_metadata.change_kind == StateChangeKind.remove:
if _is_absent(state_metadata):
return False
elif state_metadata.change_kind == StateChangeKind.add:
state_change_tracker.pop(state_name, None)
Expand All @@ -172,8 +193,7 @@ async def try_remove_state(self, state_name: str) -> bool:
async def contains_state(self, state_name: str) -> bool:
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
return state_metadata.change_kind != StateChangeKind.remove
return not _is_absent(state_change_tracker[state_name])
return await self._actor.runtime_ctx.state_provider.contains_state(
self._type_name, self._actor.id.id, state_name
)
Expand All @@ -200,6 +220,9 @@ async def add_or_update_state(
state_change_tracker = self._get_contextual_state_tracker()
if state_name in state_change_tracker:
state_metadata = state_change_tracker[state_name]
if _is_not_found(state_metadata):
state_change_tracker[state_name] = StateMetadata(value, StateChangeKind.add)
return value
if state_metadata.change_kind == StateChangeKind.remove:
state_change_tracker[state_name] = StateMetadata(value, StateChangeKind.update)
return value
Expand All @@ -225,11 +248,7 @@ async def get_state_names(self) -> List[str]:
# TODO: Get all state names from Dapr once implemented.
def append_names_sync():
state_change_tracker = self._get_contextual_state_tracker()
return [
key
for key, value in state_change_tracker.items()
if value.change_kind != StateChangeKind.remove
]
return [key for key, value in state_change_tracker.items() if not _is_absent(value)]

default_loop = asyncio.get_running_loop()
return await default_loop.run_in_executor(None, append_names_sync)
Expand Down Expand Up @@ -267,6 +286,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, cached misses included, 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
Loading
Loading