Always use state deltas instead of conditional snapshot/delta logic

This commit is contained in:
Yiorgis Gozadinos 2026-01-28 16:50:58 +02:00
parent 57fd08d600
commit cbd94f6cc9
No known key found for this signature in database
6 changed files with 82 additions and 134 deletions

View file

@ -79,17 +79,12 @@ async def stream_chat(request: Request) -> Response:
accept = request.headers.get("accept", SSE_CONTENT_TYPE) accept = request.headers.get("accept", SSE_CONTENT_TYPE)
run_input = AGUIAdapter.build_run_input(body) run_input = AGUIAdapter.build_run_input(body)
# Restore session state from incoming AG-UI state (look under namespaced key) # Restore session state from incoming AG-UI state
session_state: ChatSessionState | None = None session_state = ChatSessionState(session_id="") # New session: empty session_id
state = getattr(run_input, "state", None) state = getattr(run_input, "state", None)
if state and AGUI_STATE_KEY in state: if state and AGUI_STATE_KEY in state:
chat_state = state[AGUI_STATE_KEY] chat_state = state[AGUI_STATE_KEY]
if chat_state and chat_state.get("session_id"): if chat_state and chat_state.get("session_id"):
# Only restore state if client has a session_id (not first request)
# This ensures first request gets a full snapshot with generated UUID
# NOTE: We intentionally do NOT restore session_context from the client.
# The server maintains session_context via background summarization tasks,
# and the agent fetches it from the server-side cache (get_cached_session_context).
session_state = ChatSessionState( session_state = ChatSessionState(
session_id=chat_state["session_id"], session_id=chat_state["session_id"],
qa_history=[ qa_history=[
@ -104,8 +99,10 @@ async def stream_chat(request: Request) -> Response:
f"qa_history={len(session_state.qa_history)}, " f"qa_history={len(session_state.qa_history)}, "
f"citations={len(session_state.citation_registry)}" f"citations={len(session_state.citation_registry)}"
) )
else:
logger.info("Incoming state: new session")
else: else:
logger.info("Incoming state: new session (no session_id)") logger.info("Incoming state: new session")
deps = ChatDeps( deps = ChatDeps(
client=get_client(db_path), client=get_client(db_path),

View file

@ -1,5 +1,6 @@
import asyncio import asyncio
import math import math
import uuid
from pydantic_ai import Agent, RunContext, ToolReturn from pydantic_ai import Agent, RunContext, ToolReturn
@ -94,11 +95,9 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
limit: Number of results to return (default: 5) limit: Number of results to return (default: 5)
""" """
# Build session filter from document_filter # Build session filter from document_filter
session_filter = None session_filter = build_multi_document_filter(
if ctx.deps.session_state and ctx.deps.session_state.document_filter: ctx.deps.session_state.document_filter
session_filter = build_multi_document_filter( )
ctx.deps.session_state.document_filter
)
# Build tool filter from document_name parameter # Build tool filter from document_name parameter
tool_filter = build_document_filter(document_name) if document_name else None tool_filter = build_document_filter(document_name) if document_name else None
@ -116,12 +115,9 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
if not results: if not results:
return ToolReturn(return_value="No results found.") return ToolReturn(return_value="No results found.")
# Copy session state to work with (avoids mutating original for delta computation) new_state = ctx.deps.session_state.model_copy(deep=True)
new_state = ( if not new_state.session_id:
ctx.deps.session_state.model_copy(deep=True) new_state.session_id = str(uuid.uuid4())
if ctx.deps.session_state
else ChatSessionState()
)
# Build citation infos using the copy's registry # Build citation infos using the copy's registry
citation_infos = [] citation_infos = []
@ -190,11 +186,9 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
document_name: Optional document name/title to search within (e.g., "tbmed593", "army manual") document_name: Optional document name/title to search within (e.g., "tbmed593", "army manual")
""" """
# Build session filter from document_filter # Build session filter from document_filter
session_filter = None session_filter = build_multi_document_filter(
if ctx.deps.session_state and ctx.deps.session_state.document_filter: ctx.deps.session_state.document_filter
session_filter = build_multi_document_filter( )
ctx.deps.session_state.document_filter
)
# Build tool filter from document_name parameter # Build tool filter from document_name parameter
tool_filter = build_document_filter(document_name) if document_name else None tool_filter = build_document_filter(document_name) if document_name else None
@ -204,23 +198,19 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
# Build and run the conversational research graph # Build and run the conversational research graph
graph = build_conversational_graph(config=ctx.deps.config) graph = build_conversational_graph(config=ctx.deps.config)
session_id = ctx.deps.session_state.session_id if ctx.deps.session_state else "" session_id = ctx.deps.session_state.session_id
# Get session context from server cache for planning, fallback to initial_context # Get session context from server cache for planning, fallback to initial_context
cached_context = get_cached_session_context(session_id) if session_id else None cached_context = get_cached_session_context(session_id)
session_context = ( session_context = (
cached_context.render_markdown() cached_context.render_markdown()
if cached_context and cached_context.summary if cached_context and cached_context.summary
else ( else ctx.deps.session_state.initial_context
ctx.deps.session_state.initial_context
if ctx.deps.session_state
else None
)
) )
# Find relevant prior answers from qa_history # Find relevant prior answers from qa_history
prior_answers = [] prior_answers = []
if ctx.deps.session_state and ctx.deps.session_state.qa_history: if ctx.deps.session_state.qa_history:
embedder = get_embedder(ctx.deps.config) embedder = get_embedder(ctx.deps.config)
question_embedding = await embedder.embed_query(question) question_embedding = await embedder.embed_query(question)
@ -267,12 +257,9 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
result = await graph.run(state=state, deps=deps) result = await graph.run(state=state, deps=deps)
# Copy session state to work with (avoids mutating original for delta computation) new_state = ctx.deps.session_state.model_copy(deep=True)
new_state = ( if not new_state.session_id:
ctx.deps.session_state.model_copy(deep=True) new_state.session_id = str(uuid.uuid4())
if ctx.deps.session_state
else ChatSessionState()
)
# Build citation infos using the copy's registry # Build citation infos using the copy's registry
citation_infos = [] citation_infos = []
@ -354,12 +341,7 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
page_size = 50 page_size = 50
offset = (page - 1) * page_size offset = (page - 1) * page_size
# Build session filter from document_filter doc_filter = build_multi_document_filter(ctx.deps.session_state.document_filter)
doc_filter = None
if ctx.deps.session_state and ctx.deps.session_state.document_filter:
doc_filter = build_multi_document_filter(
ctx.deps.session_state.document_filter
)
docs = await ctx.deps.client.list_documents( docs = await ctx.deps.client.list_documents(
limit=page_size, offset=offset, filter=doc_filter limit=page_size, offset=offset, filter=doc_filter

View file

@ -4,7 +4,7 @@ from datetime import datetime
from typing import Any from typing import Any
import jsonpatch import jsonpatch
from ag_ui.core import EventType, StateDeltaEvent, StateSnapshotEvent from ag_ui.core import EventType, StateDeltaEvent
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from haiku.rag.agents.research.models import Citation, SearchAnswer from haiku.rag.agents.research.models import Citation, SearchAnswer
@ -108,14 +108,14 @@ class ChatDeps:
client: HaikuRAG client: HaikuRAG
config: AppConfig config: AppConfig
search_results: list[SearchResult] | None = None search_results: list[SearchResult] | None = None
session_state: ChatSessionState | None = None session_state: ChatSessionState = field(
default_factory=lambda: ChatSessionState(session_id="")
)
state_key: str | None = None state_key: str | None = None
@property @property
def state(self) -> dict[str, Any] | None: def state(self) -> dict[str, Any]:
"""Get current state for AG-UI protocol.""" """Get current state for AG-UI protocol."""
if self.session_state is None:
return None
snapshot = self.session_state.model_dump() snapshot = self.session_state.model_dump()
if self.state_key: if self.state_key:
return {self.state_key: snapshot} return {self.state_key: snapshot}
@ -133,29 +133,24 @@ class ChatDeps:
if isinstance(nested, dict): if isinstance(nested, dict):
state_data = nested state_data = nested
# Update session_state from incoming state # Update session_state from incoming state
if self.session_state is not None: if "qa_history" in state_data:
if "qa_history" in state_data: self.session_state.qa_history = [
self.session_state.qa_history = [ QAResponse(**qa) if isinstance(qa, dict) else qa
QAResponse(**qa) if isinstance(qa, dict) else qa for qa in state_data.get("qa_history", [])
for qa in state_data.get("qa_history", []) ]
] if "citations" in state_data:
if "citations" in state_data: self.session_state.citations = [
self.session_state.citations = [ Citation(**c) if isinstance(c, dict) else c
Citation(**c) if isinstance(c, dict) else c for c in state_data.get("citations", [])
for c in state_data.get("citations", []) ]
] if state_data.get("session_id"):
if state_data.get("session_id"): self.session_state.session_id = state_data["session_id"]
self.session_state.session_id = state_data["session_id"] if "document_filter" in state_data:
if "document_filter" in state_data: self.session_state.document_filter = state_data.get("document_filter", [])
self.session_state.document_filter = state_data.get( if "citation_registry" in state_data:
"document_filter", [] self.session_state.citation_registry = state_data["citation_registry"]
) if "initial_context" in state_data:
if "citation_registry" in state_data: self.session_state.initial_context = state_data.get("initial_context")
self.session_state.citation_registry = state_data["citation_registry"]
if "initial_context" in state_data:
self.session_state.initial_context = state_data.get("initial_context")
# NOTE: session_context is server-managed; we don't accept it from the client
# to maintain server-side ownership of conversation summarization
@dataclass @dataclass
@ -199,20 +194,14 @@ def combine_filters(filter1: str | None, filter2: str | None) -> str | None:
def emit_state_event( def emit_state_event(
current_state: ChatSessionState | None, current_state: ChatSessionState,
new_state: ChatSessionState, new_state: ChatSessionState,
state_key: str | None = None, state_key: str | None = None,
) -> StateSnapshotEvent | StateDeltaEvent | None: ) -> StateDeltaEvent | None:
"""Emit state delta against current state, or full snapshot if no current state.""" """Emit state delta against current state, or None if no changes."""
new_snapshot = new_state.model_dump(mode="json") new_snapshot = new_state.model_dump(mode="json")
wrapped_new = {state_key: new_snapshot} if state_key else new_snapshot wrapped_new = {state_key: new_snapshot} if state_key else new_snapshot
if current_state is None:
return StateSnapshotEvent(
type=EventType.STATE_SNAPSHOT,
snapshot=wrapped_new,
)
current_snapshot = current_state.model_dump(mode="json") current_snapshot = current_state.model_dump(mode="json")
wrapped_current = {state_key: current_snapshot} if state_key else current_snapshot wrapped_current = {state_key: current_snapshot} if state_key else current_snapshot

View file

@ -100,7 +100,7 @@ class ChatApp(App):
self.client: HaikuRAG | None = None self.client: HaikuRAG | None = None
self.config = get_config() self.config = get_config()
self.agent: Agent[ChatDeps, str] | None = None self.agent: Agent[ChatDeps, str] | None = None
self.session_state: ChatSessionState | None = None self.session_state = ChatSessionState()
self._is_processing = False self._is_processing = False
self._tool_call_widgets: dict[str, Any] = {} self._tool_call_widgets: dict[str, Any] = {}
self._last_citations: list[Citation] = [] self._last_citations: list[Citation] = []

View file

@ -19,7 +19,12 @@ from haiku.rag.config import Config
def extract_state_from_result(result, state_key: str = AGUI_STATE_KEY) -> dict | None: def extract_state_from_result(result, state_key: str = AGUI_STATE_KEY) -> dict | None:
"""Extract emitted state from agent result's tool return metadata.""" """Extract emitted state from agent result's tool return metadata.
For deltas, applies the patch to an empty state to get the final state.
"""
import jsonpatch
for message in result.all_messages(): for message in result.all_messages():
if hasattr(message, "parts"): if hasattr(message, "parts"):
for part in message.parts: for part in message.parts:
@ -28,9 +33,14 @@ def extract_state_from_result(result, state_key: str = AGUI_STATE_KEY) -> dict |
if isinstance(meta, StateSnapshotEvent): if isinstance(meta, StateSnapshotEvent):
return meta.snapshot.get(state_key) return meta.snapshot.get(state_key)
elif isinstance(meta, StateDeltaEvent): elif isinstance(meta, StateDeltaEvent):
# For delta, we'd need to apply the patch # Apply delta to empty state to get final state
# For now, return None and let caller handle empty_state = {
pass state_key: ChatSessionState(session_id="").model_dump(
mode="json"
)
}
patched = jsonpatch.apply_patch(empty_state, meta.delta)
return patched.get(state_key)
return None return None
@ -54,7 +64,9 @@ def test_chat_deps_initialization(temp_db_path):
assert deps.client is client assert deps.client is client
assert deps.config is Config assert deps.config is Config
assert deps.search_results is None assert deps.search_results is None
assert deps.session_state is None assert deps.session_state is not None
assert deps.session_state.qa_history == []
assert deps.session_state.citations == []
client.close() client.close()
@ -467,11 +479,9 @@ async def test_chat_agent_ask_adds_citations(allow_model_requests, temp_db_path)
) )
agent = create_chat_agent(Config) agent = create_chat_agent(Config)
# Pass session_state=None so agent creates fresh state with UUID
deps = ChatDeps( deps = ChatDeps(
client=client, client=client,
config=Config, config=Config,
session_state=None,
state_key=AGUI_STATE_KEY, state_key=AGUI_STATE_KEY,
) )
@ -506,11 +516,9 @@ async def test_chat_agent_ask_triggers_background_summarization(
) )
agent = create_chat_agent(Config) agent = create_chat_agent(Config)
# Pass session_state=None so agent creates fresh state with UUID
deps = ChatDeps( deps = ChatDeps(
client=client, client=client,
config=Config, config=Config,
session_state=None,
state_key=AGUI_STATE_KEY, state_key=AGUI_STATE_KEY,
) )
@ -565,11 +573,9 @@ async def test_chat_agent_ask_with_prior_answer_retrieval(
) )
agent = create_chat_agent(Config) agent = create_chat_agent(Config)
# First call with no session state
deps1 = ChatDeps( deps1 = ChatDeps(
client=client, client=client,
config=Config, config=Config,
session_state=None,
state_key=AGUI_STATE_KEY, state_key=AGUI_STATE_KEY,
) )

View file

@ -146,8 +146,8 @@ def test_chat_deps_state_getter_without_namespace():
assert state["session_id"] == "test-123" assert state["session_id"] == "test-123"
def test_chat_deps_state_getter_returns_none_without_session(): def test_chat_deps_state_getter_returns_default_state():
"""Test ChatDeps.state getter returns None when no session_state.""" """Test ChatDeps.state getter returns default state when not explicitly set."""
from unittest.mock import MagicMock from unittest.mock import MagicMock
from haiku.rag.agents.chat.state import ChatDeps from haiku.rag.agents.chat.state import ChatDeps
@ -158,10 +158,13 @@ def test_chat_deps_state_getter_returns_none_without_session():
deps = ChatDeps( deps = ChatDeps(
client=mock_client, client=mock_client,
config=mock_config, config=mock_config,
session_state=None,
) )
assert deps.state is None state = deps.state
assert state is not None
assert "session_id" in state
assert state["qa_history"] == []
assert state["citations"] == []
def test_chat_deps_state_setter_updates_from_namespaced_state(): def test_chat_deps_state_setter_updates_from_namespaced_state():
@ -223,8 +226,8 @@ def test_chat_deps_state_setter_handles_none():
assert deps.session_state.session_id == "original" assert deps.session_state.session_id == "original"
def test_chat_deps_state_setter_without_session_state(): def test_chat_deps_state_setter_updates_default_state():
"""Test ChatDeps.state setter does nothing when session_state is None.""" """Test ChatDeps.state setter updates the default session_state."""
from unittest.mock import MagicMock from unittest.mock import MagicMock
from haiku.rag.agents.chat.state import ChatDeps from haiku.rag.agents.chat.state import ChatDeps
@ -235,13 +238,15 @@ def test_chat_deps_state_setter_without_session_state():
deps = ChatDeps( deps = ChatDeps(
client=mock_client, client=mock_client,
config=mock_config, config=mock_config,
session_state=None,
) )
# Should not raise even with valid incoming state original_session_id = deps.session_state.session_id
deps.state = {"session_id": "test", "qa_history": [], "citations": []}
assert deps.session_state is None # Update with incoming state
deps.state = {"session_id": "updated-123", "qa_history": [], "citations": []}
assert deps.session_state.session_id == "updated-123"
assert deps.session_state.session_id != original_session_id
def test_chat_deps_state_setter_with_citation_dicts(): def test_chat_deps_state_setter_with_citation_dicts():
@ -645,37 +650,6 @@ def test_chat_session_state_model_dump_json_serializes_datetime():
assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00" assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00"
def test_emit_state_event_returns_snapshot_when_no_current_state():
"""emit_state_event returns StateSnapshotEvent when current_state is None."""
from ag_ui.core import EventType, StateSnapshotEvent
from haiku.rag.agents.chat.state import emit_state_event
new_state = ChatSessionState(session_id="test-123")
event = emit_state_event(None, new_state)
assert isinstance(event, StateSnapshotEvent)
assert event.type == EventType.STATE_SNAPSHOT
assert event.snapshot["session_id"] == "test-123"
def test_emit_state_event_returns_snapshot_with_state_key():
"""emit_state_event wraps snapshot in state_key namespace."""
from ag_ui.core import EventType, StateSnapshotEvent
from haiku.rag.agents.chat.state import AGUI_STATE_KEY, emit_state_event
new_state = ChatSessionState(session_id="test-123")
event = emit_state_event(None, new_state, state_key=AGUI_STATE_KEY)
assert isinstance(event, StateSnapshotEvent)
assert event.type == EventType.STATE_SNAPSHOT
assert AGUI_STATE_KEY in event.snapshot
assert event.snapshot[AGUI_STATE_KEY]["session_id"] == "test-123"
def test_emit_state_event_returns_none_when_no_changes(): def test_emit_state_event_returns_none_when_no_changes():
"""emit_state_event returns None when states are identical.""" """emit_state_event returns None when states are identical."""
from haiku.rag.agents.chat.state import emit_state_event from haiku.rag.agents.chat.state import emit_state_event