Introduce toolkit to reduce toolset creation ceremony
This commit is contained in:
parent
72c8f6e1b1
commit
7ff2123806
3 changed files with 220 additions and 0 deletions
|
|
@ -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",
|
||||
|
|
|
|||
107
haiku_rag_slim/haiku/rag/tools/toolkit.py
Normal file
107
haiku_rag_slim/haiku/rag/tools/toolkit.py
Normal 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)
|
||||
99
tests/tools/test_toolkit.py
Normal file
99
tests/tools/test_toolkit.py
Normal 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]
|
||||
Loading…
Reference in a new issue