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
69 changes: 60 additions & 9 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 Down Expand Up @@ -153,7 +180,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 +199,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 +226,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 +254,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 +292,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