diff --git a/haiku_rag_slim/haiku/rag/tools/__init__.py b/haiku_rag_slim/haiku/rag/tools/__init__.py index d81ef78a..3e961c25 100644 --- a/haiku_rag_slim/haiku/rag/tools/__init__.py +++ b/haiku_rag_slim/haiku/rag/tools/__init__.py @@ -17,16 +17,30 @@ from haiku.rag.tools.models import AnalysisResult, QAResult from haiku.rag.tools.prompts import build_tools_prompt from haiku.rag.tools.qa import create_qa_toolset from haiku.rag.tools.search import create_search_toolset +from haiku.rag.tools.toolkit import ( + FEATURE_ANALYSIS, + FEATURE_DOCUMENTS, + FEATURE_QA, + FEATURE_SEARCH, + Toolkit, + build_toolkit, +) __all__ = [ "AgentDeps", "AnalysisResult", + "FEATURE_ANALYSIS", + "FEATURE_DOCUMENTS", + "FEATURE_QA", + "FEATURE_SEARCH", "QAResult", "RAGDeps", "ToolContext", "ToolContextCache", + "Toolkit", "build_document_filter", "build_multi_document_filter", + "build_toolkit", "build_tools_prompt", "combine_filters", "create_analysis_toolset", diff --git a/haiku_rag_slim/haiku/rag/tools/toolkit.py b/haiku_rag_slim/haiku/rag/tools/toolkit.py new file mode 100644 index 00000000..bd85e4d7 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/tools/toolkit.py @@ -0,0 +1,107 @@ +from collections.abc import Callable +from dataclasses import dataclass, field + +from pydantic_ai import FunctionToolset + +from haiku.rag.config.models import AppConfig +from haiku.rag.tools.context import ToolContext, prepare_context +from haiku.rag.tools.prompts import build_tools_prompt + +FEATURE_SEARCH = "search" +FEATURE_DOCUMENTS = "documents" +FEATURE_QA = "qa" +FEATURE_ANALYSIS = "analysis" + + +@dataclass(frozen=True) +class Toolkit: + """Bundled toolsets, prompt, and context factory for haiku.rag agents. + + Created via build_toolkit(). Provides everything needed to compose + an agent with haiku.rag toolsets and create matching ToolContexts. + """ + + toolsets: list[FunctionToolset] = field(default_factory=list) + prompt: str = "" + features: list[str] = field(default_factory=list) + + def create_context(self, state_key: str | None = None) -> ToolContext: + """Create a ToolContext with namespaces matching this toolkit's features. + + Args: + state_key: Optional AG-UI state key to set on the context. + + Returns: + A prepared ToolContext. + """ + context = ToolContext() + prepare_context(context, features=self.features, state_key=state_key) + return context + + def prepare(self, context: ToolContext, state_key: str | None = None) -> None: + """Register namespaces on an existing ToolContext for this toolkit's features. + + Idempotent — safe to call multiple times on the same context. + + Args: + context: ToolContext to prepare. + state_key: Optional AG-UI state key to set on the context. + """ + prepare_context(context, features=self.features, state_key=state_key) + + +def build_toolkit( + config: AppConfig, + features: list[str] | None = None, + base_filter: str | None = None, + expand_context: bool = True, + on_qa_complete: Callable | None = None, +) -> Toolkit: + """Build a Toolkit with toolsets, prompt, and context factory for the given features. + + Args: + config: Application configuration. + features: List of features to enable. Defaults to ["search", "documents"]. + base_filter: Optional base SQL WHERE clause applied to all toolset factories. + expand_context: Whether to expand search results with surrounding context. + on_qa_complete: Optional callback invoked after each QA cycle. + + Returns: + A Toolkit ready for agent composition. + """ + if features is None: + features = [FEATURE_SEARCH, FEATURE_DOCUMENTS] + + toolsets: list[FunctionToolset] = [] + + if FEATURE_SEARCH in features: + from haiku.rag.tools.search import create_search_toolset + + toolsets.append( + create_search_toolset( + config, expand_context=expand_context, base_filter=base_filter + ) + ) + + if FEATURE_DOCUMENTS in features: + from haiku.rag.tools.document import create_document_toolset + + toolsets.append(create_document_toolset(config, base_filter=base_filter)) + + if FEATURE_QA in features: + from haiku.rag.tools.qa import create_qa_toolset + + toolsets.append( + create_qa_toolset( + config, base_filter=base_filter, on_ask_complete=on_qa_complete + ) + ) + + if FEATURE_ANALYSIS in features: + from haiku.rag.tools.analysis import create_analysis_toolset + + toolsets.append(create_analysis_toolset(config, base_filter=base_filter)) + + prompt = build_tools_prompt(features) + + return Toolkit(toolsets=toolsets, prompt=prompt, features=features) diff --git a/tests/tools/test_toolkit.py b/tests/tools/test_toolkit.py new file mode 100644 index 00000000..2b8dbbd9 --- /dev/null +++ b/tests/tools/test_toolkit.py @@ -0,0 +1,99 @@ +import pytest + +from haiku.rag.config import Config +from haiku.rag.tools.prompts import build_tools_prompt +from haiku.rag.tools.toolkit import ( + FEATURE_ANALYSIS, + FEATURE_DOCUMENTS, + FEATURE_QA, + FEATURE_SEARCH, + build_toolkit, +) + + +def test_build_toolkit_default_features(): + """Defaults to ["search", "documents"], producing 2 toolsets.""" + toolkit = build_toolkit(Config) + assert len(toolkit.toolsets) == 2 + assert toolkit.features == [FEATURE_SEARCH, FEATURE_DOCUMENTS] + + +def test_build_toolkit_all_features(): + """All 4 features produce 4 toolsets.""" + features = [FEATURE_SEARCH, FEATURE_DOCUMENTS, FEATURE_QA, FEATURE_ANALYSIS] + toolkit = build_toolkit(Config, features=features) + assert len(toolkit.toolsets) == 4 + assert toolkit.features == features + + +def test_build_toolkit_single_feature(): + """Single feature produces 1 toolset.""" + toolkit = build_toolkit(Config, features=[FEATURE_SEARCH]) + assert len(toolkit.toolsets) == 1 + assert toolkit.features == [FEATURE_SEARCH] + + +def test_build_toolkit_prompt_matches_features(): + """Toolkit prompt matches build_tools_prompt for the same features.""" + features = [FEATURE_SEARCH, FEATURE_DOCUMENTS, FEATURE_QA] + toolkit = build_toolkit(Config, features=features) + expected = build_tools_prompt(features) + assert toolkit.prompt == expected + + +def test_toolkit_create_context_registers_namespaces(): + """create_context registers correct namespaces for the features.""" + from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState + from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState + + toolkit = build_toolkit(Config, features=[FEATURE_SEARCH, FEATURE_QA]) + context = toolkit.create_context() + + assert context.get(SESSION_NAMESPACE, SessionState) is not None + assert context.get(QA_SESSION_NAMESPACE, QASessionState) is not None + + +def test_toolkit_create_context_no_qa_skips_qa_state(): + """create_context without QA feature skips QASessionState.""" + from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState + from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState + + toolkit = build_toolkit(Config, features=[FEATURE_SEARCH, FEATURE_DOCUMENTS]) + context = toolkit.create_context() + + assert context.get(SESSION_NAMESPACE, SessionState) is not None + assert context.get(QA_SESSION_NAMESPACE, QASessionState) is None + + +def test_toolkit_create_context_sets_state_key(): + """create_context propagates state_key to the ToolContext.""" + toolkit = build_toolkit(Config) + context = toolkit.create_context(state_key="my.state.key") + assert context.state_key == "my.state.key" + + +def test_toolkit_create_context_no_state_key(): + """create_context without state_key leaves it None.""" + toolkit = build_toolkit(Config) + context = toolkit.create_context() + assert context.state_key is None + + +def test_toolkit_prepare_existing_context(): + """prepare registers namespaces on an existing ToolContext.""" + from haiku.rag.tools.context import ToolContext + from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState + + toolkit = build_toolkit(Config, features=[FEATURE_SEARCH]) + context = ToolContext() + toolkit.prepare(context, state_key="test.key") + + assert context.get(SESSION_NAMESPACE, SessionState) is not None + assert context.state_key == "test.key" + + +def test_toolkit_frozen(): + """Toolkit is immutable after creation.""" + toolkit = build_toolkit(Config) + with pytest.raises(AttributeError): + toolkit.features = ["search"] # type: ignore[misc]