Introduce toolkit to reduce toolset creation ceremony

This commit is contained in:
Yiorgis Gozadinos 2026-02-16 11:51:02 +02:00
parent 72c8f6e1b1
commit 7ff2123806
No known key found for this signature in database
3 changed files with 220 additions and 0 deletions

View file

@ -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",

View file

@ -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)

View file

@ -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]