diff --git a/haiku_rag_slim/haiku/rag/agents/chat/agent.py b/haiku_rag_slim/haiku/rag/agents/chat/agent.py index 9f5d3b95..36400a41 100644 --- a/haiku_rag_slim/haiku/rag/agents/chat/agent.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/agent.py @@ -10,7 +10,7 @@ from haiku.rag.agents.chat.prompts import build_chat_prompt from haiku.rag.agents.chat.state import AGUI_STATE_KEY from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig -from haiku.rag.tools.context import ToolContext +from haiku.rag.tools.context import ToolContext, prepare_context from haiku.rag.tools.document import create_document_toolset from haiku.rag.tools.qa import ( QA_SESSION_NAMESPACE, @@ -18,7 +18,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, SessionContext, SessionState +from haiku.rag.tools.session import SessionContext from haiku.rag.utils import get_model FEATURE_SEARCH = "search" @@ -110,14 +110,7 @@ def prepare_chat_context( if features is None: features = DEFAULT_FEATURES - if context.get(SESSION_NAMESPACE, SessionState) is None: - context.register(SESSION_NAMESPACE, SessionState()) - if context.state_key is None: - context.state_key = AGUI_STATE_KEY - - if FEATURE_QA in features: - if context.get(QA_SESSION_NAMESPACE, QASessionState) is None: - context.register(QA_SESSION_NAMESPACE, QASessionState()) + prepare_context(context, features=features, state_key=AGUI_STATE_KEY) def create_chat_agent( diff --git a/haiku_rag_slim/haiku/rag/tools/__init__.py b/haiku_rag_slim/haiku/rag/tools/__init__.py index 84e267b0..b925b408 100644 --- a/haiku_rag_slim/haiku/rag/tools/__init__.py +++ b/haiku_rag_slim/haiku/rag/tools/__init__.py @@ -1,5 +1,11 @@ from haiku.rag.tools.analysis import create_analysis_toolset -from haiku.rag.tools.context import RAGDeps, ToolContext, ToolContextCache +from haiku.rag.tools.context import ( + RAGDeps, + ToolContext, + ToolContextCache, + prepare_context, +) +from haiku.rag.tools.deps import AgentDeps from haiku.rag.tools.document import ( DocumentInfo, DocumentListResponse, @@ -30,9 +36,11 @@ from haiku.rag.tools.session import ( ) __all__ = [ + "AgentDeps", "RAGDeps", "ToolContext", "ToolContextCache", + "prepare_context", "QAResult", "AnalysisResult", "build_document_filter", diff --git a/haiku_rag_slim/haiku/rag/tools/context.py b/haiku_rag_slim/haiku/rag/tools/context.py index 0b3daac0..7f1dc0f8 100644 --- a/haiku_rag_slim/haiku/rag/tools/context.py +++ b/haiku_rag_slim/haiku/rag/tools/context.py @@ -180,6 +180,36 @@ class ToolContext(BaseModel): return state +def prepare_context( + context: ToolContext, + features: list[str] | None = None, + state_key: str | None = None, +) -> None: + """Register required namespaces in a ToolContext based on feature flags. + + Idempotent — safe to call multiple times on the same context. + + Args: + context: ToolContext to prepare. + features: List of enabled features. Defaults to ["search", "documents"]. + state_key: Optional AG-UI state key to set on the context. + """ + from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState + from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState + + if features is None: + features = ["search", "documents"] + + if any(f in features for f in ("search", "qa", "analysis")): + context.get_or_create(SESSION_NAMESPACE, SessionState) + + if "qa" in features: + context.get_or_create(QA_SESSION_NAMESPACE, QASessionState) + + if state_key is not None: + context.state_key = state_key + + class ToolContextCache: """In-memory cache for ToolContext instances, keyed by external session/thread ID.""" diff --git a/haiku_rag_slim/haiku/rag/tools/deps.py b/haiku_rag_slim/haiku/rag/tools/deps.py new file mode 100644 index 00000000..b247cb8a --- /dev/null +++ b/haiku_rag_slim/haiku/rag/tools/deps.py @@ -0,0 +1,38 @@ +from dataclasses import dataclass +from typing import Any + +from haiku.rag.client import HaikuRAG +from haiku.rag.tools.context import ToolContext + + +@dataclass +class AgentDeps: + """Generic dependencies for agents using haiku.rag toolsets. + + Implements RAGDeps protocol and AG-UI state protocol. + """ + + client: HaikuRAG + tool_context: ToolContext + state_key: str | None = None + + @property + def state(self) -> dict[str, Any]: + """Get current state for AG-UI protocol.""" + snapshot = self.tool_context.build_state_snapshot() + if self.state_key: + return {self.state_key: snapshot} + return snapshot + + @state.setter + def state(self, value: dict[str, Any] | None) -> None: + """Set state from AG-UI protocol.""" + if value is None: + return + + data: dict[str, Any] = value + if self.state_key and self.state_key in value: + nested = value[self.state_key] + if isinstance(nested, dict): + data = nested + self.tool_context.restore_state_snapshot(data) diff --git a/tests/tools/test_context.py b/tests/tools/test_context.py index 4b673e06..14ad4de1 100644 --- a/tests/tools/test_context.py +++ b/tests/tools/test_context.py @@ -1,6 +1,8 @@ from pydantic import BaseModel -from haiku.rag.tools.context import ToolContext +from haiku.rag.tools.context import ToolContext, prepare_context +from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState +from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState class TestState(BaseModel): @@ -440,3 +442,47 @@ def test_restore_state_snapshot_ignores_unknown_fields(): ns1 = ctx.get("ns1", TestState) assert ns1 is not None assert ns1.value == 10 + + +# --- prepare_context tests --- + + +def test_prepare_context_default_features(): + """Default features register SessionState only.""" + ctx = ToolContext() + prepare_context(ctx) + assert ctx.get(SESSION_NAMESPACE, SessionState) is not None + assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is None + + +def test_prepare_context_with_qa(): + """QA feature registers both SessionState and QASessionState.""" + ctx = ToolContext() + prepare_context(ctx, features=["search", "qa"]) + assert ctx.get(SESSION_NAMESPACE, SessionState) is not None + assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is not None + + +def test_prepare_context_sets_state_key(): + """state_key is set on context when provided.""" + ctx = ToolContext() + prepare_context(ctx, state_key="my_app") + assert ctx.state_key == "my_app" + + +def test_prepare_context_no_state_key_by_default(): + """state_key is not set when not provided.""" + ctx = ToolContext() + prepare_context(ctx) + assert ctx.state_key is None + + +def test_prepare_context_idempotent(): + """Calling prepare_context twice doesn't create duplicate state.""" + ctx = ToolContext() + prepare_context(ctx, features=["search", "qa"]) + session1 = ctx.get(SESSION_NAMESPACE, SessionState) + qa1 = ctx.get(QA_SESSION_NAMESPACE, QASessionState) + prepare_context(ctx, features=["search", "qa"]) + assert ctx.get(SESSION_NAMESPACE, SessionState) is session1 + assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is qa1 diff --git a/tests/tools/test_deps.py b/tests/tools/test_deps.py new file mode 100644 index 00000000..2608ac30 --- /dev/null +++ b/tests/tools/test_deps.py @@ -0,0 +1,102 @@ +from unittest.mock import MagicMock + +import pytest + +from haiku.rag.tools.context import ToolContext +from haiku.rag.tools.deps import AgentDeps +from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState +from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState + + +@pytest.fixture +def mock_client(): + return MagicMock() + + +def test_agent_deps_state_getter_empty(mock_client): + """state returns empty dict when no namespaces are registered.""" + ctx = ToolContext() + deps = AgentDeps(client=mock_client, tool_context=ctx) + assert deps.state == {} + + +def test_agent_deps_state_getter_with_session(mock_client): + """state returns flat snapshot of registered namespaces.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx) + snapshot = deps.state + assert "citations" in snapshot + assert "citation_registry" in snapshot + + +def test_agent_deps_state_getter_with_state_key(mock_client): + """state wraps snapshot under state_key when set.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="my_app") + snapshot = deps.state + assert "my_app" in snapshot + assert "citations" in snapshot["my_app"] + + +def test_agent_deps_state_setter_restores(mock_client): + """state setter restores namespace fields from flat dict.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx) + deps.state = {"document_filter": ["doc1", "doc2"]} + session = ctx.get(SESSION_NAMESPACE, SessionState) + assert session is not None + assert session.document_filter == ["doc1", "doc2"] + + +def test_agent_deps_state_setter_with_state_key(mock_client): + """state setter extracts data from namespaced key.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="my_app") + deps.state = {"my_app": {"document_filter": ["doc1"]}} + session = ctx.get(SESSION_NAMESPACE, SessionState) + assert session is not None + assert session.document_filter == ["doc1"] + + +def test_agent_deps_state_setter_ignores_none(mock_client): + """state setter is a no-op when value is None.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx) + deps.state = None + session = ctx.get(SESSION_NAMESPACE, SessionState) + assert session is not None + assert session.document_filter == [] + + +def test_agent_deps_state_roundtrip(mock_client): + """Build snapshot then restore produces equivalent state.""" + ctx = ToolContext() + ctx.register(SESSION_NAMESPACE, SessionState(document_filter=["doc1"])) + ctx.register(QA_SESSION_NAMESPACE, QASessionState()) + deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="app") + + snapshot = deps.state + + ctx2 = ToolContext() + ctx2.register(SESSION_NAMESPACE, SessionState()) + ctx2.register(QA_SESSION_NAMESPACE, QASessionState()) + deps2 = AgentDeps(client=mock_client, tool_context=ctx2, state_key="app") + deps2.state = snapshot + + session = ctx2.get(SESSION_NAMESPACE, SessionState) + assert session is not None + assert session.document_filter == ["doc1"] + + +def test_agent_deps_satisfies_rag_deps_protocol(mock_client): + """AgentDeps satisfies the RAGDeps protocol.""" + from haiku.rag.tools.context import RAGDeps + + ctx = ToolContext() + deps = AgentDeps(client=mock_client, tool_context=ctx) + assert isinstance(deps, RAGDeps)