Simplify search agent

This commit is contained in:
Yiorgis Gozadinos 2025-09-15 15:28:50 +03:00
parent 4df5d31b18
commit 53905c2d69
No known key found for this signature in database
3 changed files with 91 additions and 97 deletions

View file

@ -22,18 +22,18 @@ class ResearchPlan(BaseModel):
sub_questions: list[str] = Field( sub_questions: list[str] = Field(
description="Decomposed sub-questions to investigate" description="Decomposed sub-questions to investigate"
) )
search_strategies: list[str] = Field(
description="Different search approaches to use"
)
success_criteria: list[str] = Field(description="Criteria for successful research")
class ResearchOrchestrator(BaseResearchAgent): class ResearchOrchestrator(BaseResearchAgent):
"""Orchestrator agent that coordinates the research workflow.""" """Orchestrator agent that coordinates the research workflow."""
def __init__( def __init__(
self, provider: str = Config.RERANK_PROVIDER, model: str = Config.RERANK_MODEL self, provider: str | None = Config.RESEARCH_PROVIDER, model: str | None = None
): ):
# Use provided values or fall back to config defaults
provider = provider or Config.RESEARCH_PROVIDER or Config.QA_PROVIDER
model = model or Config.RESEARCH_MODEL or Config.QA_MODEL
super().__init__(provider, model, output_type=ResearchPlan) super().__init__(provider, model, output_type=ResearchPlan)
self.search_agent = SearchSpecialistAgent(provider, model) self.search_agent = SearchSpecialistAgent(provider, model)
@ -55,7 +55,8 @@ class ResearchOrchestrator(BaseResearchAgent):
- Breaks down complex questions into manageable parts - Breaks down complex questions into manageable parts
- Identifies multiple search strategies - Identifies multiple search strategies
- Defines clear success criteria - Defines clear success criteria
- Ensures thorough investigation""" - Ensures thorough investigation
/no_think"""
def register_tools(self) -> None: def register_tools(self) -> None:
"""Register orchestration tools.""" """Register orchestration tools."""
@ -63,13 +64,21 @@ class ResearchOrchestrator(BaseResearchAgent):
@self.agent.tool @self.agent.tool
async def delegate_search( async def delegate_search(
ctx: RunContext[ResearchDependencies], queries: list[str], limit: int = 5 ctx: RunContext[ResearchDependencies], queries: list[str], limit: int = 5
) -> Any: ) -> list[Any]:
"""Delegate search to the search specialist agent.""" """Delegate search to the search specialist agent for multiple queries."""
# Pass the context to maintain usage tracking all_results = []
result = await self.search_agent.run(
f"Search for: {', '.join(queries)}", deps=ctx.deps, usage=ctx.usage # Search for each query
) # The search agent will automatically store results in context
return result for query in queries:
result = await self.search_agent.run(
f"Search for: {query} with limit {limit}",
deps=ctx.deps,
usage=ctx.usage,
)
all_results.append(result)
return all_results
@self.agent.tool @self.agent.tool
async def delegate_analysis( async def delegate_analysis(
@ -158,10 +167,8 @@ class ResearchOrchestrator(BaseResearchAgent):
f"Create a research plan for: {question}", deps=deps f"Create a research plan for: {question}", deps=deps
) )
if hasattr(plan_result, "output") and isinstance( assert plan_result.output and isinstance(plan_result.output, ResearchPlan)
plan_result.output, ResearchPlan context.sub_questions = plan_result.output.sub_questions
):
context.sub_questions = plan_result.output.sub_questions
# Execute research iterations # Execute research iterations
for iteration in range(max_iterations): for iteration in range(max_iterations):
@ -176,16 +183,18 @@ class ResearchOrchestrator(BaseResearchAgent):
else: else:
# Fall back to original question with variation # Fall back to original question with variation
search_prompt = f"Additional search for: {question}" search_prompt = f"Additional search for: {question}"
# Search phase - directly call the search agent
# Search phase await self.search_agent.run(search_prompt, deps=deps)
await self.run(search_prompt, deps=deps)
# Analysis phase (only if we have results) # Analysis phase (only if we have results)
if context.search_results: if context.search_results:
await self.run("Analyze the gathered information", deps=deps) await self.analysis_agent.run(
"Analyze the gathered information", deps=deps
)
# Clarification phase - evaluate completeness # Clarification phase - evaluate completeness
clarification_result = await self.run( clarification_result = await self.clarification_agent.run(
f"Evaluate the completeness of research for: {question}. " f"Evaluate the completeness of research for: {question}. "
f"Consider all information gathered so far and determine if we have sufficient " f"Consider all information gathered so far and determine if we have sufficient "
f"information to provide a comprehensive answer.", f"information to provide a comprehensive answer.",
@ -202,7 +211,9 @@ class ResearchOrchestrator(BaseResearchAgent):
break break
# Generate final report # Generate final report
report_result = await self.run("Generate the final research report", deps=deps) report_result = await self.synthesis_agent.run(
"Generate the final research report", deps=deps
)
return ( return (
report_result.output if hasattr(report_result, "output") else report_result report_result.output if hasattr(report_result, "output") else report_result
) )

View file

@ -1,26 +1,26 @@
"""Search specialist agent for advanced document retrieval.""" """Search specialist agent for document retrieval."""
from pydantic_ai import RunContext from pydantic_ai import RunContext
from haiku.rag.research.base import BaseResearchAgent, SearchResult from haiku.rag.research.base import BaseResearchAgent
from haiku.rag.research.dependencies import ResearchDependencies from haiku.rag.research.dependencies import ResearchDependencies
from haiku.rag.store.models.chunk import Chunk
class SearchSpecialistAgent(BaseResearchAgent): class SearchSpecialistAgent(BaseResearchAgent):
"""Agent specialized in advanced document search and retrieval.""" """Agent specialized in document search and retrieval."""
def __init__(self, provider: str, model: str): def __init__(self, provider: str, model: str):
super().__init__(provider, model, output_type=list[SearchResult]) # No specific output type needed - the tool handles everything
super().__init__(provider, model)
def get_system_prompt(self) -> str: def get_system_prompt(self) -> str:
return """You are a search specialist agent focused on document retrieval. return """You are a search specialist agent focused on document retrieval from a knowledge base that uses hybrid (semantic and full-text search) search.
Your role is to: Your role is to:
1. Generate multiple search queries from different perspectives 1. Understand the search query and context
2. Identify key terms and synonyms for comprehensive search 2. Execute targeted searches to find relevant documents
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.""" Use the search tool to perform the searches on the knowledge base."""
def register_tools(self) -> None: def register_tools(self) -> None:
"""Register search-specific tools.""" """Register search-specific tools."""
@ -28,46 +28,30 @@ class SearchSpecialistAgent(BaseResearchAgent):
@self.agent.tool @self.agent.tool
async def search( async def search(
ctx: RunContext[ResearchDependencies], ctx: RunContext[ResearchDependencies],
queries: str | list[str], query: str,
limit: int = 5, limit: int = 5,
) -> list[SearchResult]: ) -> list[tuple[Chunk, float]]:
"""Execute search with single or multiple query variants.""" """Execute search and return raw results from client."""
# Normalize to list # Use the default hybrid search
query_list = [queries] if isinstance(queries, str) else queries search_results = await ctx.deps.client.search(query, limit=limit)
all_results = [] # Expand context for better relevance
seen_chunk_ids = set() expanded = await ctx.deps.client.expand_context(search_results)
for query in query_list: # Store in context (convert to SearchResult for context storage)
# Use the default hybrid search from haiku.rag.research.base import SearchResult
search_results = await ctx.deps.client.search(query, limit=limit)
# Expand context for better relevance results_for_context = []
expanded = await ctx.deps.client.expand_context(search_results) for chunk, score in expanded:
results_for_context.append(
SearchResult(
content=chunk.content,
score=score,
document_uri=chunk.document_uri or "",
metadata={"chunk_id": chunk.id} if chunk.id else {},
)
)
ctx.deps.context.add_search_result(query, results_for_context)
for chunk, score in expanded: # Return raw chunk, score tuples
# Avoid duplicates based on chunk ID return expanded
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

@ -101,12 +101,14 @@ async def test_search_single_query(mock_client, research_deps):
# Test the tool # Test the tool
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage()) ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
results = await search_tool(ctx, queries="climate change") results = await search_tool(ctx, query="climate change")
# Verify results # Verify results - should be list of (Chunk, float) tuples
assert isinstance(results, list)
assert len(results) == 2 assert len(results) == 2
assert results[0].content == "Climate change is a global phenomenon" assert results[0][0].content == "Climate change is a global phenomenon"
assert results[0].metadata["chunk_id"] == "chunk1" assert results[0][0].id == "chunk1"
assert results[0][1] == 0.8 # score
# Verify mock was called # Verify mock was called
mock_client.search.assert_called_once_with("climate change", limit=5) mock_client.search.assert_called_once_with("climate change", limit=5)
@ -114,20 +116,19 @@ async def test_search_single_query(mock_client, research_deps):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_search_multiple_queries_deduplication(mock_client, research_deps): async def test_search_with_limit(mock_client, research_deps):
"""Test deduplication when searching with multiple queries.""" """Test that search respects the limit parameter."""
# Create chunks with duplicate IDs across queries # Create more chunks than the limit
chunks_q1 = [ mock_chunks = [
create_mock_chunk("chunk1", "Content 1", 0.9), create_mock_chunk("chunk1", "Content 1", 0.9),
create_mock_chunk("chunk2", "Content 2", 0.7), create_mock_chunk("chunk2", "Content 2", 0.8),
] create_mock_chunk("chunk3", "Content 3", 0.7),
chunks_q2 = [ create_mock_chunk("chunk4", "Content 4", 0.6),
create_mock_chunk("chunk1", "Content 1", 0.9), # Duplicate create_mock_chunk("chunk5", "Content 5", 0.5),
create_mock_chunk("chunk3", "Content 3", 0.8),
] ]
mock_client.search.side_effect = [[chunks_q1[0]], [chunks_q2[0]]] mock_client.search.return_value = mock_chunks[:3]
mock_client.expand_context.side_effect = [chunks_q1, chunks_q2] mock_client.expand_context.return_value = mock_chunks[:3]
agent = SearchSpecialistAgent(provider="openai", model="gpt-4") agent = SearchSpecialistAgent(provider="openai", model="gpt-4")
@ -135,20 +136,18 @@ async def test_search_multiple_queries_deduplication(mock_client, research_deps)
search_tool = get_agent_tool(agent, "search") search_tool = get_agent_tool(agent, "search")
assert search_tool is not None assert search_tool is not None
# Test the tool # Test the tool with limit
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage()) ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
results = await search_tool(ctx, queries=["query1", "query2"]) results = await search_tool(ctx, query="test query", limit=3)
# Check deduplication - chunk1 should appear only once # Verify results respect limit
chunk_ids = [r.metadata["chunk_id"] for r in results] assert isinstance(results, list)
assert chunk_ids.count("chunk1") == 1 assert len(results) == 3
assert "chunk2" in chunk_ids assert all(isinstance(r, tuple) and len(r) == 2 for r in results)
assert "chunk3" in chunk_ids assert all(isinstance(r[0], Chunk) and isinstance(r[1], float) for r in results)
# Verify sorting by score # Verify mock was called with correct limit
assert all( mock_client.search.assert_called_once_with("test query", limit=3)
results[i].score >= results[i + 1].score for i in range(len(results) - 1)
)
@pytest.mark.asyncio @pytest.mark.asyncio
@ -167,7 +166,7 @@ async def test_search_updates_context(mock_client, research_deps):
# Test the tool # Test the tool
ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage()) ctx = RunContext(deps=research_deps, model=TestModel(), usage=RunUsage())
await search_tool(ctx, queries="test query") await search_tool(ctx, query="test query")
# Verify context was updated # Verify context was updated
assert len(research_deps.context.search_results) == 1 assert len(research_deps.context.search_results) == 1