diff --git a/reflex/istate/manager/disk.py b/reflex/istate/manager/disk.py index 15dcc53f2ed..9272a98cd5f 100644 --- a/reflex/istate/manager/disk.py +++ b/reflex/istate/manager/disk.py @@ -340,6 +340,15 @@ async def set_state( state=state, timestamp=time.time(), ) + else: + # A write for this token is already queued; replace the queued + # state with the latest value so a stale snapshot is not + # flushed to disk. Preserve the original timestamp so the + # item still flushes at its originally scheduled time. + self._write_queue[token] = dataclasses.replace( + self._write_queue[token], + state=state, + ) else: # Immediate write to disk. await self.set_state_for_substate(token, state) diff --git a/tests/units/test_state.py b/tests/units/test_state.py index 7106287a43d..1fe0875f3da 100644 --- a/tests/units/test_state.py +++ b/tests/units/test_state.py @@ -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, @@ -4431,6 +4431,34 @@ 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_set_state_updates_queued_write(tmp_path, token): + """Test that a second set_state call for an already-queued token replaces + the queued state, so the latest value is what gets flushed to disk. + + Args: + tmp_path: A temporary directory (pytest fixture). + token: A token. + """ + state_manager = StateManagerDisk() + object.__setattr__(state_manager, "states_directory", tmp_path) + state_manager._write_debounce_seconds = 60 + + state_token = StateToken(ident=token, cls=int) + + await state_manager.set_state(state_token, 1) + await state_manager.set_state(state_token, 2) + + assert state_manager._write_queue[state_token].state == 2 + + await state_manager.close() + + reloaded_state_manager = StateManagerDisk() + object.__setattr__(reloaded_state_manager, "states_directory", tmp_path) + assert await reloaded_state_manager.load_state(state_token) == 2 + await reloaded_state_manager.close() + + class Obj(Base): """A object containing a callable for testing fallback pickle."""