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)
|
console: Optional Rich console instance (creates new one if not provided)
|
||||||
"""
|
"""
|
||||||
self.console = console or Console()
|
self.console = console or Console()
|
||||||
self._state: dict | None = None
|
|
||||||
|
|
||||||
async def render(self, events: AsyncIterator[AGUIEvent]) -> Any | None:
|
async def render(self, events: AsyncIterator[AGUIEvent]) -> Any | None:
|
||||||
"""Process events and render to console, return final result.
|
"""Process events and render to console, return final result.
|
||||||
|
|
@ -58,11 +57,9 @@ class AGUIConsoleRenderer:
|
||||||
elif event_type == "TEXT_MESSAGE_END":
|
elif event_type == "TEXT_MESSAGE_END":
|
||||||
pass # End of streaming message, no output needed
|
pass # End of streaming message, no output needed
|
||||||
elif event_type == "STATE_SNAPSHOT":
|
elif event_type == "STATE_SNAPSHOT":
|
||||||
new_state = event.get("snapshot")
|
self._render_state_snapshot(event)
|
||||||
self._render_state_snapshot(new_state)
|
|
||||||
self._state = new_state
|
|
||||||
elif event_type == "STATE_DELTA":
|
elif event_type == "STATE_DELTA":
|
||||||
self._apply_state_delta(event)
|
self._render_state_delta(event)
|
||||||
elif event_type == "ACTIVITY_SNAPSHOT":
|
elif event_type == "ACTIVITY_SNAPSHOT":
|
||||||
self._render_activity(event)
|
self._render_activity(event)
|
||||||
elif event_type == "ACTIVITY_DELTA":
|
elif event_type == "ACTIVITY_DELTA":
|
||||||
|
|
@ -109,33 +106,20 @@ class AGUIConsoleRenderer:
|
||||||
if content:
|
if content:
|
||||||
self.console.print(f"[yellow][ACTIVITY][/yellow] {content}")
|
self.console.print(f"[yellow][ACTIVITY][/yellow] {content}")
|
||||||
|
|
||||||
def _render_state_snapshot(self, new_state: dict | None) -> None:
|
def _render_state_snapshot(self, event: AGUIEvent) -> None:
|
||||||
"""Render state snapshot showing only what changed."""
|
"""Render full state snapshot."""
|
||||||
if not new_state:
|
snapshot = event.get("snapshot")
|
||||||
return
|
if not snapshot:
|
||||||
|
|
||||||
old_state = self._state or {}
|
|
||||||
diff = self._compute_diff(old_state, new_state)
|
|
||||||
|
|
||||||
if not diff:
|
|
||||||
return
|
return
|
||||||
|
|
||||||
self.console.print("[blue][STATE_SNAPSHOT][/blue]")
|
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:
|
def _render_state_delta(self, event: AGUIEvent) -> None:
|
||||||
"""Compute difference between old and new state."""
|
"""Render state delta operations."""
|
||||||
diff = {}
|
delta = event.get("delta", [])
|
||||||
for key, new_value in new.items():
|
if not delta:
|
||||||
old_value = old.get(key)
|
return
|
||||||
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 _apply_state_delta(self, _event: AGUIEvent) -> None:
|
self.console.print("[blue][STATE_DELTA][/blue]")
|
||||||
"""Apply state delta to current state."""
|
self.console.print(delta, style="dim")
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ from haiku.rag.graph.agui.events import (
|
||||||
emit_run_error,
|
emit_run_error,
|
||||||
emit_run_finished,
|
emit_run_finished,
|
||||||
emit_run_started,
|
emit_run_started,
|
||||||
|
emit_state_delta,
|
||||||
emit_state_snapshot,
|
emit_state_snapshot,
|
||||||
emit_step_finished,
|
emit_step_finished,
|
||||||
emit_step_started,
|
emit_step_started,
|
||||||
|
|
@ -35,12 +36,18 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
|
||||||
ResultT: The result type returned by the graph
|
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.
|
"""Initialize the emitter.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
thread_id: Optional thread ID (generated from input hash if not provided)
|
thread_id: Optional thread ID (generated from input hash if not provided)
|
||||||
run_id: Optional run ID (random UUID 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._queue: asyncio.Queue[AGUIEvent | None] = asyncio.Queue()
|
||||||
self._closed = False
|
self._closed = False
|
||||||
|
|
@ -48,6 +55,7 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
|
||||||
self._run_id = run_id or str(uuid4())
|
self._run_id = run_id or str(uuid4())
|
||||||
self._last_state: StateT | None = None
|
self._last_state: StateT | None = None
|
||||||
self._current_step: str | None = None
|
self._current_step: str | None = None
|
||||||
|
self._use_deltas = use_deltas
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def thread_id(self) -> str:
|
def thread_id(self) -> str:
|
||||||
|
|
@ -73,7 +81,8 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
|
||||||
# RunStarted (state snapshot follows immediately with full state)
|
# RunStarted (state snapshot follows immediately with full state)
|
||||||
self._emit(emit_run_started(self._thread_id, self._run_id))
|
self._emit(emit_run_started(self._thread_id, self._run_id))
|
||||||
self._emit(emit_state_snapshot(initial_state))
|
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:
|
def start_step(self, step_name: str) -> None:
|
||||||
"""Emit StepStarted event.
|
"""Emit StepStarted event.
|
||||||
|
|
@ -100,14 +109,19 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
|
||||||
self._emit(emit_text_message(message, role))
|
self._emit(emit_text_message(message, role))
|
||||||
|
|
||||||
def update_state(self, new_state: StateT) -> None:
|
def update_state(self, new_state: StateT) -> None:
|
||||||
"""Emit StateSnapshot for state change.
|
"""Emit StateDelta or StateSnapshot for state change.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
new_state: The updated state
|
new_state: The updated state
|
||||||
"""
|
"""
|
||||||
# Always emit full snapshot (not delta) for complete state visibility
|
if self._use_deltas and self._last_state is not None:
|
||||||
self._emit(emit_state_snapshot(new_state))
|
# Emit delta for incremental updates
|
||||||
self._last_state = new_state
|
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(
|
def update_activity(
|
||||||
self, activity_type: str, content: str, message_id: str | None = None
|
self, activity_type: str, content: str, message_id: str | None = None
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,7 @@ async def stream_graph(
|
||||||
graph: Any,
|
graph: Any,
|
||||||
state: BaseModel,
|
state: BaseModel,
|
||||||
deps: GraphDeps,
|
deps: GraphDeps,
|
||||||
|
use_deltas: bool = True,
|
||||||
) -> AsyncIterator[AGUIEvent]:
|
) -> AsyncIterator[AGUIEvent]:
|
||||||
"""Run a graph and yield AG-UI events as they occur.
|
"""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
|
graph: The pydantic-graph Graph to execute
|
||||||
state: Initial state (Pydantic BaseModel)
|
state: Initial state (Pydantic BaseModel)
|
||||||
deps: Graph dependencies with agui_emitter support
|
deps: Graph dependencies with agui_emitter support
|
||||||
|
use_deltas: Whether to emit state deltas instead of full snapshots (default: True)
|
||||||
|
|
||||||
Yields:
|
Yields:
|
||||||
AG-UI event dictionaries
|
AG-UI event dictionaries
|
||||||
|
|
@ -46,7 +48,7 @@ async def stream_graph(
|
||||||
raise TypeError("deps must have an 'agui_emitter' attribute")
|
raise TypeError("deps must have an 'agui_emitter' attribute")
|
||||||
|
|
||||||
# Create AG-UI emitter
|
# Create AG-UI emitter
|
||||||
emitter: AGUIEmitter[Any, Any] = AGUIEmitter()
|
emitter: AGUIEmitter[Any, Any] = AGUIEmitter(use_deltas=use_deltas)
|
||||||
deps.agui_emitter = emitter
|
deps.agui_emitter = emitter
|
||||||
|
|
||||||
async def _execute() -> None:
|
async def _execute() -> None:
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ from haiku.rag.graph.agui.events import (
|
||||||
emit_run_error,
|
emit_run_error,
|
||||||
emit_run_finished,
|
emit_run_finished,
|
||||||
emit_run_started,
|
emit_run_started,
|
||||||
|
emit_state_delta,
|
||||||
emit_state_snapshot,
|
emit_state_snapshot,
|
||||||
emit_step_finished,
|
emit_step_finished,
|
||||||
emit_step_started,
|
emit_step_started,
|
||||||
|
|
@ -49,32 +50,18 @@ async def test_renderer_basic_flow():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_renderer_state_diff():
|
async def test_renderer_multiple_snapshots():
|
||||||
"""Test that renderer computes state diffs correctly."""
|
"""Test that renderer handles multiple snapshots without errors."""
|
||||||
renderer = AGUIConsoleRenderer()
|
renderer = AGUIConsoleRenderer()
|
||||||
|
|
||||||
# First state
|
events = [
|
||||||
old_state = {"value": 1, "text": "old", "nested": {"a": 1, "b": 2}}
|
emit_state_snapshot(SimpleState(value=1)),
|
||||||
# Second state with changes
|
emit_state_snapshot(SimpleState(value=2)),
|
||||||
new_state = {"value": 2, "text": "old", "nested": {"a": 1, "b": 3, "c": 4}}
|
]
|
||||||
|
|
||||||
diff = renderer._compute_diff(old_state, new_state)
|
# Should render both snapshots without error
|
||||||
|
result = await renderer.render(async_gen(events))
|
||||||
assert diff == {
|
assert result is None # No run finished event
|
||||||
"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 == {}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -99,19 +86,18 @@ async def test_renderer_handles_all_event_types():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_renderer_initial_state():
|
async def test_renderer_state_snapshots():
|
||||||
"""Test that initial state is rendered."""
|
"""Test that state snapshots are rendered."""
|
||||||
renderer = AGUIConsoleRenderer()
|
renderer = AGUIConsoleRenderer()
|
||||||
|
|
||||||
events = [
|
events = [
|
||||||
emit_state_snapshot(SimpleState(value=1)),
|
emit_state_snapshot(SimpleState(value=1)),
|
||||||
emit_state_snapshot(SimpleState(value=2)),
|
emit_state_snapshot(SimpleState(value=2)),
|
||||||
|
emit_run_finished("t1", "r1", {"done": True}),
|
||||||
]
|
]
|
||||||
|
|
||||||
await renderer.render(async_gen(events))
|
result = await renderer.render(async_gen(events))
|
||||||
|
assert result == {"done": True}
|
||||||
# After processing, internal state should be the last state
|
|
||||||
assert renderer._state == {"value": 2}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -129,18 +115,23 @@ async def test_renderer_no_result():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_renderer_nested_state_diff():
|
async def test_renderer_snapshot_then_delta():
|
||||||
"""Test nested state diff computation."""
|
"""Test that renderer handles snapshot followed by deltas."""
|
||||||
renderer = AGUIConsoleRenderer()
|
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)
|
result = await renderer.render(async_gen(events))
|
||||||
|
assert result == {"complete": True}
|
||||||
# Should only show the changed nested value
|
|
||||||
assert diff == {"level1": {"level2": {"value": 2}}}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -156,3 +147,39 @@ async def test_renderer_with_empty_state():
|
||||||
|
|
||||||
result = await renderer.render(async_gen(events))
|
result = await renderer.render(async_gen(events))
|
||||||
assert result == {"result": "ok"}
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_emitter_state_updates():
|
async def test_emitter_state_updates():
|
||||||
"""Test state update events."""
|
"""Test state update events with snapshots."""
|
||||||
emitter: AGUIEmitter[TestState, TestResult] = AGUIEmitter()
|
emitter: AGUIEmitter[TestState, TestResult] = AGUIEmitter(use_deltas=False)
|
||||||
|
|
||||||
state1 = TestState(value=1, text="first")
|
state1 = TestState(value=1, text="first")
|
||||||
state2 = TestState(value=2, text="second")
|
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"}
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_emitter_activity_events():
|
async def test_emitter_activity_events():
|
||||||
"""Test activity events."""
|
"""Test activity events."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue