From 0c6d73409d3ecb9ec1c58ea2ecd541d47fe3afa2 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 12 Feb 2026 17:13:01 +0200 Subject: [PATCH] Move SessionContext from agents/chat/state.py to tools/session.py, change QASessionState.session_context to use it --- .../haiku/rag/agents/chat/__init__.py | 2 -- haiku_rag_slim/haiku/rag/agents/chat/agent.py | 16 ++++----- .../haiku/rag/agents/chat/context.py | 11 ++++-- haiku_rag_slim/haiku/rag/agents/chat/state.py | 17 +++------ haiku_rag_slim/haiku/rag/chat/app.py | 16 ++++----- .../haiku/rag/chat/widgets/context_modal.py | 2 +- haiku_rag_slim/haiku/rag/tools/__init__.py | 2 ++ haiku_rag_slim/haiku/rag/tools/qa.py | 22 +++++++----- haiku_rag_slim/haiku/rag/tools/session.py | 8 +++++ tests/agents/chat/test_chat_agent.py | 36 ++++++++++++------- tests/agents/chat/test_context.py | 2 +- tests/agents/chat/test_state.py | 7 ++-- 12 files changed, 79 insertions(+), 62 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/agents/chat/__init__.py b/haiku_rag_slim/haiku/rag/agents/chat/__init__.py index 03b1c395..5d068881 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/__init__.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/__init__.py @@ -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", ] diff --git a/haiku_rag_slim/haiku/rag/agents/chat/agent.py b/haiku_rag_slim/haiku/rag/agents/chat/agent.py index 51bcb317..10384af3 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/agent.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/agent.py @@ -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", diff --git a/haiku_rag_slim/haiku/rag/agents/chat/context.py b/haiku_rag_slim/haiku/rag/agents/chat/context.py index f09b4dba..f77603ad 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/context.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/context.py @@ -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 diff --git a/haiku_rag_slim/haiku/rag/agents/chat/state.py b/haiku_rag_slim/haiku/rag/agents/chat/state.py index 08afa90d..51c83c14 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/state.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/state.py @@ -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 diff --git a/haiku_rag_slim/haiku/rag/chat/app.py b/haiku_rag_slim/haiku/rag/chat/app.py index c6ddc441..7fcef2ec 100644 --- a/haiku_rag_slim/haiku/rag/chat/app.py +++ b/haiku_rag_slim/haiku/rag/chat/app.py @@ -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( diff --git a/haiku_rag_slim/haiku/rag/chat/widgets/context_modal.py b/haiku_rag_slim/haiku/rag/chat/widgets/context_modal.py index c8ec876d..1e155a72 100644 --- a/haiku_rag_slim/haiku/rag/chat/widgets/context_modal.py +++ b/haiku_rag_slim/haiku/rag/chat/widgets/context_modal.py @@ -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 diff --git a/haiku_rag_slim/haiku/rag/tools/__init__.py b/haiku_rag_slim/haiku/rag/tools/__init__.py index 00e703d0..84e267b0 100644 --- a/haiku_rag_slim/haiku/rag/tools/__init__.py +++ b/haiku_rag_slim/haiku/rag/tools/__init__.py @@ -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", diff --git a/haiku_rag_slim/haiku/rag/tools/qa.py b/haiku_rag_slim/haiku/rag/tools/qa.py index d9482443..daee8f43 100644 --- a/haiku_rag_slim/haiku/rag/tools/qa.py +++ b/haiku_rag_slim/haiku/rag/tools/qa.py @@ -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, diff --git a/haiku_rag_slim/haiku/rag/tools/session.py b/haiku_rag_slim/haiku/rag/tools/session.py index fcd1e73d..b87ea3c8 100644 --- a/haiku_rag_slim/haiku/rag/tools/session.py +++ b/haiku_rag_slim/haiku/rag/tools/session.py @@ -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. diff --git a/tests/agents/chat/test_chat_agent.py b/tests/agents/chat/test_chat_agent.py index a051ee48..4d7a1544 100644 --- a/tests/agents/chat/test_chat_agent.py +++ b/tests/agents/chat/test_chat_agent.py @@ -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 diff --git a/tests/agents/chat/test_context.py b/tests/agents/chat/test_context.py index 5cf0875b..4af88877 100644 --- a/tests/agents/chat/test_context.py +++ b/tests/agents/chat/test_context.py @@ -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") diff --git a/tests/agents/chat/test_state.py b/tests/agents/chat/test_state.py index 77fcdc77..34c50651 100644 --- a/tests/agents/chat/test_state.py +++ b/tests/agents/chat/test_state.py @@ -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():