Compute and emit state deltas instead of full state snapshots by default

This commit is contained in:
Yiorgis Gozadinos 2025-11-12 10:50:00 +02:00
parent 657245f686
commit 68d1127e46
No known key found for this signature in database
5 changed files with 135 additions and 76 deletions

View file

@ -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")

View file

@ -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

View file

@ -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:

View file

@ -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}

View file

@ -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."""