DocumentToolset, QAToolset, AnalysisToolset

This commit is contained in:
Yiorgis Gozadinos 2026-02-02 11:25:25 +02:00
parent 2f229f91e1
commit 2be3389222
No known key found for this signature in database
8 changed files with 835 additions and 1 deletions

View file

@ -10,7 +10,7 @@ from haiku.rag.agents.research.models import (
from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
from haiku.rag.config.models import AppConfig, ModelConfig
from haiku.rag.tools import ToolContext
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
from haiku.rag.utils import get_model

View file

@ -1,10 +1,24 @@
from haiku.rag.tools.analysis import (
ANALYSIS_NAMESPACE,
AnalysisState,
create_analysis_toolset,
)
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.document import (
DOCUMENT_NAMESPACE,
DocumentInfo,
DocumentListResponse,
DocumentState,
create_document_toolset,
find_document,
)
from haiku.rag.tools.filters import (
build_document_filter,
build_multi_document_filter,
combine_filters,
)
from haiku.rag.tools.models import AnalysisResult, QAResult
from haiku.rag.tools.qa import QA_NAMESPACE, QAState, create_qa_toolset
from haiku.rag.tools.search import SEARCH_NAMESPACE, SearchState, create_search_toolset
__all__ = [
@ -17,4 +31,16 @@ __all__ = [
"SEARCH_NAMESPACE",
"SearchState",
"create_search_toolset",
"DOCUMENT_NAMESPACE",
"DocumentInfo",
"DocumentListResponse",
"DocumentState",
"create_document_toolset",
"find_document",
"QA_NAMESPACE",
"QAState",
"create_qa_toolset",
"ANALYSIS_NAMESPACE",
"AnalysisState",
"create_analysis_toolset",
]

View file

@ -0,0 +1,95 @@
from pydantic import BaseModel
from pydantic_ai import FunctionToolset
from haiku.rag.agents.rlm.agent import create_rlm_agent
from haiku.rag.agents.rlm.dependencies import RLMContext, RLMDeps
from haiku.rag.agents.rlm.models import CodeExecution
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.filters import build_document_filter, combine_filters
from haiku.rag.tools.models import AnalysisResult
ANALYSIS_NAMESPACE = "haiku.rag.analysis"
class AnalysisState(BaseModel):
"""State for analysis toolset.
Tracks code executions across tool invocations.
"""
code_executions: list[CodeExecution] = []
def create_analysis_toolset(
client: HaikuRAG,
config: AppConfig,
context: ToolContext | None = None,
base_filter: str | None = None,
tool_name: str = "analyze",
) -> FunctionToolset:
"""Create a toolset with code analysis capabilities via RLM agent.
Args:
client: HaikuRAG client for document operations.
config: Application configuration.
context: Optional ToolContext for state accumulation.
If provided, code executions are tracked in AnalysisState.
base_filter: Optional base SQL WHERE clause applied to searches.
tool_name: Name for the analyze tool. Defaults to "analyze".
Returns:
FunctionToolset with an analyze tool.
"""
# Get or create analysis state if context provided
state: AnalysisState | None = None
if context is not None:
state = context.get_or_create(ANALYSIS_NAMESPACE, AnalysisState)
async def analyze(
task: str,
document_name: str | None = None,
) -> AnalysisResult:
"""Execute a computational task via code execution.
Uses the RLM (Recursive Language Model) agent to write and execute
Python code to answer the task.
Args:
task: A specific, actionable instruction describing what to compute.
document_name: Optional document name/title to focus on.
Returns:
AnalysisResult with answer and execution metadata.
"""
# Build filter from base_filter and document_name
doc_filter = build_document_filter(document_name) if document_name else None
effective_filter = combine_filters(base_filter, doc_filter)
# Create RLM context and deps
rlm_context = RLMContext(filter=effective_filter)
deps = RLMDeps(
client=client,
config=config,
context=rlm_context,
)
# Run RLM agent
rlm_agent = create_rlm_agent(config)
result = await rlm_agent.run(task, deps=deps)
# Track code executions in state
code_executions = rlm_context.code_executions
if state is not None:
state.code_executions.extend(code_executions)
return AnalysisResult(
answer=result.output.answer,
code_executed=len(code_executions) > 0,
execution_count=len(code_executions),
)
toolset = FunctionToolset()
toolset.add_function(analyze, name=tool_name)
return toolset

View file

@ -0,0 +1,208 @@
from pydantic import BaseModel
from pydantic_ai import Agent, FunctionToolset
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.tools.context import ToolContext
from haiku.rag.utils import get_model
DOCUMENT_NAMESPACE = "haiku.rag.document"
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."""
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
class DocumentState(BaseModel):
"""State for document toolset.
Tracks documents accessed during tool invocations.
"""
accessed_documents: list[DocumentInfo] = []
async def find_document(client: HaikuRAG, query: str):
"""Find a document by exact URI, partial URI, or partial title match."""
# Try exact URI match first
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(" ", "")
# Try partial URI match (with and without spaces)
docs = await client.list_documents(
limit=1,
filter=f"LOWER(uri) LIKE LOWER('%{escaped_query}%') OR LOWER(uri) LIKE LOWER('%{no_spaces}%')",
)
if docs:
return docs[0]
# Try partial title match (with and without spaces)
docs = await client.list_documents(
limit=1,
filter=f"LOWER(title) LIKE LOWER('%{escaped_query}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')",
)
if docs:
return docs[0]
return None
def create_document_toolset(
client: HaikuRAG,
config: AppConfig,
context: ToolContext | None = None,
base_filter: str | None = None,
) -> FunctionToolset:
"""Create a toolset with document management capabilities.
Args:
client: HaikuRAG client for document operations.
config: Application configuration (used for summarization LLM).
context: Optional ToolContext for state tracking.
If provided, accessed documents are tracked in DocumentState.
base_filter: Optional base SQL WHERE clause applied to list operations.
Returns:
FunctionToolset with list_documents, get_document, summarize_document tools.
"""
# Get or create document state if context provided
state: DocumentState | None = None
if context is not None:
state = context.get_or_create(DOCUMENT_NAMESPACE, DocumentState)
async def list_documents(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.
"""
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(
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(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.
"""
doc = await find_document(client, query)
if doc is None:
return f"Document not found: {query}"
# Track accessed document in state
if state is not None:
state.accessed_documents.append(
DocumentInfo(
title=doc.title or "Untitled",
uri=doc.uri or "",
created=doc.created_at.strftime("%Y-%m-%d"),
)
)
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(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.
"""
doc = await find_document(client, query)
if doc is None:
return f"Document not found: {query}"
# Track accessed document in state
if state is not None:
state.accessed_documents.append(
DocumentInfo(
title=doc.title or "Untitled",
uri=doc.uri or "",
created=doc.created_at.strftime("%Y-%m-%d"),
)
)
# Use LLM to generate summary
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()
toolset.add_function(list_documents)
toolset.add_function(get_document)
toolset.add_function(summarize_document)
return toolset

View file

@ -0,0 +1,122 @@
from pydantic import BaseModel
from pydantic_ai import FunctionToolset
from haiku.rag.agents.research.dependencies import ResearchContext
from haiku.rag.agents.research.graph import build_research_graph
from haiku.rag.agents.research.models import Citation, SearchAnswer
from haiku.rag.agents.research.state import ResearchDeps, ResearchState
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.filters import build_document_filter, combine_filters
from haiku.rag.tools.models import QAResult
QA_NAMESPACE = "haiku.rag.qa"
class QAState(BaseModel):
"""State for QA toolset.
Tracks Q&A history across tool invocations.
"""
history: list[QAResult] = []
def create_qa_toolset(
client: HaikuRAG,
config: AppConfig,
context: ToolContext | None = None,
base_filter: str | None = None,
tool_name: str = "ask",
session_context: str | None = None,
prior_answers: list[SearchAnswer] | None = None,
) -> FunctionToolset:
"""Create a toolset with Q&A capabilities using research graph.
Args:
client: HaikuRAG client for search operations.
config: Application configuration.
context: Optional ToolContext for state accumulation.
If provided, Q&A results are accumulated in QAState.
base_filter: Optional base SQL WHERE clause applied to searches.
tool_name: Name for the ask tool. Defaults to "ask".
session_context: Optional session context for the research graph.
prior_answers: Optional list of prior answers for context.
Returns:
FunctionToolset with an ask tool.
"""
# Get or create QA state if context provided
state: QAState | None = None
if context is not None:
state = context.get_or_create(QA_NAMESPACE, QAState)
async def ask(
question: str,
document_name: str | None = None,
) -> QAResult:
"""Answer a question using the knowledge base.
Uses a research graph for searching and synthesizing answers.
Args:
question: The question to answer.
document_name: Optional document name/title to search within.
Returns:
QAResult with answer, confidence, and citations.
"""
# Build filter from base_filter and document_name
doc_filter = build_document_filter(document_name) if document_name else None
effective_filter = combine_filters(base_filter, doc_filter)
# Build and run the research graph
graph = build_research_graph(config=config, output_mode="conversational")
research_context = ResearchContext(
original_question=question,
session_context=session_context,
qa_responses=prior_answers or [],
)
research_state = ResearchState(
context=research_context,
max_iterations=1,
search_filter=effective_filter,
max_concurrency=config.research.max_concurrency,
)
deps = ResearchDeps(client=client)
result = await graph.run(state=research_state, deps=deps)
# Convert to QAResult
citations = [
Citation(
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(result.citations)
]
qa_result = QAResult(
question=question,
answer=result.answer,
confidence=result.confidence,
citations=citations,
)
# Accumulate in state if context provided
if state is not None:
state.history.append(qa_result)
return qa_result
toolset = FunctionToolset()
toolset.add_function(ask, name=tool_name)
return toolset

View file

@ -0,0 +1,94 @@
import pytest
from haiku.rag.tools import ToolContext
from haiku.rag.tools.analysis import (
ANALYSIS_NAMESPACE,
AnalysisState,
create_analysis_toolset,
)
class TestAnalysisState:
"""Tests for AnalysisState model."""
def test_analysis_state_defaults(self):
"""AnalysisState initializes with empty code_executions."""
state = AnalysisState()
assert state.code_executions == []
def test_analysis_state_serialization(self):
"""AnalysisState serializes and deserializes correctly."""
from haiku.rag.agents.rlm.models import CodeExecution
state = AnalysisState()
state.code_executions.append(
CodeExecution(
code="print('hello')",
stdout="hello\n",
stderr="",
success=True,
)
)
data = state.model_dump()
restored = AnalysisState.model_validate(data)
assert len(restored.code_executions) == 1
assert restored.code_executions[0].code == "print('hello')"
class TestAnalysisToolset:
"""Tests for create_analysis_toolset."""
def test_create_analysis_toolset_returns_function_toolset(
self, analysis_client, analysis_config
):
"""create_analysis_toolset returns a FunctionToolset."""
from pydantic_ai import FunctionToolset
toolset = create_analysis_toolset(analysis_client, analysis_config)
assert isinstance(toolset, FunctionToolset)
def test_analysis_toolset_has_analyze_tool(self, analysis_client, analysis_config):
"""The toolset includes an 'analyze' tool."""
toolset = create_analysis_toolset(analysis_client, analysis_config)
assert "analyze" in toolset.tools
def test_analysis_toolset_registers_state(self, analysis_client, analysis_config):
"""Toolset registers AnalysisState under ANALYSIS_NAMESPACE."""
context = ToolContext()
create_analysis_toolset(analysis_client, analysis_config, context=context)
state = context.get(ANALYSIS_NAMESPACE)
assert state is not None
assert isinstance(state, AnalysisState)
def test_analysis_toolset_custom_tool_name(self, analysis_client, analysis_config):
"""Toolset supports custom tool name."""
toolset = create_analysis_toolset(
analysis_client, analysis_config, tool_name="run_code"
)
assert "run_code" in toolset.tools
assert "analyze" not in toolset.tools
@pytest.fixture
def analysis_client(temp_db_path):
"""Create a HaikuRAG client for analysis tests."""
import asyncio
from haiku.rag.client import HaikuRAG
async def setup():
rag = HaikuRAG(temp_db_path, create=True)
await rag.__aenter__()
return rag
return asyncio.get_event_loop().run_until_complete(setup())
@pytest.fixture
def analysis_config():
"""Default AppConfig for analysis tests."""
from haiku.rag.config import Config
return Config

View file

@ -0,0 +1,194 @@
import pytest
from haiku.rag.tools import ToolContext
from haiku.rag.tools.document import (
DOCUMENT_NAMESPACE,
DocumentInfo,
DocumentListResponse,
DocumentState,
create_document_toolset,
)
class TestDocumentModels:
"""Tests for document models."""
def test_document_info(self):
"""DocumentInfo holds basic document metadata."""
info = DocumentInfo(title="Test Doc", uri="test://doc", created="2024-01-01")
assert info.title == "Test Doc"
assert info.uri == "test://doc"
assert info.created == "2024-01-01"
def test_document_list_response(self):
"""DocumentListResponse holds paginated results."""
response = DocumentListResponse(
documents=[
DocumentInfo(title="Doc 1", uri="test://1", created="2024-01-01"),
DocumentInfo(title="Doc 2", uri="test://2", created="2024-01-02"),
],
page=1,
total_pages=3,
total_documents=125,
)
assert len(response.documents) == 2
assert response.page == 1
assert response.total_pages == 3
assert response.total_documents == 125
def test_document_state_defaults(self):
"""DocumentState initializes with empty accessed list."""
state = DocumentState()
assert state.accessed_documents == []
class TestDocumentToolset:
"""Tests for create_document_toolset."""
def test_create_document_toolset_returns_function_toolset(
self, doc_client, doc_config
):
"""create_document_toolset returns a FunctionToolset."""
from pydantic_ai import FunctionToolset
toolset = create_document_toolset(doc_client, doc_config)
assert isinstance(toolset, FunctionToolset)
def test_document_toolset_has_expected_tools(self, doc_client, doc_config):
"""The toolset includes list_documents, get_document, summarize_document."""
toolset = create_document_toolset(doc_client, doc_config)
assert "list_documents" in toolset.tools
assert "get_document" in toolset.tools
assert "summarize_document" in toolset.tools
def test_document_toolset_registers_state(self, doc_client, doc_config):
"""Toolset registers DocumentState under DOCUMENT_NAMESPACE."""
context = ToolContext()
create_document_toolset(doc_client, doc_config, context=context)
state = context.get(DOCUMENT_NAMESPACE)
assert state is not None
assert isinstance(state, DocumentState)
class TestDocumentToolExecution:
"""Tests for document tool execution."""
@pytest.mark.asyncio
async def test_list_documents_returns_paginated_results(
self, doc_client, doc_config
):
"""list_documents returns DocumentListResponse."""
toolset = create_document_toolset(doc_client, doc_config)
list_tool = toolset.tools["list_documents"]
result = await list_tool.function()
assert isinstance(result, DocumentListResponse)
assert result.total_documents == 2
assert len(result.documents) == 2
assert result.page == 1
@pytest.mark.asyncio
async def test_list_documents_pagination(self, doc_client, doc_config):
"""list_documents supports pagination."""
toolset = create_document_toolset(doc_client, doc_config)
list_tool = toolset.tools["list_documents"]
result = await list_tool.function(page=2)
# With only 2 documents and page_size=50, page 2 should be empty
assert result.page == 2
assert len(result.documents) == 0
@pytest.mark.asyncio
async def test_get_document_by_title(self, doc_client, doc_config):
"""get_document finds document by title."""
toolset = create_document_toolset(doc_client, doc_config)
get_tool = toolset.tools["get_document"]
result = await get_tool.function("Python Guide")
assert "Python Guide" in result
assert "Python is a programming language" in result
@pytest.mark.asyncio
async def test_get_document_by_uri(self, doc_client, doc_config):
"""get_document finds document by URI."""
toolset = create_document_toolset(doc_client, doc_config)
get_tool = toolset.tools["get_document"]
result = await get_tool.function("test://python")
assert "Python Guide" in result
@pytest.mark.asyncio
async def test_get_document_not_found(self, doc_client, doc_config):
"""get_document returns appropriate message when not found."""
toolset = create_document_toolset(doc_client, doc_config)
get_tool = toolset.tools["get_document"]
result = await get_tool.function("nonexistent")
assert "Document not found" in result
@pytest.mark.asyncio
async def test_get_document_tracks_in_state(self, doc_client, doc_config):
"""get_document tracks accessed documents in state."""
context = ToolContext()
toolset = create_document_toolset(doc_client, doc_config, context=context)
get_tool = toolset.tools["get_document"]
await get_tool.function("Python Guide")
state = context.get(DOCUMENT_NAMESPACE)
assert isinstance(state, DocumentState)
assert len(state.accessed_documents) == 1
assert state.accessed_documents[0].title == "Python Guide"
@pytest.mark.asyncio
async def test_list_documents_with_base_filter(self, doc_client, doc_config):
"""list_documents respects base_filter."""
toolset = create_document_toolset(
doc_client, doc_config, base_filter="title LIKE '%Python%'"
)
list_tool = toolset.tools["list_documents"]
result = await list_tool.function()
assert result.total_documents == 1
assert result.documents[0].title == "Python Guide"
@pytest.fixture
def doc_client(temp_db_path):
"""Create a HaikuRAG client with test documents."""
import asyncio
from haiku.rag.client import HaikuRAG
async def setup():
rag = HaikuRAG(temp_db_path, create=True)
await rag.__aenter__()
await rag.create_document(
"Python is a programming language. It is widely used for web development.",
uri="test://python",
title="Python Guide",
)
await rag.create_document(
"JavaScript runs in the browser. It powers interactive web pages.",
uri="test://javascript",
title="JavaScript Guide",
)
return rag
return asyncio.get_event_loop().run_until_complete(setup())
@pytest.fixture
def doc_config():
"""Default AppConfig for document tests."""
from haiku.rag.config import Config
return Config

95
tests/tools/test_qa.py Normal file
View file

@ -0,0 +1,95 @@
import pytest
from haiku.rag.tools import QAResult, ToolContext
from haiku.rag.tools.qa import QA_NAMESPACE, QAState, create_qa_toolset
class TestQAState:
"""Tests for QAState model."""
def test_qa_state_defaults(self):
"""QAState initializes with empty history."""
state = QAState()
assert state.history == []
def test_qa_state_add_result(self):
"""Can add QAResult to history."""
state = QAState()
result = QAResult(question="What is Python?", answer="A programming language.")
state.history.append(result)
assert len(state.history) == 1
assert state.history[0].question == "What is Python?"
def test_qa_state_serialization(self):
"""QAState serializes and deserializes correctly."""
state = QAState()
state.history.append(
QAResult(
question="Test?",
answer="Answer.",
confidence=0.95,
)
)
data = state.model_dump()
restored = QAState.model_validate(data)
assert len(restored.history) == 1
assert restored.history[0].confidence == 0.95
class TestQAToolset:
"""Tests for create_qa_toolset."""
def test_create_qa_toolset_returns_function_toolset(
self, qa_client_simple, qa_config
):
"""create_qa_toolset returns a FunctionToolset."""
from pydantic_ai import FunctionToolset
toolset = create_qa_toolset(qa_client_simple, qa_config)
assert isinstance(toolset, FunctionToolset)
def test_qa_toolset_has_ask_tool(self, qa_client_simple, qa_config):
"""The toolset includes an 'ask' tool."""
toolset = create_qa_toolset(qa_client_simple, qa_config)
assert "ask" in toolset.tools
def test_qa_toolset_registers_state(self, qa_client_simple, qa_config):
"""Toolset registers QAState under QA_NAMESPACE."""
context = ToolContext()
create_qa_toolset(qa_client_simple, qa_config, context=context)
state = context.get(QA_NAMESPACE)
assert state is not None
assert isinstance(state, QAState)
def test_qa_toolset_custom_tool_name(self, qa_client_simple, qa_config):
"""Toolset supports custom tool name."""
toolset = create_qa_toolset(
qa_client_simple, qa_config, tool_name="answer_question"
)
assert "answer_question" in toolset.tools
assert "ask" not in toolset.tools
@pytest.fixture
def qa_client_simple(temp_db_path):
"""Create a HaikuRAG client without documents for basic tests."""
import asyncio
from haiku.rag.client import HaikuRAG
async def setup():
rag = HaikuRAG(temp_db_path, create=True)
await rag.__aenter__()
return rag
return asyncio.get_event_loop().run_until_complete(setup())
@pytest.fixture
def qa_config():
"""Default AppConfig for QA tests."""
from haiku.rag.config import Config
return Config