Compute and emit state deltas instead of full state snapshots by default
This commit is contained in:
parent
657245f686
commit
68d1127e46
5 changed files with 135 additions and 76 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Reference in a new issue