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:
|
||||
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
|
||||
3. NEVER make up information - always use tools to get facts from the knowledge base
|
||||
2. For questions: Use the "ask" tool EXACTLY ONCE - it handles query expansion internally
|
||||
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:
|
||||
- "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.
|
||||
- "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.
|
||||
- "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" - 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]:
|
||||
|
|
@ -102,36 +104,74 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
|||
async def search(
|
||||
ctx: RunContext[ChatDeps],
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
document_filter: str | None = None,
|
||||
) -> str:
|
||||
"""Search the knowledge base for relevant documents.
|
||||
|
||||
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:
|
||||
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:
|
||||
ctx.deps.agui_emitter.log(f"Searching: {query}")
|
||||
|
||||
results = await ctx.deps.client.search(
|
||||
query, limit=limit, filter=document_filter
|
||||
)
|
||||
results = await ctx.deps.client.expand_context(results)
|
||||
# Build context from conversation history
|
||||
context = None
|
||||
if ctx.deps.session_state and ctx.deps.session_state.qa_history:
|
||||
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
|
||||
ctx.deps.search_results = results
|
||||
|
||||
if not results:
|
||||
return "No results found for your query."
|
||||
return "No results found."
|
||||
|
||||
# Format results for the agent
|
||||
parts = [r.format_for_agent() for r in results]
|
||||
return "\n\n".join(parts)
|
||||
# Build citation infos for frontend display
|
||||
citation_infos = [
|
||||
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
|
||||
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