174 lines
5.2 KiB
Python
174 lines
5.2 KiB
Python
from pydantic import BaseModel
|
|
from pydantic_ai import Agent, FunctionToolset, RunContext
|
|
|
|
from haiku.rag.client import HaikuRAG
|
|
from haiku.rag.config.models import AppConfig
|
|
from haiku.rag.tools.context import RAGDeps
|
|
from haiku.rag.utils import get_model
|
|
|
|
DOCUMENT_SUMMARY_PROMPT = """Generate a summary of the document content provided below.
|
|
|
|
Start with a one-paragraph overview, then list the main topics covered, and highlight any key findings or conclusions.
|
|
|
|
Guidelines:
|
|
- Aim for 1-2 paragraphs for short documents, 3-4 paragraphs for longer ones
|
|
- Focus on factual content and key information
|
|
- Do not include meta-commentary like "This document discusses..." or "The document covers..."
|
|
- Do not speculate beyond what's in the content
|
|
|
|
Document content:
|
|
{content}"""
|
|
|
|
|
|
class DocumentInfo(BaseModel):
|
|
"""Document info for list_documents response."""
|
|
|
|
id: str | None = None
|
|
title: str
|
|
uri: str
|
|
created: str
|
|
|
|
|
|
class DocumentListResponse(BaseModel):
|
|
"""Response from list_documents tool."""
|
|
|
|
documents: list[DocumentInfo]
|
|
page: int
|
|
total_pages: int
|
|
total_documents: int
|
|
|
|
|
|
async def find_document(client: HaikuRAG, query: str):
|
|
"""Find a document by exact URI, partial URI, or partial title match."""
|
|
doc = await client.get_document_by_uri(query)
|
|
if doc is not None:
|
|
return doc
|
|
|
|
escaped_query = query.replace("'", "''")
|
|
# Also try without spaces for matching "TB MED 593" to "tbmed593"
|
|
no_spaces = escaped_query.replace(" ", "")
|
|
|
|
docs = await client.list_documents(
|
|
limit=1,
|
|
filter=f"LOWER(uri) LIKE LOWER('%{escaped_query}%') OR LOWER(uri) LIKE LOWER('%{no_spaces}%')",
|
|
)
|
|
if docs and docs[0].id:
|
|
return await client.get_document_by_id(docs[0].id)
|
|
|
|
docs = await client.list_documents(
|
|
limit=1,
|
|
filter=f"LOWER(title) LIKE LOWER('%{escaped_query}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')",
|
|
)
|
|
if docs and docs[0].id:
|
|
return await client.get_document_by_id(docs[0].id)
|
|
|
|
return None
|
|
|
|
|
|
def create_document_toolset(
|
|
config: AppConfig,
|
|
base_filter: str | None = None,
|
|
) -> FunctionToolset[RAGDeps]:
|
|
"""Create a toolset with document management capabilities.
|
|
|
|
Args:
|
|
config: Application configuration (used for summarization LLM).
|
|
base_filter: Optional base SQL WHERE clause applied to list operations.
|
|
|
|
Returns:
|
|
FunctionToolset with list_documents, get_document, summarize_document tools.
|
|
"""
|
|
|
|
async def list_documents(
|
|
ctx: RunContext[RAGDeps], page: int = 1
|
|
) -> DocumentListResponse:
|
|
"""List available documents in the knowledge base.
|
|
|
|
Args:
|
|
page: Page number (default: 1, 50 documents per page)
|
|
|
|
Returns:
|
|
Paginated list of documents with metadata.
|
|
"""
|
|
client = ctx.deps.client
|
|
|
|
page_size = 50
|
|
offset = (page - 1) * page_size
|
|
|
|
docs = await client.list_documents(
|
|
limit=page_size, offset=offset, filter=base_filter
|
|
)
|
|
total = await client.count_documents(filter=base_filter)
|
|
total_pages = (total + page_size - 1) // page_size if total > 0 else 1
|
|
|
|
return DocumentListResponse(
|
|
documents=[
|
|
DocumentInfo(
|
|
id=doc.id,
|
|
title=doc.title or "Untitled",
|
|
uri=doc.uri or "",
|
|
created=doc.created_at.strftime("%Y-%m-%d"),
|
|
)
|
|
for doc in docs
|
|
],
|
|
page=page,
|
|
total_pages=total_pages,
|
|
total_documents=total,
|
|
)
|
|
|
|
async def get_document(ctx: RunContext[RAGDeps], query: str) -> str:
|
|
"""Retrieve a specific document by title or URI.
|
|
|
|
Args:
|
|
query: The document title or URI to look up.
|
|
|
|
Returns:
|
|
Document content and metadata, or not found message.
|
|
"""
|
|
client = ctx.deps.client
|
|
|
|
doc = await find_document(client, query)
|
|
|
|
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}\n"
|
|
f"- Created: {doc.created_at.strftime('%Y-%m-%d %H:%M')}\n\n"
|
|
f"**Content:**\n{doc.content}"
|
|
)
|
|
|
|
async def summarize_document(ctx: RunContext[RAGDeps], query: str) -> str:
|
|
"""Generate a summary of a specific document.
|
|
|
|
Args:
|
|
query: The document title or URI to summarize.
|
|
|
|
Returns:
|
|
Generated summary or not found message.
|
|
"""
|
|
client = ctx.deps.client
|
|
|
|
doc = await find_document(client, query)
|
|
|
|
if doc is None:
|
|
return f"Document not found: {query}"
|
|
|
|
summary_model = get_model(config.qa.model, config)
|
|
summary_agent: Agent[None, str] = Agent(
|
|
summary_model,
|
|
output_type=str,
|
|
)
|
|
result = await summary_agent.run(
|
|
DOCUMENT_SUMMARY_PROMPT.format(content=doc.content or "")
|
|
)
|
|
|
|
return f"**Summary of {doc.title or doc.uri}:**\n\n{result.output}"
|
|
|
|
toolset: FunctionToolset[RAGDeps] = FunctionToolset()
|
|
toolset.add_function(list_documents)
|
|
toolset.add_function(get_document)
|
|
toolset.add_function(summarize_document)
|
|
return toolset
|