From d88f2f003acab9f833c01279168e6a957adcc20d Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 11 Feb 2026 14:31:37 +0200 Subject: [PATCH] Remove incoming_* fields, simplify delta computation --- haiku_rag_slim/haiku/rag/agents/chat/agent.py | 38 ++++++------------- haiku_rag_slim/haiku/rag/agents/chat/state.py | 31 ++++----------- haiku_rag_slim/haiku/rag/tools/qa.py | 7 ---- haiku_rag_slim/haiku/rag/tools/session.py | 3 +- tests/agents/chat/test_chat_agent.py | 10 +---- 5 files changed, 23 insertions(+), 66 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/agents/chat/agent.py b/haiku_rag_slim/haiku/rag/agents/chat/agent.py index ef0aa486..4639aaa2 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/agent.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/agent.py @@ -59,11 +59,7 @@ class ChatDeps: """ session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState) qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState) - snapshot = build_chat_state_snapshot( - session_state, - qa_session_state, - incoming=False, - ) + snapshot = build_chat_state_snapshot(session_state, qa_session_state) if self.state_key: return {self.state_key: snapshot} return snapshot @@ -95,19 +91,15 @@ class ChatDeps: for c in state_data.get("citations", []) ] - # Track what the client sent (for delta computation) - incoming_session_id = state_data.get("session_id", "") - - if incoming_session_id: - self.session_id = incoming_session_id + # Restore session_id from client or generate one + client_session_id = state_data.get("session_id", "") + if client_session_id: + self.session_id = client_session_id elif not self.session_id: - # Generate session_id now so ask() tool can use it self.session_id = str(uuid.uuid4()) - # Sync session_id to SessionState (track incoming for delta computation) if session_state is not None: session_state.session_id = self.session_id - session_state.incoming_session_id = incoming_session_id qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState) if qa_session_state is not None: @@ -119,28 +111,22 @@ class ChatDeps: for qa in state_data.get("qa_history", []) ] - # Track what client sent for delta computation - incoming_session_context = state_data.get("session_context") - if isinstance(incoming_session_context, dict): - qa_session_state.incoming_session_context = SessionContext( - **incoming_session_context - ) - qa_session_state.session_context = ( - qa_session_state.incoming_session_context.summary - ) - elif incoming_session_context is None: - qa_session_state.incoming_session_context = None + # Restore session_context from client + session_context = state_data.get("session_context") + if isinstance(session_context, dict): + qa_session_state.session_context = SessionContext( + **session_context + ).summary + elif session_context is None: qa_session_state.session_context = None # Check cache for fresher session_context from background summarization - # Cache is authoritative so background summaries show up on next request if self.session_id: cached = get_cached_session_context(self.session_id) if cached and cached.summary: qa_session_state.session_context = cached.summary # Handle initial_context -> session_context for first message - # Only applies if session_context is still empty after restoring and cache check if "initial_context" in state_data: initial = state_data.get("initial_context") if initial and not qa_session_state.session_context: diff --git a/haiku_rag_slim/haiku/rag/agents/chat/state.py b/haiku_rag_slim/haiku/rag/agents/chat/state.py index 84d22c28..a777eb6f 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/state.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/state.py @@ -46,29 +46,22 @@ def _rebuild_models(qa_history_entry_cls: type) -> None: def build_chat_state_snapshot( session_state: "SessionState | None", qa_state: "QASessionState | None", - *, - incoming: bool = False, ) -> dict[str, Any]: - """Build a combined AG-UI chat state snapshot. + """Build a combined AG-UI chat state snapshot from current values. Args: session_state: SessionState from ToolContext. qa_state: QASessionState from ToolContext. - incoming: If True, use client-sent values where applicable. Returns: - Snapshot dict, optionally wrapped by state_key. + Snapshot dict. """ snapshot: dict[str, Any] = {"session_id": ""} if session_state is not None: snapshot.update( { - "session_id": ( - session_state.incoming_session_id - if incoming - else session_state.session_id - ), + "session_id": session_state.session_id, "document_filter": session_state.document_filter.copy(), "citation_registry": session_state.citation_registry.copy(), "citations": [c.model_dump() for c in session_state.citations], @@ -77,20 +70,12 @@ def build_chat_state_snapshot( if qa_state is not None: snapshot["qa_history"] = [qa.model_dump() for qa in qa_state.qa_history] - if incoming: - if qa_state.incoming_session_context is not None: - snapshot["session_context"] = ( - qa_state.incoming_session_context.model_dump(mode="json") - ) - else: - snapshot["session_context"] = None + if qa_state.session_context: + snapshot["session_context"] = SessionContext( + summary=qa_state.session_context + ).model_dump(mode="json") else: - if qa_state.session_context: - snapshot["session_context"] = SessionContext( - summary=qa_state.session_context - ).model_dump(mode="json") - else: - snapshot["session_context"] = None + snapshot["session_context"] = None return snapshot diff --git a/haiku_rag_slim/haiku/rag/tools/qa.py b/haiku_rag_slim/haiku/rag/tools/qa.py index 20a77cbd..4882323e 100644 --- a/haiku_rag_slim/haiku/rag/tools/qa.py +++ b/haiku_rag_slim/haiku/rag/tools/qa.py @@ -9,7 +9,6 @@ from haiku.rag.agents.chat.context import ( trigger_background_summarization, ) from haiku.rag.agents.chat.state import ( - SessionContext, build_chat_state_delta, build_chat_state_snapshot, ) @@ -83,9 +82,6 @@ class QASessionState(BaseModel): qa_history: list[QAHistoryEntry] = [] session_context: str | None = None - incoming_session_context: SessionContext | None = Field( - default=None, exclude=True - ) # Track what client sent QA_SESSION_NAMESPACE = "haiku.rag.qa_session" @@ -288,12 +284,10 @@ def create_qa_toolset( qa_session_state = context.get(QA_SESSION_NAMESPACE, QASessionState) state_key = context.state_key - # Use incoming values (what client sent) so delta shows server-side updates if session_state is not None: old_state_snapshot = build_chat_state_snapshot( session_state, qa_session_state, - incoming=True, ) qa_result = await run_qa_core( @@ -311,7 +305,6 @@ def create_qa_toolset( new_state_snapshot = build_chat_state_snapshot( session_state, qa_session_state, - incoming=False, ) state_event = build_chat_state_delta( diff --git a/haiku_rag_slim/haiku/rag/tools/session.py b/haiku_rag_slim/haiku/rag/tools/session.py index 46ea9f01..1f84837e 100644 --- a/haiku_rag_slim/haiku/rag/tools/session.py +++ b/haiku_rag_slim/haiku/rag/tools/session.py @@ -2,7 +2,7 @@ from typing import Any import jsonpatch from ag_ui.core import EventType, StateDeltaEvent -from pydantic import BaseModel, Field +from pydantic import BaseModel from haiku.rag.agents.research.models import Citation @@ -20,7 +20,6 @@ class SessionState(BaseModel): """ session_id: str = "" - incoming_session_id: str = Field(default="", exclude=True) # Track what client sent document_filter: list[str] = [] citation_registry: dict[str, int] = {} citations: list[Citation] = [] diff --git a/tests/agents/chat/test_chat_agent.py b/tests/agents/chat/test_chat_agent.py index 25b06a1c..22964690 100644 --- a/tests/agents/chat/test_chat_agent.py +++ b/tests/agents/chat/test_chat_agent.py @@ -128,8 +128,7 @@ def test_chat_deps_state_setter_handles_initial_context(): def test_chat_deps_state_setter_parses_session_context_dict(): - """Test ChatDeps.state setter parses session_context dict into SessionContext model.""" - from haiku.rag.agents.chat.state import SessionContext + """Test ChatDeps.state setter parses session_context dict and extracts summary.""" from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState context = ToolContext() @@ -155,15 +154,10 @@ def test_chat_deps_state_setter_parses_session_context_dict(): deps.state = incoming_state - # session_context dict should be parsed into SessionContext model + # session_context dict should be parsed and summary extracted qa_session_state = context.get(QA_SESSION_NAMESPACE) assert isinstance(qa_session_state, QASessionState) assert qa_session_state.session_context == "Previous conversation summary" - assert isinstance(qa_session_state.incoming_session_context, SessionContext) - assert ( - qa_session_state.incoming_session_context.summary - == "Previous conversation summary" - ) def test_chat_deps_state_setter_generates_session_id():