Move SessionContext from agents/chat/state.py to tools/session.py, change QASessionState.session_context to use it

This commit is contained in:
Yiorgis Gozadinos 2026-02-12 17:13:01 +02:00
parent 05ebd781d4
commit 0c6d73409d
No known key found for this signature in database
12 changed files with 79 additions and 62 deletions

View file

@ -14,7 +14,6 @@ from haiku.rag.agents.chat.prompts import build_chat_prompt
from haiku.rag.agents.chat.state import (
AGUI_STATE_KEY,
ChatSessionState,
SessionContext,
)
__all__ = [
@ -31,5 +30,4 @@ __all__ = [
"trigger_background_summarization",
"ChatDeps",
"ChatSessionState",
"SessionContext",
]

View file

@ -10,7 +10,6 @@ from haiku.rag.agents.chat.prompts import build_chat_prompt
from haiku.rag.agents.chat.state import (
AGUI_STATE_KEY,
ChatSessionState,
SessionContext,
build_chat_state_snapshot,
)
from haiku.rag.client import HaikuRAG
@ -23,7 +22,7 @@ from haiku.rag.tools.qa import (
create_qa_toolset,
)
from haiku.rag.tools.search import create_search_toolset
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionContext, SessionState
from haiku.rag.utils import get_model
FEATURE_SEARCH = "search"
@ -99,18 +98,18 @@ class ChatDeps:
# Prefer server's session_context (background summarizer may
# have updated it since the client's last snapshot).
if not qa_session_state.session_context:
if qa_session_state.session_context is None:
session_context = state_data.get("session_context")
if isinstance(session_context, dict):
qa_session_state.session_context = SessionContext(
**session_context
).summary
qa_session_state.session_context = SessionContext(**session_context)
# Handle initial_context -> session_context for first message
if "initial_context" in state_data:
initial = state_data.get("initial_context")
if initial and not qa_session_state.session_context:
qa_session_state.session_context = initial
if initial and qa_session_state.session_context is None:
qa_session_state.session_context = SessionContext(
summary=initial
)
def prepare_chat_context(
@ -236,7 +235,6 @@ __all__ = [
"trigger_background_summarization",
"ChatDeps",
"ChatSessionState",
"SessionContext",
"AGUI_STATE_KEY",
"FEATURE_SEARCH",
"FEATURE_DOCUMENTS",

View file

@ -5,8 +5,8 @@ from typing import TYPE_CHECKING
from pydantic_ai import Agent
from haiku.rag.agents.chat.prompts import SESSION_SUMMARY_PROMPT
from haiku.rag.agents.chat.state import SessionContext
from haiku.rag.config.models import AppConfig
from haiku.rag.tools.session import SessionContext
from haiku.rag.utils import get_model
if TYPE_CHECKING:
@ -96,14 +96,19 @@ async def _update_context_background(
) -> None:
"""Background task to update session context after an ask."""
try:
current_summary = (
qa_session_state.session_context.summary
if qa_session_state.session_context is not None
else None
)
result = await update_session_context(
qa_history=list(qa_session_state.qa_history),
config=config,
current_context=qa_session_state.session_context,
current_context=current_summary,
)
if result.summary:
qa_session_state.session_context = result.summary
qa_session_state.session_context = result
except asyncio.CancelledError:
pass

View file

@ -1,9 +1,9 @@
from datetime import datetime
from typing import TYPE_CHECKING, Any
from pydantic import BaseModel
from haiku.rag.agents.research.models import Citation
from haiku.rag.tools.session import SessionContext
if TYPE_CHECKING:
from haiku.rag.tools.qa import QAHistoryEntry, QASessionState
@ -12,13 +12,6 @@ if TYPE_CHECKING:
AGUI_STATE_KEY = "haiku.rag.chat"
class SessionContext(BaseModel):
"""Compressed summary of conversation history for research graph."""
summary: str = ""
last_updated: datetime | None = None
class ChatSessionState(BaseModel):
"""State shared between frontend and agent via AG-UI."""
@ -66,10 +59,10 @@ 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 qa_state.session_context:
snapshot["session_context"] = SessionContext(
summary=qa_state.session_context
).model_dump(mode="json")
if qa_state.session_context is not None:
snapshot["session_context"] = qa_state.session_context.model_dump(
mode="json"
)
else:
snapshot["session_context"] = None

View file

@ -247,8 +247,12 @@ class ChatApp(App):
QA_SESSION_NAMESPACE, QASessionState
)
if qa_session_state is not None:
if not qa_session_state.session_context and self._initial_context:
qa_session_state.session_context = self._initial_context
if qa_session_state.session_context is None and self._initial_context:
from haiku.rag.tools.session import SessionContext
qa_session_state.session_context = SessionContext(
summary=self._initial_context
)
deps = ChatDeps(
config=self.config,
@ -362,17 +366,13 @@ class ChatApp(App):
async def action_show_context(self) -> None:
"""Show context modal (edit initial context or view session context)."""
from haiku.rag.agents.chat.state import SessionContext
from haiku.rag.chat.widgets.context_modal import ContextModal
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
session_context = None
qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
if qa_session_state and qa_session_state.session_context:
session_context = SessionContext(
summary=qa_session_state.session_context,
last_updated=datetime.now(),
)
if qa_session_state and qa_session_state.session_context is not None:
session_context = qa_session_state.session_context
await self.push_screen(
ContextModal(

View file

@ -8,7 +8,7 @@ from textual.screen import ModalScreen
from textual.widgets import Button, Markdown, Static, TextArea
if TYPE_CHECKING:
from haiku.rag.agents.chat.state import SessionContext
from haiku.rag.tools.session import SessionContext
class ContextModal(ModalScreen): # pragma: no cover

View file

@ -23,6 +23,7 @@ from haiku.rag.tools.qa import (
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
from haiku.rag.tools.session import (
SESSION_NAMESPACE,
SessionContext,
SessionState,
compute_combined_state_delta,
compute_state_delta,
@ -52,6 +53,7 @@ __all__ = [
"run_qa_core",
"create_analysis_toolset",
"SESSION_NAMESPACE",
"SessionContext",
"SessionState",
"compute_state_delta",
"compute_combined_state_delta",

View file

@ -4,8 +4,6 @@ from ag_ui.core import EventType, StateSnapshotEvent
from pydantic import BaseModel, Field
from pydantic_ai import FunctionToolset, RunContext, ToolReturn
from haiku.rag.agents.chat.context import trigger_background_summarization
from haiku.rag.agents.chat.state import build_chat_state_snapshot
from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import build_research_graph
from haiku.rag.agents.research.models import Citation, SearchAnswer
@ -22,6 +20,7 @@ from haiku.rag.tools.filters import (
from haiku.rag.tools.models import QAResult
from haiku.rag.tools.session import (
SESSION_NAMESPACE,
SessionContext,
SessionState,
)
@ -65,17 +64,20 @@ class QAHistoryEntry(BaseModel):
)
# Resolve ChatSessionState forward reference to QAHistoryEntry
from haiku.rag.agents.chat.state import _rebuild_models # noqa: E402
def _resolve_chat_state_forward_refs() -> None:
from haiku.rag.agents.chat.state import _rebuild_models
_rebuild_models(QAHistoryEntry)
_rebuild_models(QAHistoryEntry)
_resolve_chat_state_forward_refs()
class QASessionState(BaseModel):
"""Extended session state for QA with embedding cache."""
qa_history: list[QAHistoryEntry] = []
session_context: str | None = None
session_context: SessionContext | None = None
QA_SESSION_NAMESPACE = "haiku.rag.qa_session"
@ -111,8 +113,8 @@ async def run_qa_core(
)
effective_session_context = session_context
if qa_session_state is not None and qa_session_state.session_context:
effective_session_context = qa_session_state.session_context
if qa_session_state is not None and qa_session_state.session_context is not None:
effective_session_context = qa_session_state.session_context.summary
effective_prior_answers = prior_answers or []
if qa_session_state is not None and qa_session_state.qa_history:
@ -203,6 +205,8 @@ async def run_qa_core(
# Enforce FIFO limit
if len(qa_session_state.qa_history) > MAX_QA_HISTORY:
qa_session_state.qa_history = qa_session_state.qa_history[-MAX_QA_HISTORY:]
from haiku.rag.agents.chat.context import trigger_background_summarization
trigger_background_summarization(
qa_session_state=qa_session_state,
config=config,
@ -265,6 +269,8 @@ def create_qa_toolset(
)
if session_state is not None:
from haiku.rag.agents.chat.state import build_chat_state_snapshot
snapshot = build_chat_state_snapshot(
session_state,
qa_session_state,

View file

@ -1,3 +1,4 @@
from datetime import datetime
from typing import Any
import jsonpatch
@ -9,6 +10,13 @@ from haiku.rag.agents.research.models import Citation
SESSION_NAMESPACE = "haiku.rag.session"
class SessionContext(BaseModel):
"""Compressed summary of conversation history for research graph."""
summary: str = ""
last_updated: datetime | None = None
class SessionState(BaseModel):
"""Session-level state for AG-UI integration.

View file

@ -16,7 +16,7 @@ from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
from haiku.rag.tools import ToolContext
from haiku.rag.tools.qa import MAX_QA_HISTORY, QAHistoryEntry
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionContext, SessionState
def extract_state_from_result(result, state_key: str = AGUI_STATE_KEY) -> dict | None:
@ -127,7 +127,10 @@ def test_chat_deps_state_setter_handles_initial_context(temp_db_path):
# initial_context should be copied to qa_session_state.session_context
qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState)
assert qa_session_state.session_context == "Background info about the project"
assert qa_session_state.session_context is not None
assert (
qa_session_state.session_context.summary == "Background info about the project"
)
client.close()
@ -160,10 +163,12 @@ def test_chat_deps_state_setter_parses_session_context_dict(temp_db_path):
deps.state = incoming_state
# session_context dict should be parsed and summary extracted
# session_context dict should be parsed into SessionContext
qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState)
assert qa_session_state.session_context == "Previous conversation summary"
assert qa_session_state.session_context is not None
assert isinstance(qa_session_state.session_context, SessionContext)
assert qa_session_state.session_context.summary == "Previous conversation summary"
client.close()
@ -174,7 +179,9 @@ def test_chat_deps_state_setter_preserves_server_session_context(temp_db_path):
client = HaikuRAG(temp_db_path, create=True)
context = ToolContext()
qa_state = QASessionState()
qa_state.session_context = "Fresh summary from background summarizer"
qa_state.session_context = SessionContext(
summary="Fresh summary from background summarizer"
)
context.register(QA_SESSION_NAMESPACE, qa_state)
context.register(SESSION_NAMESPACE, SessionState())
@ -201,8 +208,10 @@ def test_chat_deps_state_setter_preserves_server_session_context(temp_db_path):
# Server's fresher session_context should be preserved
qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState)
assert qa_session_state.session_context is not None
assert (
qa_session_state.session_context == "Fresh summary from background summarizer"
qa_session_state.session_context.summary
== "Fresh summary from background summarizer"
)
client.close()
@ -552,7 +561,7 @@ async def test_chat_agent_ask_triggers_background_summarization(
)
# Patch internal trigger to avoid concurrent HTTP calls during VCR
with patch("haiku.rag.tools.qa.trigger_background_summarization"):
with patch("haiku.rag.agents.chat.context.trigger_background_summarization"):
result = await agent.run(
"What is the highest count class in the DocLayNet dataset?",
deps=deps,
@ -572,7 +581,7 @@ async def test_chat_agent_ask_triggers_background_summarization(
# Verify session_context was populated by background task
assert qa_session_state.session_context is not None
assert qa_session_state.session_context != ""
assert qa_session_state.session_context.summary != ""
@pytest.mark.asyncio
@ -632,15 +641,16 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
# initial_context should be transferred to QASessionState
qa_session = context.get(QA_SESSION_NAMESPACE, QASessionState)
assert qa_session is not None
assert qa_session.session_context is not None
assert (
qa_session.session_context
qa_session.session_context.summary
== "The user is researching the DocLayNet dataset for a paper on document layout analysis."
)
# Patch the internal summarization trigger in the ask tool to avoid
# concurrent HTTP calls that break VCR cassette replay ordering.
with patch(
"haiku.rag.tools.qa.trigger_background_summarization",
"haiku.rag.agents.chat.context.trigger_background_summarization",
):
# First question about class labels
result1 = await agent.run(
@ -656,7 +666,7 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
await _summarization_tasks[key]
assert qa_session.session_context is not None
assert qa_session.session_context != ""
assert qa_session.session_context.summary != ""
# qa_history should have one entry
qa_session = context.get(QA_SESSION_NAMESPACE, QASessionState)
@ -665,7 +675,7 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
# Second related question - uses prior answers and updated session context
with patch(
"haiku.rag.tools.qa.trigger_background_summarization",
"haiku.rag.agents.chat.context.trigger_background_summarization",
):
result2 = await agent.run(
"How were the annotations created and how many annotators were involved?",
@ -687,7 +697,7 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
# Session context should be updated with newer summary
assert qa_session.session_context is not None
assert qa_session.session_context != ""
assert qa_session.session_context.summary != ""
@pytest.mark.asyncio

View file

@ -3,10 +3,10 @@ from pathlib import Path
import pytest
from haiku.rag.agents.chat.state import SessionContext
from haiku.rag.agents.research.models import Citation
from haiku.rag.config import Config
from haiku.rag.tools.qa import QAHistoryEntry
from haiku.rag.tools.session import SessionContext
@pytest.fixture(scope="module")

View file

@ -1,8 +1,5 @@
from haiku.rag.agents.chat.state import (
ChatSessionState,
SessionContext,
)
from haiku.rag.tools.session import SessionState
from haiku.rag.agents.chat.state import ChatSessionState
from haiku.rag.tools.session import SessionContext, SessionState
def test_max_qa_history_constant():