From 68d1127e4614f1e891a0af46f510fed1ff64801c Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 12 Nov 2025 10:50:00 +0200 Subject: [PATCH] Compute and emit state deltas instead of full state snapshots by default --- .../haiku/rag/graph/agui/cli_renderer.py | 44 +++----- .../haiku/rag/graph/agui/emitter.py | 26 +++-- haiku_rag_slim/haiku/rag/graph/agui/stream.py | 4 +- tests/graph/agui/test_cli_renderer.py | 101 +++++++++++------- tests/graph/agui/test_emitter.py | 36 ++++++- 5 files changed, 135 insertions(+), 76 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py b/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py index 52e223dc..03a60bb4 100644 --- a/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py @@ -22,7 +22,6 @@ class AGUIConsoleRenderer: console: Optional Rich console instance (creates new one if not provided) """ self.console = console or Console() - self._state: dict | None = None async def render(self, events: AsyncIterator[AGUIEvent]) -> Any | None: """Process events and render to console, return final result. @@ -58,11 +57,9 @@ class AGUIConsoleRenderer: elif event_type == "TEXT_MESSAGE_END": pass # End of streaming message, no output needed elif event_type == "STATE_SNAPSHOT": - new_state = event.get("snapshot") - self._render_state_snapshot(new_state) - self._state = new_state + self._render_state_snapshot(event) elif event_type == "STATE_DELTA": - self._apply_state_delta(event) + self._render_state_delta(event) elif event_type == "ACTIVITY_SNAPSHOT": self._render_activity(event) elif event_type == "ACTIVITY_DELTA": @@ -109,33 +106,20 @@ class AGUIConsoleRenderer: if content: self.console.print(f"[yellow][ACTIVITY][/yellow] {content}") - def _render_state_snapshot(self, new_state: dict | None) -> None: - """Render state snapshot showing only what changed.""" - if not new_state: - return - - old_state = self._state or {} - diff = self._compute_diff(old_state, new_state) - - if not diff: + def _render_state_snapshot(self, event: AGUIEvent) -> None: + """Render full state snapshot.""" + snapshot = event.get("snapshot") + if not snapshot: return self.console.print("[blue][STATE_SNAPSHOT][/blue]") - self.console.print(diff, style="dim") + self.console.print(snapshot, style="dim") - def _compute_diff(self, old: dict, new: dict) -> dict: - """Compute difference between old and new state.""" - diff = {} - for key, new_value in new.items(): - old_value = old.get(key) - if old_value != new_value: - if isinstance(new_value, dict) and isinstance(old_value, dict): - nested_diff = self._compute_diff(old_value, new_value) - if nested_diff: - diff[key] = nested_diff - else: - diff[key] = new_value - return diff + def _render_state_delta(self, event: AGUIEvent) -> None: + """Render state delta operations.""" + delta = event.get("delta", []) + if not delta: + return - def _apply_state_delta(self, _event: AGUIEvent) -> None: - """Apply state delta to current state.""" + self.console.print("[blue][STATE_DELTA][/blue]") + self.console.print(delta, style="dim") diff --git a/haiku_rag_slim/haiku/rag/graph/agui/emitter.py b/haiku_rag_slim/haiku/rag/graph/agui/emitter.py index 2c284d46..09201559 100644 --- a/haiku_rag_slim/haiku/rag/graph/agui/emitter.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/emitter.py @@ -13,6 +13,7 @@ from haiku.rag.graph.agui.events import ( emit_run_error, emit_run_finished, emit_run_started, + emit_state_delta, emit_state_snapshot, emit_step_finished, emit_step_started, @@ -35,12 +36,18 @@ class AGUIEmitter[StateT: BaseModel, ResultT]: ResultT: The result type returned by the graph """ - def __init__(self, thread_id: str | None = None, run_id: str | None = None): + def __init__( + self, + thread_id: str | None = None, + run_id: str | None = None, + use_deltas: bool = True, + ): """Initialize the emitter. Args: thread_id: Optional thread ID (generated from input hash if not provided) run_id: Optional run ID (random UUID if not provided) + use_deltas: Whether to emit state deltas instead of full snapshots (default: True) """ self._queue: asyncio.Queue[AGUIEvent | None] = asyncio.Queue() self._closed = False @@ -48,6 +55,7 @@ class AGUIEmitter[StateT: BaseModel, ResultT]: self._run_id = run_id or str(uuid4()) self._last_state: StateT | None = None self._current_step: str | None = None + self._use_deltas = use_deltas @property def thread_id(self) -> str: @@ -73,7 +81,8 @@ class AGUIEmitter[StateT: BaseModel, ResultT]: # RunStarted (state snapshot follows immediately with full state) self._emit(emit_run_started(self._thread_id, self._run_id)) self._emit(emit_state_snapshot(initial_state)) - self._last_state = initial_state + # Store a deep copy to detect future changes + self._last_state = initial_state.model_copy(deep=True) def start_step(self, step_name: str) -> None: """Emit StepStarted event. @@ -100,14 +109,19 @@ class AGUIEmitter[StateT: BaseModel, ResultT]: self._emit(emit_text_message(message, role)) def update_state(self, new_state: StateT) -> None: - """Emit StateSnapshot for state change. + """Emit StateDelta or StateSnapshot for state change. Args: new_state: The updated state """ - # Always emit full snapshot (not delta) for complete state visibility - self._emit(emit_state_snapshot(new_state)) - self._last_state = new_state + if self._use_deltas and self._last_state is not None: + # Emit delta for incremental updates + self._emit(emit_state_delta(self._last_state, new_state)) + else: + # Emit full snapshot for initial state or when deltas disabled + self._emit(emit_state_snapshot(new_state)) + # Store a deep copy to detect future changes + self._last_state = new_state.model_copy(deep=True) def update_activity( self, activity_type: str, content: str, message_id: str | None = None diff --git a/haiku_rag_slim/haiku/rag/graph/agui/stream.py b/haiku_rag_slim/haiku/rag/graph/agui/stream.py index 75f6a3e5..ee6f1ac5 100644 --- a/haiku_rag_slim/haiku/rag/graph/agui/stream.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/stream.py @@ -21,6 +21,7 @@ async def stream_graph( graph: Any, state: BaseModel, deps: GraphDeps, + use_deltas: bool = True, ) -> AsyncIterator[AGUIEvent]: """Run a graph and yield AG-UI events as they occur. @@ -34,6 +35,7 @@ async def stream_graph( graph: The pydantic-graph Graph to execute state: Initial state (Pydantic BaseModel) deps: Graph dependencies with agui_emitter support + use_deltas: Whether to emit state deltas instead of full snapshots (default: True) Yields: AG-UI event dictionaries @@ -46,7 +48,7 @@ async def stream_graph( raise TypeError("deps must have an 'agui_emitter' attribute") # Create AG-UI emitter - emitter: AGUIEmitter[Any, Any] = AGUIEmitter() + emitter: AGUIEmitter[Any, Any] = AGUIEmitter(use_deltas=use_deltas) deps.agui_emitter = emitter async def _execute() -> None: diff --git a/tests/graph/agui/test_cli_renderer.py b/tests/graph/agui/test_cli_renderer.py index 738b0bcc..2f21e904 100644 --- a/tests/graph/agui/test_cli_renderer.py +++ b/tests/graph/agui/test_cli_renderer.py @@ -10,6 +10,7 @@ from haiku.rag.graph.agui.events import ( emit_run_error, emit_run_finished, emit_run_started, + emit_state_delta, emit_state_snapshot, emit_step_finished, emit_step_started, @@ -49,32 +50,18 @@ async def test_renderer_basic_flow(): @pytest.mark.asyncio -async def test_renderer_state_diff(): - """Test that renderer computes state diffs correctly.""" +async def test_renderer_multiple_snapshots(): + """Test that renderer handles multiple snapshots without errors.""" renderer = AGUIConsoleRenderer() - # First state - old_state = {"value": 1, "text": "old", "nested": {"a": 1, "b": 2}} - # Second state with changes - new_state = {"value": 2, "text": "old", "nested": {"a": 1, "b": 3, "c": 4}} + events = [ + emit_state_snapshot(SimpleState(value=1)), + emit_state_snapshot(SimpleState(value=2)), + ] - diff = renderer._compute_diff(old_state, new_state) - - assert diff == { - "value": 2, - "nested": {"b": 3, "c": 4}, # Only changed/new fields in nested - } - - -@pytest.mark.asyncio -async def test_renderer_state_diff_no_changes(): - """Test that no diff is computed when states are equal.""" - renderer = AGUIConsoleRenderer() - - state = {"value": 1, "text": "test"} - diff = renderer._compute_diff(state, state) - - assert diff == {} + # Should render both snapshots without error + result = await renderer.render(async_gen(events)) + assert result is None # No run finished event @pytest.mark.asyncio @@ -99,19 +86,18 @@ async def test_renderer_handles_all_event_types(): @pytest.mark.asyncio -async def test_renderer_initial_state(): - """Test that initial state is rendered.""" +async def test_renderer_state_snapshots(): + """Test that state snapshots are rendered.""" renderer = AGUIConsoleRenderer() events = [ emit_state_snapshot(SimpleState(value=1)), emit_state_snapshot(SimpleState(value=2)), + emit_run_finished("t1", "r1", {"done": True}), ] - await renderer.render(async_gen(events)) - - # After processing, internal state should be the last state - assert renderer._state == {"value": 2} + result = await renderer.render(async_gen(events)) + assert result == {"done": True} @pytest.mark.asyncio @@ -129,18 +115,23 @@ async def test_renderer_no_result(): @pytest.mark.asyncio -async def test_renderer_nested_state_diff(): - """Test nested state diff computation.""" +async def test_renderer_snapshot_then_delta(): + """Test that renderer handles snapshot followed by deltas.""" renderer = AGUIConsoleRenderer() - old = {"level1": {"level2": {"value": 1, "text": "old"}, "other": "same"}} + state1 = SimpleState(value=1) + state2 = SimpleState(value=2) + state3 = SimpleState(value=3) - new = {"level1": {"level2": {"value": 2, "text": "old"}, "other": "same"}} + events = [ + emit_state_snapshot(state1), # Initial snapshot + emit_state_delta(state1, state2), # Delta to value=2 + emit_state_delta(state2, state3), # Delta to value=3 + emit_run_finished("t1", "r1", {"complete": True}), + ] - diff = renderer._compute_diff(old, new) - - # Should only show the changed nested value - assert diff == {"level1": {"level2": {"value": 2}}} + result = await renderer.render(async_gen(events)) + assert result == {"complete": True} @pytest.mark.asyncio @@ -156,3 +147,39 @@ async def test_renderer_with_empty_state(): result = await renderer.render(async_gen(events)) assert result == {"result": "ok"} + + +@pytest.mark.asyncio +async def test_renderer_state_delta(): + """Test that renderer renders state deltas.""" + renderer = AGUIConsoleRenderer() + + state1 = SimpleState(value=1) + state2 = SimpleState(value=2) + + events = [ + emit_state_snapshot(state1), # Initial state + emit_state_delta(state1, state2), # Delta update + emit_run_finished("t1", "r1", None), + ] + + result = await renderer.render(async_gen(events)) + assert result is None + + +@pytest.mark.asyncio +async def test_renderer_state_delta_without_initial(): + """Test that renderer handles delta without initial state gracefully.""" + renderer = AGUIConsoleRenderer() + + state1 = SimpleState(value=1) + state2 = SimpleState(value=2) + + # Send delta without initial snapshot - should still render it + events = [ + emit_state_delta(state1, state2), + emit_run_finished("t1", "r1", {"ok": True}), + ] + + result = await renderer.render(async_gen(events)) + assert result == {"ok": True} diff --git a/tests/graph/agui/test_emitter.py b/tests/graph/agui/test_emitter.py index 5e6148d1..51c89488 100644 --- a/tests/graph/agui/test_emitter.py +++ b/tests/graph/agui/test_emitter.py @@ -67,8 +67,8 @@ async def test_emitter_step_events(): @pytest.mark.asyncio async def test_emitter_state_updates(): - """Test state update events.""" - emitter: AGUIEmitter[TestState, TestResult] = AGUIEmitter() + """Test state update events with snapshots.""" + emitter: AGUIEmitter[TestState, TestResult] = AGUIEmitter(use_deltas=False) state1 = TestState(value=1, text="first") state2 = TestState(value=2, text="second") @@ -87,6 +87,38 @@ async def test_emitter_state_updates(): assert state_events[1]["snapshot"] == {"value": 2, "text": "second"} +@pytest.mark.asyncio +async def test_emitter_state_deltas(): + """Test state update events with deltas.""" + emitter: AGUIEmitter[TestState, TestResult] = AGUIEmitter(use_deltas=True) + + state1 = TestState(value=1, text="first") + state2 = TestState(value=2, text="second") + + emitter.update_state(state1) + emitter.update_state(state2) + await emitter.close() + + events = [] + async for event in emitter: + events.append(event) + + # First update should be a snapshot (no previous state) + snapshot_events = [e for e in events if e["type"] == "STATE_SNAPSHOT"] + assert len(snapshot_events) == 1 + assert snapshot_events[0]["snapshot"] == {"value": 1, "text": "first"} + + # Second update should be a delta + delta_events = [e for e in events if e["type"] == "STATE_DELTA"] + assert len(delta_events) == 1 + # Delta should contain replace operations for changed fields + delta = delta_events[0]["delta"] + assert isinstance(delta, list) + assert len(delta) == 2 # Two fields changed + assert any(op["path"] == "/value" and op["value"] == 2 for op in delta) + assert any(op["path"] == "/text" and op["value"] == "second" for op in delta) + + @pytest.mark.asyncio async def test_emitter_activity_events(): """Test activity events."""