Remove incoming_* fields, simplify delta computation

This commit is contained in:
Yiorgis Gozadinos 2026-02-11 14:31:37 +02:00
parent d67df09cce
commit d88f2f003a
No known key found for this signature in database
5 changed files with 23 additions and 66 deletions

View file

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

View file

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

View file

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

View file

@ -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] = []

View file

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