Search agent

This commit is contained in:
Yiorgis Gozadinos 2025-09-12 12:36:48 +03:00
parent b568293661
commit b03ee08f18
No known key found for this signature in database
2 changed files with 247 additions and 0 deletions

View file

@ -0,0 +1,73 @@
"""Search specialist agent for advanced document retrieval."""
from pydantic_ai import RunContext
from haiku.rag.research.base import BaseResearchAgent, SearchResult
from haiku.rag.research.dependencies import ResearchDependencies
class SearchSpecialistAgent(BaseResearchAgent):
"""Agent specialized in advanced document search and retrieval."""
def __init__(self, provider: str, model: str):
super().__init__(provider, model, output_type=list[SearchResult])
def get_system_prompt(self) -> str:
return """You are a search specialist agent focused on document retrieval.
Your role is to:
1. Generate multiple search queries from different perspectives
2. Identify key terms and synonyms for comprehensive search
3. Execute searches and rank results by relevance
4. Return the most relevant documents for the research question
Use the search tools to explore the knowledge base thoroughly."""
def register_tools(self) -> None:
"""Register search-specific tools."""
@self.agent.tool
async def search(
ctx: RunContext[ResearchDependencies],
queries: str | list[str],
limit: int = 5,
) -> list[SearchResult]:
"""Execute search with single or multiple query variants."""
# Normalize to list
query_list = [queries] if isinstance(queries, str) else queries
all_results = []
seen_chunk_ids = set()
for query in query_list:
# Use the default hybrid search
search_results = await ctx.deps.client.search(query, limit=limit)
# Expand context for better relevance
expanded = await ctx.deps.client.expand_context(search_results)
for chunk, score in expanded:
# Avoid duplicates based on chunk ID
if chunk.id and chunk.id not in seen_chunk_ids:
seen_chunk_ids.add(chunk.id)
all_results.append(
SearchResult(
content=chunk.content,
score=score,
document_uri=chunk.document_uri or "",
metadata={"query": query, "chunk_id": chunk.id},
)
)
# Sort by score and limit results
all_results.sort(key=lambda x: x.score, reverse=True)
final_results = all_results[: limit * len(query_list)]
# Store in context
query_summary = (
query_list[0]
if len(query_list) == 1
else f"Multi-query: {len(query_list)} variants"
)
ctx.deps.context.add_search_result(query_summary, final_results)
return final_results

View file

@ -0,0 +1,174 @@
"""Tests for the search specialist agent."""
from unittest.mock import AsyncMock, create_autospec
import pytest
from pydantic_ai import RunContext
from pydantic_ai.models.test import TestModel
from pydantic_ai.usage import RunUsage
from haiku.rag.client import HaikuRAG
from haiku.rag.research.dependencies import ResearchContext, ResearchDependencies
from haiku.rag.research.search_agent import SearchSpecialistAgent
from haiku.rag.store.models.chunk import Chunk
@pytest.fixture
def mock_client():
"""Create a mock HaikuRAG client."""
client = create_autospec(HaikuRAG, instance=True)
client.search = AsyncMock()
client.expand_context = AsyncMock()
return client
@pytest.fixture
def research_context():
"""Create a research context for testing."""
return ResearchContext(
original_question="What is climate change?",
)
@pytest.fixture
def research_deps(mock_client, research_context):
"""Create research dependencies for testing."""
return ResearchDependencies(client=mock_client, context=research_context)
def create_mock_chunk(chunk_id: str, content: str, score: float = 0.8):
"""Helper to create mock chunk objects."""
return Chunk(
id=chunk_id,
document_id=f"doc_{chunk_id}",
content=content,
document_uri=f"doc_{chunk_id}.md",
metadata={},
), score
def get_agent_tool(agent, tool_name: str):
"""Helper to get a tool from an agent by name."""
tools = agent.agent._function_toolset.tools
if tool_name in tools:
return tools[tool_name].function
return None
@pytest.mark.asyncio
async def test_search_agent_has_search_tool():
"""Test that the search agent registers a search tool."""
test_model = TestModel()
# Use a valid provider for initialization
agent = SearchSpecialistAgent(provider="openai", model="gpt-4")
# Run agent with TestModel to check tools
with agent.agent.override(model=test_model):
await agent.agent.run(
"test",
deps=ResearchDependencies(
client=create_autospec(HaikuRAG, instance=True),
context=ResearchContext(original_question="test"),
),
)
# Verify the search tool was registered
assert test_model.last_model_request_parameters is not None
tools = test_model.last_model_request_parameters.function_tools
assert tools is not None
assert len(tools) == 1
assert tools[0].name == "search"
@pytest.mark.asyncio
async def test_search_single_query(mock_client, research_deps):
"""Test that search tool is called with single query."""
# Setup mock responses
mock_chunks = [
create_mock_chunk("chunk1", "Climate change is a global phenomenon"),
create_mock_chunk("chunk2", "Rising temperatures affect ecosystems"),
]
mock_client.search.return_value = mock_chunks[:1]
mock_client.expand_context.return_value = mock_chunks
# Create agent
agent = SearchSpecialistAgent(provider="openai", model="gpt-4")
# Get the search tool
search_tool = get_agent_tool(agent, "search")
assert search_tool is not None
# Test the tool
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
results = await search_tool(ctx, queries="climate change")
# Verify results
assert len(results) == 2
assert results[0].content == "Climate change is a global phenomenon"
assert results[0].metadata["chunk_id"] == "chunk1"
# Verify mock was called
mock_client.search.assert_called_once_with("climate change", limit=5)
mock_client.expand_context.assert_called_once()
@pytest.mark.asyncio
async def test_search_multiple_queries_deduplication(mock_client, research_deps):
"""Test deduplication when searching with multiple queries."""
# Create chunks with duplicate IDs across queries
chunks_q1 = [
create_mock_chunk("chunk1", "Content 1", 0.9),
create_mock_chunk("chunk2", "Content 2", 0.7),
]
chunks_q2 = [
create_mock_chunk("chunk1", "Content 1", 0.9), # Duplicate
create_mock_chunk("chunk3", "Content 3", 0.8),
]
mock_client.search.side_effect = [[chunks_q1[0]], [chunks_q2[0]]]
mock_client.expand_context.side_effect = [chunks_q1, chunks_q2]
agent = SearchSpecialistAgent(provider="openai", model="gpt-4")
# Get the search tool
search_tool = get_agent_tool(agent, "search")
assert search_tool is not None
# Test the tool
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
results = await search_tool(ctx, queries=["query1", "query2"])
# Check deduplication - chunk1 should appear only once
chunk_ids = [r.metadata["chunk_id"] for r in results]
assert chunk_ids.count("chunk1") == 1
assert "chunk2" in chunk_ids
assert "chunk3" in chunk_ids
# Verify sorting by score
assert all(
results[i].score >= results[i + 1].score for i in range(len(results) - 1)
)
@pytest.mark.asyncio
async def test_search_updates_context(mock_client, research_deps):
"""Test that search results are stored in context."""
mock_chunks = [create_mock_chunk("chunk1", "Test content")]
mock_client.search.return_value = []
mock_client.expand_context.return_value = mock_chunks
agent = SearchSpecialistAgent(provider="openai", model="gpt-4")
# Get the search tool
search_tool = get_agent_tool(agent, "search")
assert search_tool is not None
# Test the tool
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
await search_tool(ctx, queries="test query")
# Verify context was updated
assert len(research_deps.context.search_results) == 1
assert research_deps.context.search_results[0]["query"] == "test query"