Add SearchAgent for internal query expansion
SearchAgent generates 2-4 diverse search queries internally, runs them against the knowledge base, deduplicates by chunk_id, and returns consolidated results. This prevents the outer chat agent from making multiple search calls. Also updates system prompt to enforce single tool calls per message. 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
parent
49f5c20757
commit
8b7c679159
2 changed files with 151 additions and 18 deletions
|
|
@ -77,14 +77,16 @@ You have access to a knowledge base of documents. Use your tools to search and a
|
||||||
|
|
||||||
CRITICAL RULES:
|
CRITICAL RULES:
|
||||||
1. For greetings or casual chat: respond directly WITHOUT using any tools
|
1. For greetings or casual chat: respond directly WITHOUT using any tools
|
||||||
2. For questions: ALWAYS use the "ask" tool - it provides answers with proper citations
|
2. For questions: Use the "ask" tool EXACTLY ONCE - it handles query expansion internally
|
||||||
3. NEVER make up information - always use tools to get facts from the knowledge base
|
3. For searches: Use the "search" tool EXACTLY ONCE - it handles multi-query expansion internally
|
||||||
|
4. NEVER call the same tool multiple times for a single user message
|
||||||
|
5. NEVER make up information - always use tools to get facts from the knowledge base
|
||||||
|
|
||||||
How to decide which tool to use:
|
How to decide which tool to use:
|
||||||
- "ask" - DEFAULT CHOICE for any question. Use this for questions like "What is X?", "How does Y work?", "Explain Z", etc. Returns answers with citations. The ask tool maintains conversation context, so follow-up questions benefit from previous answers.
|
- "ask" - DEFAULT CHOICE for any question. Call it ONCE with the user's question. It internally handles query decomposition and returns answers with citations.
|
||||||
- "search" - ONLY use when explicitly exploring/browsing the knowledge base, or when the user asks to "search for" or "find" something without needing an answer.
|
- "search" - Use when the user explicitly asks to search/find/explore documents. Call it ONCE. After calling search, just output the list of results returned by the tool verbatim. Do NOT summarize or add commentary.
|
||||||
|
|
||||||
Be friendly and conversational. When you use tools, summarize the key findings for the user."""
|
Be friendly and conversational. When you use the "ask" tool, summarize the key findings for the user."""
|
||||||
|
|
||||||
|
|
||||||
def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
||||||
|
|
@ -102,36 +104,74 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
||||||
async def search(
|
async def search(
|
||||||
ctx: RunContext[ChatDeps],
|
ctx: RunContext[ChatDeps],
|
||||||
query: str,
|
query: str,
|
||||||
limit: int = 5,
|
|
||||||
document_filter: str | None = None,
|
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Search the knowledge base for relevant documents.
|
"""Search the knowledge base for relevant documents.
|
||||||
|
|
||||||
Use this when you need to find documents or explore the knowledge base.
|
Use this when you need to find documents or explore the knowledge base.
|
||||||
Returns relevant chunks with metadata.
|
Results are displayed to the user - just list the titles found.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
query: The search query
|
query: The search query
|
||||||
limit: Maximum number of results (default 5)
|
|
||||||
document_filter: Optional SQL WHERE clause to filter documents (e.g. "id IN ('doc1', 'doc2')")
|
|
||||||
"""
|
"""
|
||||||
|
from search_agent import SearchAgent
|
||||||
|
|
||||||
if ctx.deps.agui_emitter:
|
if ctx.deps.agui_emitter:
|
||||||
ctx.deps.agui_emitter.log(f"Searching: {query}")
|
ctx.deps.agui_emitter.log(f"Searching: {query}")
|
||||||
|
|
||||||
results = await ctx.deps.client.search(
|
# Build context from conversation history
|
||||||
query, limit=limit, filter=document_filter
|
context = None
|
||||||
)
|
if ctx.deps.session_state and ctx.deps.session_state.qa_history:
|
||||||
results = await ctx.deps.client.expand_context(results)
|
context = format_conversation_context(ctx.deps.session_state.qa_history)
|
||||||
|
|
||||||
|
# Use search agent for query expansion and deduplication
|
||||||
|
search_agent = SearchAgent(ctx.deps.client, ctx.deps.config)
|
||||||
|
results = await search_agent.search(query, context=context)
|
||||||
|
|
||||||
# Store for potential citation resolution
|
# Store for potential citation resolution
|
||||||
ctx.deps.search_results = results
|
ctx.deps.search_results = results
|
||||||
|
|
||||||
if not results:
|
if not results:
|
||||||
return "No results found for your query."
|
return "No results found."
|
||||||
|
|
||||||
# Format results for the agent
|
# Build citation infos for frontend display
|
||||||
parts = [r.format_for_agent() for r in results]
|
citation_infos = [
|
||||||
return "\n\n".join(parts)
|
CitationInfo(
|
||||||
|
index=i + 1,
|
||||||
|
document_id=r.document_id or "",
|
||||||
|
chunk_id=r.chunk_id or "",
|
||||||
|
document_uri=r.document_uri or "",
|
||||||
|
document_title=r.document_title,
|
||||||
|
page_numbers=r.page_numbers or [],
|
||||||
|
headings=r.headings,
|
||||||
|
content=r.content,
|
||||||
|
)
|
||||||
|
for i, r in enumerate(results)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Emit search results as citations
|
||||||
|
if ctx.deps.agui_emitter:
|
||||||
|
ctx.deps.agui_emitter.update_state(
|
||||||
|
ChatSessionState(
|
||||||
|
session_id=(
|
||||||
|
ctx.deps.session_state.session_id
|
||||||
|
if ctx.deps.session_state
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
citations=citation_infos,
|
||||||
|
qa_history=(
|
||||||
|
ctx.deps.session_state.qa_history
|
||||||
|
if ctx.deps.session_state
|
||||||
|
else []
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Return simple list of titles for the agent to present
|
||||||
|
titles = []
|
||||||
|
for i, r in enumerate(results):
|
||||||
|
title = r.document_title or r.document_uri or "Unknown"
|
||||||
|
titles.append(f"[{i + 1}] {title}")
|
||||||
|
return f"Found {len(results)} results:\n" + "\n".join(titles)
|
||||||
|
|
||||||
@agent.tool
|
@agent.tool
|
||||||
async def ask(
|
async def ask(
|
||||||
|
|
|
||||||
93
app/backend/search_agent.py
Normal file
93
app/backend/search_agent.py
Normal file
|
|
@ -0,0 +1,93 @@
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from pydantic_ai import Agent, RunContext
|
||||||
|
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.config.models import AppConfig
|
||||||
|
from haiku.rag.store.models import SearchResult
|
||||||
|
from haiku.rag.utils import get_model
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SearchDeps:
|
||||||
|
"""Dependencies for search agent."""
|
||||||
|
|
||||||
|
client: HaikuRAG
|
||||||
|
config: AppConfig
|
||||||
|
search_results: list[SearchResult] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
SEARCH_SYSTEM_PROMPT = """You are a search query optimizer. Given a user's search request:
|
||||||
|
|
||||||
|
1. Generate 2-4 diverse search queries that cover different aspects/phrasings of the request
|
||||||
|
2. For each query, call the run_search tool
|
||||||
|
3. After all searches complete, respond with "Search complete"
|
||||||
|
|
||||||
|
Be thorough but focused. Generate queries that will find relevant results without being redundant."""
|
||||||
|
|
||||||
|
|
||||||
|
class SearchAgent:
|
||||||
|
"""Agent that generates multiple queries and consolidates results."""
|
||||||
|
|
||||||
|
def __init__(self, client: HaikuRAG, config: AppConfig):
|
||||||
|
self._client = client
|
||||||
|
self._config = config
|
||||||
|
|
||||||
|
model = get_model(config.qa.model, config)
|
||||||
|
self._agent: Agent[SearchDeps, str] = Agent(
|
||||||
|
model,
|
||||||
|
deps_type=SearchDeps,
|
||||||
|
output_type=str,
|
||||||
|
instructions=SEARCH_SYSTEM_PROMPT,
|
||||||
|
)
|
||||||
|
|
||||||
|
@self._agent.tool
|
||||||
|
async def run_search(
|
||||||
|
ctx: RunContext[SearchDeps],
|
||||||
|
query: str,
|
||||||
|
) -> str:
|
||||||
|
"""Run a single search query against the knowledge base.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: The search query
|
||||||
|
"""
|
||||||
|
limit = ctx.deps.config.search.limit
|
||||||
|
results = await ctx.deps.client.search(query, limit=limit)
|
||||||
|
results = await ctx.deps.client.expand_context(results)
|
||||||
|
ctx.deps.search_results.extend(results)
|
||||||
|
|
||||||
|
if not results:
|
||||||
|
return f"No results for: {query}"
|
||||||
|
return f"Found {len(results)} results for: {query}"
|
||||||
|
|
||||||
|
async def search(
|
||||||
|
self,
|
||||||
|
query: str,
|
||||||
|
context: str | None = None,
|
||||||
|
) -> list[SearchResult]:
|
||||||
|
"""Execute search with query expansion and deduplication.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: The user's search request
|
||||||
|
context: Optional conversation context
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Deduplicated list of SearchResult sorted by score
|
||||||
|
"""
|
||||||
|
prompt = query
|
||||||
|
if context:
|
||||||
|
prompt = f"Context: {context}\n\nSearch request: {query}"
|
||||||
|
|
||||||
|
deps = SearchDeps(client=self._client, config=self._config)
|
||||||
|
await self._agent.run(prompt, deps=deps)
|
||||||
|
|
||||||
|
# Deduplicate by chunk_id, keeping highest score
|
||||||
|
seen: dict[str, SearchResult] = {}
|
||||||
|
for result in deps.search_results:
|
||||||
|
chunk_id = result.chunk_id or ""
|
||||||
|
if chunk_id not in seen or result.score > seen[chunk_id].score:
|
||||||
|
seen[chunk_id] = result
|
||||||
|
|
||||||
|
# Sort by score descending and limit to config
|
||||||
|
limit = self._config.search.limit
|
||||||
|
return sorted(seen.values(), key=lambda r: r.score, reverse=True)[:limit]
|
||||||
Loading…
Reference in a new issue