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

View file

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

View file

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

View file

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

View file

@ -8,7 +8,7 @@ from textual.screen import ModalScreen
from textual.widgets import Button, Markdown, Static, TextArea from textual.widgets import Button, Markdown, Static, TextArea
if TYPE_CHECKING: if TYPE_CHECKING:
from haiku.rag.agents.chat.state import SessionContext from haiku.rag.tools.session import SessionContext
class ContextModal(ModalScreen): # pragma: no cover 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.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
from haiku.rag.tools.session import ( from haiku.rag.tools.session import (
SESSION_NAMESPACE, SESSION_NAMESPACE,
SessionContext,
SessionState, SessionState,
compute_combined_state_delta, compute_combined_state_delta,
compute_state_delta, compute_state_delta,
@ -52,6 +53,7 @@ __all__ = [
"run_qa_core", "run_qa_core",
"create_analysis_toolset", "create_analysis_toolset",
"SESSION_NAMESPACE", "SESSION_NAMESPACE",
"SessionContext",
"SessionState", "SessionState",
"compute_state_delta", "compute_state_delta",
"compute_combined_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 import BaseModel, Field
from pydantic_ai import FunctionToolset, RunContext, ToolReturn 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.dependencies import ResearchContext
from haiku.rag.agents.research.graph import build_research_graph from haiku.rag.agents.research.graph import build_research_graph
from haiku.rag.agents.research.models import Citation, SearchAnswer 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.models import QAResult
from haiku.rag.tools.session import ( from haiku.rag.tools.session import (
SESSION_NAMESPACE, SESSION_NAMESPACE,
SessionContext,
SessionState, SessionState,
) )
@ -65,17 +64,20 @@ class QAHistoryEntry(BaseModel):
) )
# Resolve ChatSessionState forward reference to QAHistoryEntry def _resolve_chat_state_forward_refs() -> None:
from haiku.rag.agents.chat.state import _rebuild_models # noqa: E402 from haiku.rag.agents.chat.state import _rebuild_models
_rebuild_models(QAHistoryEntry) _rebuild_models(QAHistoryEntry)
_resolve_chat_state_forward_refs()
class QASessionState(BaseModel): class QASessionState(BaseModel):
"""Extended session state for QA with embedding cache.""" """Extended session state for QA with embedding cache."""
qa_history: list[QAHistoryEntry] = [] qa_history: list[QAHistoryEntry] = []
session_context: str | None = None session_context: SessionContext | None = None
QA_SESSION_NAMESPACE = "haiku.rag.qa_session" QA_SESSION_NAMESPACE = "haiku.rag.qa_session"
@ -111,8 +113,8 @@ async def run_qa_core(
) )
effective_session_context = session_context effective_session_context = session_context
if qa_session_state is not None and 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 effective_session_context = qa_session_state.session_context.summary
effective_prior_answers = prior_answers or [] effective_prior_answers = prior_answers or []
if qa_session_state is not None and qa_session_state.qa_history: 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 # Enforce FIFO limit
if len(qa_session_state.qa_history) > MAX_QA_HISTORY: if len(qa_session_state.qa_history) > MAX_QA_HISTORY:
qa_session_state.qa_history = 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( trigger_background_summarization(
qa_session_state=qa_session_state, qa_session_state=qa_session_state,
config=config, config=config,
@ -265,6 +269,8 @@ def create_qa_toolset(
) )
if session_state is not None: if session_state is not None:
from haiku.rag.agents.chat.state import build_chat_state_snapshot
snapshot = build_chat_state_snapshot( snapshot = build_chat_state_snapshot(
session_state, session_state,
qa_session_state, qa_session_state,

View file

@ -1,3 +1,4 @@
from datetime import datetime
from typing import Any from typing import Any
import jsonpatch import jsonpatch
@ -9,6 +10,13 @@ from haiku.rag.agents.research.models import Citation
SESSION_NAMESPACE = "haiku.rag.session" 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): class SessionState(BaseModel):
"""Session-level state for AG-UI integration. """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.config import Config
from haiku.rag.tools import ToolContext from haiku.rag.tools import ToolContext
from haiku.rag.tools.qa import MAX_QA_HISTORY, QAHistoryEntry 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: 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 # initial_context should be copied to qa_session_state.session_context
qa_session_state = context.get(QA_SESSION_NAMESPACE) qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState) 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() client.close()
@ -160,10 +163,12 @@ def test_chat_deps_state_setter_parses_session_context_dict(temp_db_path):
deps.state = incoming_state 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) qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState) 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() 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) client = HaikuRAG(temp_db_path, create=True)
context = ToolContext() context = ToolContext()
qa_state = QASessionState() 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(QA_SESSION_NAMESPACE, qa_state)
context.register(SESSION_NAMESPACE, SessionState()) 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 # Server's fresher session_context should be preserved
qa_session_state = context.get(QA_SESSION_NAMESPACE) qa_session_state = context.get(QA_SESSION_NAMESPACE)
assert isinstance(qa_session_state, QASessionState) assert isinstance(qa_session_state, QASessionState)
assert qa_session_state.session_context is not None
assert ( assert (
qa_session_state.session_context == "Fresh summary from background summarizer" qa_session_state.session_context.summary
== "Fresh summary from background summarizer"
) )
client.close() 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 # 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( result = await agent.run(
"What is the highest count class in the DocLayNet dataset?", "What is the highest count class in the DocLayNet dataset?",
deps=deps, deps=deps,
@ -572,7 +581,7 @@ async def test_chat_agent_ask_triggers_background_summarization(
# Verify session_context was populated by background task # Verify session_context was populated by background task
assert qa_session_state.session_context is not None 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 @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 # initial_context should be transferred to QASessionState
qa_session = context.get(QA_SESSION_NAMESPACE, QASessionState) qa_session = context.get(QA_SESSION_NAMESPACE, QASessionState)
assert qa_session is not None assert qa_session is not None
assert qa_session.session_context is not None
assert ( assert (
qa_session.session_context qa_session.session_context.summary
== "The user is researching the DocLayNet dataset for a paper on document layout analysis." == "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 # Patch the internal summarization trigger in the ask tool to avoid
# concurrent HTTP calls that break VCR cassette replay ordering. # concurrent HTTP calls that break VCR cassette replay ordering.
with patch( with patch(
"haiku.rag.tools.qa.trigger_background_summarization", "haiku.rag.agents.chat.context.trigger_background_summarization",
): ):
# First question about class labels # First question about class labels
result1 = await agent.run( 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] await _summarization_tasks[key]
assert qa_session.session_context is not None 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_history should have one entry
qa_session = context.get(QA_SESSION_NAMESPACE, QASessionState) 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 # Second related question - uses prior answers and updated session context
with patch( with patch(
"haiku.rag.tools.qa.trigger_background_summarization", "haiku.rag.agents.chat.context.trigger_background_summarization",
): ):
result2 = await agent.run( result2 = await agent.run(
"How were the annotations created and how many annotators were involved?", "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 # Session context should be updated with newer summary
assert qa_session.session_context is not None assert qa_session.session_context is not None
assert qa_session.session_context != "" assert qa_session.session_context.summary != ""
@pytest.mark.asyncio @pytest.mark.asyncio

View file

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

View file

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