314 lines
11 KiB
Python
314 lines
11 KiB
Python
from dataclasses import dataclass
|
|
from typing import TYPE_CHECKING
|
|
|
|
from pydantic import BaseModel
|
|
from pydantic_ai import Agent, RunContext, format_as_xml
|
|
|
|
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
|
|
|
|
if TYPE_CHECKING:
|
|
from haiku.rag.graph.agui.emitter import AGUIEmitter
|
|
|
|
|
|
class CitationInfo(BaseModel):
|
|
"""Citation info for frontend display."""
|
|
|
|
index: int
|
|
document_id: str
|
|
chunk_id: str
|
|
document_uri: str
|
|
document_title: str | None = None
|
|
page_numbers: list[int] = []
|
|
headings: list[str] | None = None
|
|
content: str
|
|
|
|
|
|
class QAResponse(BaseModel):
|
|
"""A Q&A pair from conversation history."""
|
|
|
|
question: str
|
|
answer: str
|
|
sources: list[str] = []
|
|
|
|
|
|
class ChatSessionState(BaseModel):
|
|
"""State shared between frontend and agent via AG-UI."""
|
|
|
|
session_id: str = ""
|
|
citations: list[CitationInfo] = []
|
|
qa_history: list[QAResponse] = []
|
|
|
|
|
|
def format_conversation_context(qa_history: list[QAResponse]) -> str:
|
|
"""Format conversation history as XML for inclusion in prompts."""
|
|
if not qa_history:
|
|
return ""
|
|
|
|
context_data = {
|
|
"previous_qa": [
|
|
{
|
|
"question": qa.question,
|
|
"answer": qa.answer,
|
|
"sources": qa.sources,
|
|
}
|
|
for qa in qa_history
|
|
],
|
|
}
|
|
return format_as_xml(context_data, root_tag="conversation_context")
|
|
|
|
|
|
@dataclass
|
|
class ChatDeps:
|
|
"""Dependencies for chat agent."""
|
|
|
|
client: HaikuRAG
|
|
config: AppConfig
|
|
agui_emitter: "AGUIEmitter | None" = None
|
|
search_results: list[SearchResult] | None = None
|
|
session_state: ChatSessionState | None = None
|
|
|
|
|
|
CHAT_SYSTEM_PROMPT = """You are a helpful research assistant powered by haiku.rag, a knowledge base system.
|
|
|
|
You have access to a knowledge base of documents. Use your tools to search and answer questions.
|
|
|
|
CRITICAL RULES:
|
|
1. For greetings or casual chat: respond directly WITHOUT using any tools
|
|
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:
|
|
- "get_document" - Use when the user references a SPECIFIC document by name, title, or URI (e.g., "summarize document X", "get the paper about Y", "fetch 2412.00566"). Retrieves the full document content.
|
|
- "ask" - Use for general questions about topics in the knowledge base when no specific document is named. It searches across all documents 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 the "ask" tool, summarize the key findings for the user."""
|
|
|
|
|
|
def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]:
|
|
"""Create the chat agent with search and ask tools."""
|
|
model = get_model(config.qa.model, config)
|
|
|
|
agent: Agent[ChatDeps, str] = Agent(
|
|
model,
|
|
deps_type=ChatDeps,
|
|
output_type=str,
|
|
instructions=CHAT_SYSTEM_PROMPT,
|
|
)
|
|
|
|
@agent.tool
|
|
async def search(
|
|
ctx: RunContext[ChatDeps],
|
|
query: str,
|
|
) -> str:
|
|
"""Search the knowledge base for relevant documents.
|
|
|
|
Use this when you need to find documents or explore the knowledge base.
|
|
Results are displayed to the user - just list the titles found.
|
|
|
|
Args:
|
|
query: The search query
|
|
"""
|
|
from search_agent import SearchAgent
|
|
|
|
if ctx.deps.agui_emitter:
|
|
ctx.deps.agui_emitter.log(f"Searching: {query}")
|
|
|
|
# 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."
|
|
|
|
# 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(
|
|
ctx: RunContext[ChatDeps],
|
|
question: str,
|
|
document_filter: str | None = None,
|
|
) -> str:
|
|
"""Answer a specific question using the knowledge base.
|
|
|
|
Use this for direct questions that need a focused answer with citations.
|
|
|
|
Args:
|
|
question: The question to answer
|
|
document_filter: Optional SQL WHERE clause to filter documents (e.g. "id IN ('doc1', 'doc2')")
|
|
"""
|
|
if ctx.deps.agui_emitter:
|
|
ctx.deps.agui_emitter.log(f"Answering: {question}")
|
|
|
|
# Build context-aware system prompt if we have history
|
|
system_prompt = None
|
|
if ctx.deps.session_state and ctx.deps.session_state.qa_history:
|
|
from haiku.rag.qa.prompts import QA_SYSTEM_PROMPT
|
|
|
|
context_xml = format_conversation_context(ctx.deps.session_state.qa_history)
|
|
system_prompt = (
|
|
f"{QA_SYSTEM_PROMPT}\n\n"
|
|
f"{context_xml}\n\n"
|
|
"Use this conversation context to provide informed answers. "
|
|
"Reference previous answers when relevant."
|
|
)
|
|
|
|
answer, citations = await ctx.deps.client.ask(
|
|
question, system_prompt=system_prompt, filter=document_filter
|
|
)
|
|
|
|
# Accumulate Q&A in session state
|
|
if ctx.deps.session_state is not None:
|
|
sources = (
|
|
[c.document_title or c.document_uri for c in citations]
|
|
if citations
|
|
else []
|
|
)
|
|
qa_response = QAResponse(
|
|
question=question,
|
|
answer=answer,
|
|
sources=list(dict.fromkeys(sources)), # dedupe preserving order
|
|
)
|
|
ctx.deps.session_state.qa_history.append(qa_response)
|
|
|
|
# Build citation infos for frontend
|
|
citation_infos = []
|
|
if citations:
|
|
citation_infos = [
|
|
CitationInfo(
|
|
index=i + 1,
|
|
document_id=c.document_id,
|
|
chunk_id=c.chunk_id,
|
|
document_uri=c.document_uri,
|
|
document_title=c.document_title,
|
|
page_numbers=c.page_numbers,
|
|
headings=c.headings,
|
|
content=c.content,
|
|
)
|
|
for i, c in enumerate(citations)
|
|
]
|
|
|
|
# Emit updated state with citations AND accumulated qa_history
|
|
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 []
|
|
),
|
|
)
|
|
)
|
|
|
|
# Format answer with citation references
|
|
if citations:
|
|
citation_refs = " ".join(f"[{i + 1}]" for i in range(len(citations)))
|
|
return f"{answer}\n\nSources: {citation_refs}"
|
|
|
|
return answer
|
|
|
|
@agent.tool
|
|
async def get_document(
|
|
ctx: RunContext[ChatDeps],
|
|
query: str,
|
|
) -> str:
|
|
"""Retrieve a specific document by title or URI.
|
|
|
|
Use this when the user wants to fetch/get/retrieve a specific document.
|
|
|
|
Args:
|
|
query: The document title or URI to look up
|
|
"""
|
|
if ctx.deps.agui_emitter:
|
|
ctx.deps.agui_emitter.log(f"Fetching document: {query}")
|
|
|
|
# Try exact URI match first
|
|
doc = await ctx.deps.client.get_document_by_uri(query)
|
|
|
|
escaped_query = query.replace("'", "''")
|
|
|
|
# If not found, try partial URI match
|
|
if doc is None:
|
|
docs = await ctx.deps.client.list_documents(
|
|
limit=1, filter=f"LOWER(uri) LIKE LOWER('%{escaped_query}%')"
|
|
)
|
|
if docs:
|
|
doc = docs[0]
|
|
|
|
# If still not found, try partial title match
|
|
if doc is None:
|
|
docs = await ctx.deps.client.list_documents(
|
|
limit=1, filter=f"LOWER(title) LIKE LOWER('%{escaped_query}%')"
|
|
)
|
|
if docs:
|
|
doc = docs[0]
|
|
|
|
if doc is None:
|
|
return f"Document not found: {query}"
|
|
|
|
return (
|
|
f"**{doc.title or 'Untitled'}**\n\n"
|
|
f"- ID: {doc.id}\n"
|
|
f"- URI: {doc.uri or 'N/A'}\n"
|
|
f"- Created: {doc.created_at.strftime('%Y-%m-%d %H:%M')}\n\n"
|
|
f"**Content:**\n{doc.content}"
|
|
)
|
|
|
|
return agent
|