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
1 change: 1 addition & 0 deletions news/6807.bugfix.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Fixed debounced disk state writes so non-`BaseState` values stay cache-coherent and flush the latest queued value.
17 changes: 11 additions & 6 deletions reflex/istate/manager/disk.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,17 +332,22 @@ async def set_state(
context: The state modification context.
"""
token = self._coerce_token(token)
is_base_state_token = isinstance(token, BaseStateToken)
if self._write_debounce_seconds > 0:
# Deferred write to reduce disk IO overhead.
if token not in self._write_queue:
self._write_queue[token] = QueueItem(
token=token,
state=state,
timestamp=time.time(),
)
if not is_base_state_token:
self.states[token.cache_key] = state
queued_item = self._write_queue.get(token)
self._write_queue[token] = QueueItem[TOKEN_TYPE](
token=token,
state=state,
timestamp=queued_item.timestamp if queued_item else time.time(),
)
else:
# Immediate write to disk.
await self.set_state_for_substate(token, state)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
if not is_base_state_token:
self.states[token.cache_key] = state
# Ensure the processing task is scheduled to handle expirations and any deferred writes.
await self._schedule_process_write_queue()

Expand Down
65 changes: 64 additions & 1 deletion tests/units/test_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@
from reflex.istate.manager.disk import StateManagerDisk
from reflex.istate.manager.memory import StateManagerMemory
from reflex.istate.manager.redis import StateManagerRedis
from reflex.istate.manager.token import BaseStateToken
from reflex.istate.manager.token import BaseStateToken, StateToken
from reflex.istate.proxy import StateProxy
from reflex.state import (
BaseState,
Expand Down Expand Up @@ -4432,6 +4432,69 @@ async def test_state_manager_disk_close_resets_write_queue_task():
assert state_manager._write_queue_task is None


@pytest.mark.asyncio
async def test_state_manager_disk_debounced_set_state_updates_non_base_state_cache(
tmp_path, monkeypatch
):
"""Test that debounced non-BaseState writes are visible before disk flush."""
monkeypatch.setattr(prerequisites, "get_states_dir", lambda: tmp_path)
state_manager = StateManagerDisk(_write_debounce_seconds=60)
token = StateToken(ident="client", cls=int)

await state_manager.set_state(token, 1)

assert await state_manager.get_state(token) == 1

await state_manager.close()


@pytest.mark.asyncio
async def test_state_manager_disk_debounced_set_state_flushes_latest_non_base_state(
tmp_path, monkeypatch
):
"""Test that debounced non-BaseState writes flush the latest queued value."""
monkeypatch.setattr(prerequisites, "get_states_dir", lambda: tmp_path)
state_manager = StateManagerDisk(_write_debounce_seconds=60)
token = StateToken(ident="client", cls=int)

await state_manager.set_state(token, 1)
first_timestamp = state_manager._write_queue[token].timestamp
await state_manager.set_state(token, 2)

assert state_manager._write_queue[token].timestamp == first_timestamp

await state_manager.close()

fresh_state_manager = StateManagerDisk(_write_debounce_seconds=0)
assert await fresh_state_manager.get_state(token) == 2

await fresh_state_manager.close()


@pytest.mark.asyncio
async def test_state_manager_disk_immediate_set_state_failure_keeps_previous_cache(
tmp_path, monkeypatch, mocker
):
"""Test failed immediate writes do not cache an unpersisted value."""
monkeypatch.setattr(prerequisites, "get_states_dir", lambda: tmp_path)
state_manager = StateManagerDisk(_write_debounce_seconds=0)
token = StateToken(ident="client", cls=int)

await state_manager.set_state(token, 1)
mocker.patch.object(
state_manager,
"set_state_for_substate",
side_effect=RuntimeError("write failed"),
)

with pytest.raises(RuntimeError, match="write failed"):
await state_manager.set_state(token, 2)

assert await state_manager.get_state(token) == 1

await state_manager.close()


class Obj(Base):
"""A object containing a callable for testing fallback pickle."""

Expand Down
Loading