Create SearchToolset, refactor QA Agent to use it

This commit is contained in:
Yiorgis Gozadinos 2026-02-02 11:01:43 +02:00
parent a3863882cd
commit 2f229f91e1
No known key found for this signature in database
4 changed files with 366 additions and 43 deletions

View file

@ -1,5 +1,4 @@
from pydantic import BaseModel from pydantic_ai import Agent
from pydantic_ai import Agent, RunContext
from pydantic_ai.output import ToolOutput from pydantic_ai.output import ToolOutput
from haiku.rag.agents.qa.prompts import QA_SYSTEM_PROMPT from haiku.rag.agents.qa.prompts import QA_SYSTEM_PROMPT
@ -9,18 +8,13 @@ from haiku.rag.agents.research.models import (
resolve_citations, resolve_citations,
) )
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
from haiku.rag.config.models import AppConfig, ModelConfig from haiku.rag.config.models import AppConfig, ModelConfig
from haiku.rag.store.models import SearchResult from haiku.rag.tools import ToolContext
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
from haiku.rag.utils import get_model from haiku.rag.utils import get_model
class Dependencies(BaseModel):
model_config = {"arbitrary_types_allowed": True}
client: HaikuRAG
search_results: list[SearchResult] = []
search_filter: str | None = None
class QuestionAnswerAgent: class QuestionAnswerAgent:
def __init__( def __init__(
self, self,
@ -30,41 +24,14 @@ class QuestionAnswerAgent:
system_prompt: str | None = None, system_prompt: str | None = None,
): ):
self._client = client self._client = client
model_obj = get_model(model_config, config) self._config = config or Config
self._agent: Agent[None, RawSearchAnswer] = Agent(
self._agent: Agent[Dependencies, RawSearchAnswer] = Agent( model=get_model(model_config, self._config),
model=model_obj,
deps_type=Dependencies,
output_type=ToolOutput(RawSearchAnswer, max_retries=3), output_type=ToolOutput(RawSearchAnswer, max_retries=3),
instructions=system_prompt or QA_SYSTEM_PROMPT, instructions=system_prompt or QA_SYSTEM_PROMPT,
retries=3, retries=3,
) )
@self._agent.tool
async def search_documents(
ctx: RunContext[Dependencies],
query: str,
limit: int | None = None,
) -> str:
"""Search the knowledge base for relevant documents.
Returns results with chunk IDs and rank positions.
Reference results by their chunk_id in cited_chunks.
"""
results = await ctx.deps.client.search(
query, limit=limit, filter=ctx.deps.search_filter
)
results = await ctx.deps.client.expand_context(results)
# Store results for citation resolution
ctx.deps.search_results = results
# Format with rank instead of raw score to avoid confusing LLMs
total = len(results)
parts = [
r.format_for_agent(rank=i + 1, total=total)
for i, r in enumerate(results)
]
return "\n\n".join(parts) if parts else "No results found."
async def answer( async def answer(
self, question: str, filter: str | None = None self, question: str, filter: str | None = None
) -> tuple[str, list[Citation]]: ) -> tuple[str, list[Citation]]:
@ -77,8 +44,24 @@ class QuestionAnswerAgent:
Returns: Returns:
Tuple of (answer text, list of resolved citations) Tuple of (answer text, list of resolved citations)
""" """
deps = Dependencies(client=self._client, search_filter=filter) # Create context and search toolset for this run
result = await self._agent.run(question, deps=deps) context = ToolContext()
search_toolset = create_search_toolset(
self._client,
self._config,
context=context,
base_filter=filter,
tool_name="search_documents",
)
result = await self._agent.run(question, toolsets=[search_toolset])
output = result.output output = result.output
citations = resolve_citations(output.cited_chunks, deps.search_results)
# Get search results from context for citation resolution
search_state = context.get(SEARCH_NAMESPACE)
search_results = (
search_state.results if isinstance(search_state, SearchState) else []
)
citations = resolve_citations(output.cited_chunks, search_results)
return output.answer, citations return output.answer, citations

View file

@ -5,6 +5,7 @@ from haiku.rag.tools.filters import (
combine_filters, combine_filters,
) )
from haiku.rag.tools.models import AnalysisResult, QAResult from haiku.rag.tools.models import AnalysisResult, QAResult
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
__all__ = [ __all__ = [
"ToolContext", "ToolContext",
@ -13,4 +14,7 @@ __all__ = [
"build_document_filter", "build_document_filter",
"build_multi_document_filter", "build_multi_document_filter",
"combine_filters", "combine_filters",
"SEARCH_NAMESPACE",
"SearchState",
"create_search_toolset",
] ]

View file

@ -0,0 +1,91 @@
from pydantic import BaseModel
from pydantic_ai import FunctionToolset
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.store.models import SearchResult
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.filters import combine_filters
SEARCH_NAMESPACE = "haiku.rag.search"
class SearchState(BaseModel):
"""State for search toolset.
Accumulates search results across tool invocations.
"""
results: list[SearchResult] = []
def create_search_toolset(
client: HaikuRAG,
config: AppConfig,
context: ToolContext | None = None,
expand_context: bool = True,
base_filter: str | None = None,
tool_name: str = "search",
) -> FunctionToolset:
"""Create a toolset with search capabilities.
Args:
client: HaikuRAG client for search operations.
config: Application configuration.
context: Optional ToolContext for state accumulation.
If provided, search results are accumulated in SearchState.
expand_context: Whether to expand search results with surrounding context.
Defaults to True.
base_filter: Optional base SQL WHERE clause applied to all searches.
Combined with any filter passed to the search tool.
tool_name: Name for the search tool. Defaults to "search".
Returns:
FunctionToolset with a search tool.
"""
# Get or create search state if context provided
state: SearchState | None = None
if context is not None:
state = context.get_or_create(SEARCH_NAMESPACE, SearchState)
async def search(
query: str,
limit: int | None = None,
filter: str | None = None,
) -> str:
"""Search the knowledge base for relevant documents.
Args:
query: The search query (what to search for).
limit: Number of results to return (default: from config).
filter: Optional SQL WHERE clause to filter documents.
Returns:
Formatted search results with content and metadata.
"""
effective_limit = limit or config.search.limit
effective_filter = combine_filters(base_filter, filter)
results = await client.search(
query, limit=effective_limit, filter=effective_filter
)
if expand_context:
results = await client.expand_context(results)
# Accumulate results in state if context provided
if state is not None:
state.results.extend(results)
if not results:
return "No results found."
# Format results for agent context
total = len(results)
formatted = [
r.format_for_agent(rank=i + 1, total=total) for i, r in enumerate(results)
]
return "\n\n".join(formatted)
toolset = FunctionToolset()
toolset.add_function(search, name=tool_name)
return toolset

245
tests/tools/test_search.py Normal file
View file

@ -0,0 +1,245 @@
import pytest
from haiku.rag.tools import ToolContext
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
class TestSearchState:
"""Tests for SearchState model."""
def test_search_state_defaults(self):
"""SearchState initializes with empty results."""
state = SearchState()
assert state.results == []
def test_search_state_add_results(self):
"""Can add results to SearchState."""
from haiku.rag.store.models import SearchResult
state = SearchState()
result = SearchResult(content="test content", score=0.9, chunk_id="chunk1")
state.results.append(result)
assert len(state.results) == 1
assert state.results[0].chunk_id == "chunk1"
def test_search_state_serialization(self):
"""SearchState serializes and deserializes correctly."""
from haiku.rag.store.models import SearchResult
state = SearchState()
state.results.append(
SearchResult(
content="test",
score=0.8,
chunk_id="c1",
document_title="Doc Title",
)
)
# Serialize
data = state.model_dump()
assert "results" in data
assert len(data["results"]) == 1
# Deserialize
restored = SearchState.model_validate(data)
assert len(restored.results) == 1
assert restored.results[0].chunk_id == "c1"
class TestSearchToolset:
"""Tests for create_search_toolset."""
def test_create_search_toolset_returns_function_toolset(
self, search_client, search_config
):
"""create_search_toolset returns a FunctionToolset."""
from pydantic_ai import FunctionToolset
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
assert isinstance(toolset, FunctionToolset)
def test_search_toolset_has_search_tool(self, search_client, search_config):
"""The toolset includes a 'search' tool."""
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
# toolset.tools is a dict with tool names as keys
assert "search" in toolset.tools
def test_search_toolset_registers_state(self, search_client, search_config):
"""Toolset registers SearchState under SEARCH_NAMESPACE."""
context = ToolContext()
create_search_toolset(search_client, search_config, context)
state = context.get(SEARCH_NAMESPACE)
assert state is not None
assert isinstance(state, SearchState)
def test_search_toolset_uses_existing_state(self, search_client, search_config):
"""Toolset uses existing state if already registered."""
from haiku.rag.store.models import SearchResult
context = ToolContext()
existing_state = SearchState()
existing_state.results.append(
SearchResult(content="pre-existing", score=0.5, chunk_id="pre1")
)
context.register(SEARCH_NAMESPACE, existing_state)
create_search_toolset(search_client, search_config, context)
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
assert len(state.results) == 1
assert state.results[0].chunk_id == "pre1"
class TestSearchToolExecution:
"""Tests for search tool execution."""
@pytest.mark.asyncio
async def test_search_returns_formatted_results(self, search_client, search_config):
"""Search tool returns formatted results."""
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
# Get the search function
search_tool = toolset.tools["search"]
result = await search_tool.function("Python")
assert "Python" in result or "programming" in result
assert "No results found" not in result
@pytest.mark.asyncio
async def test_search_accumulates_in_state(self, search_client, search_config):
"""Search tool accumulates results in SearchState."""
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
# Run search
search_tool = toolset.tools["search"]
await search_tool.function("Python")
# Check state was updated
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
assert len(state.results) > 0
assert any("Python" in r.content for r in state.results)
@pytest.mark.asyncio
async def test_search_with_no_results(self, temp_db_path, search_config):
"""Search tool returns appropriate message when no results."""
from haiku.rag.client import HaikuRAG
# Use empty database
async with HaikuRAG(temp_db_path, create=True) as empty_client:
context = ToolContext()
toolset = create_search_toolset(empty_client, search_config, context)
search_tool = toolset.tools["search"]
result = await search_tool.function("anything")
assert result == "No results found."
@pytest.mark.asyncio
async def test_search_with_filter(self, search_client, search_config):
"""Search tool respects filter parameter."""
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
search_tool = toolset.tools["search"]
# Filter to only Python documents
await search_tool.function("programming", filter="title LIKE '%Python%'")
# Should find Python but not JavaScript
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
for r in state.results:
assert "JavaScript" not in (r.document_title or "")
@pytest.mark.asyncio
async def test_search_without_context(self, search_client, search_config):
"""Search tool works without ToolContext."""
toolset = create_search_toolset(search_client, search_config, context=None)
search_tool = toolset.tools["search"]
result = await search_tool.function("Python")
# Should still return results
assert "Python" in result or "programming" in result
@pytest.mark.asyncio
async def test_search_multiple_accumulates(self, search_client, search_config):
"""Multiple searches accumulate results in state."""
context = ToolContext()
toolset = create_search_toolset(search_client, search_config, context)
search_tool = toolset.tools["search"]
await search_tool.function("Python")
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
first_count = len(state.results)
await search_tool.function("JavaScript")
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
second_count = len(state.results)
assert second_count > first_count
@pytest.mark.asyncio
async def test_search_with_base_filter(self, search_client, search_config):
"""Search toolset respects base_filter parameter."""
context = ToolContext()
# Create toolset with base_filter for Python documents only
toolset = create_search_toolset(
search_client,
search_config,
context,
base_filter="title LIKE '%Python%'",
)
search_tool = toolset.tools["search"]
await search_tool.function("programming")
# Should only find Python documents
state = context.get(SEARCH_NAMESPACE)
assert isinstance(state, SearchState)
assert len(state.results) > 0
for r in state.results:
assert "JavaScript" not in (r.document_title or "")
@pytest.fixture
def search_client(temp_db_path):
"""Create a HaikuRAG client with test data for search tests."""
import asyncio
from haiku.rag.client import HaikuRAG
async def setup():
rag = HaikuRAG(temp_db_path, create=True)
await rag.__aenter__()
await rag.create_document(
"Python is a programming language. It is widely used for web development.",
uri="test://python",
title="Python Guide",
)
await rag.create_document(
"JavaScript runs in the browser. It powers interactive web pages.",
uri="test://javascript",
title="JavaScript Guide",
)
return rag
return asyncio.get_event_loop().run_until_complete(setup())
@pytest.fixture
def search_config():
"""Default AppConfig for search tests."""
from haiku.rag.config import Config
return Config