flatten skill architecture: replace ask/analyze/research with direct tools
This commit is contained in:
parent
20e40d75f9
commit
d52f453c44
14 changed files with 375 additions and 892 deletions
|
|
@ -21,8 +21,8 @@ import { FilterIcon } from "../lib/icons";
|
||||||
import type { RAGState } from "../lib/sessionStorage";
|
import type { RAGState } from "../lib/sessionStorage";
|
||||||
import {
|
import {
|
||||||
createSession,
|
createSession,
|
||||||
deriveCitationsHistory,
|
|
||||||
getActiveSessionId,
|
getActiveSessionId,
|
||||||
|
getLatestCitations,
|
||||||
getSession,
|
getSession,
|
||||||
normalizeRAGState,
|
normalizeRAGState,
|
||||||
updateSessionMessages,
|
updateSessionMessages,
|
||||||
|
|
@ -292,7 +292,7 @@ function MessageViewWithCitations({
|
||||||
isRunning?: boolean;
|
isRunning?: boolean;
|
||||||
}) {
|
}) {
|
||||||
const ragState = useContext(ChatStateContext);
|
const ragState = useContext(ChatStateContext);
|
||||||
const citationsHistory = ragState ? deriveCitationsHistory(ragState) : [];
|
const latestCitations = ragState ? getLatestCitations(ragState) : [];
|
||||||
|
|
||||||
// Collect completed tool_call_ids from skill_tool_result activity messages
|
// Collect completed tool_call_ids from skill_tool_result activity messages
|
||||||
const completedToolCallIds = useMemo(() => {
|
const completedToolCallIds = useMemo(() => {
|
||||||
|
|
@ -326,7 +326,6 @@ function MessageViewWithCitations({
|
||||||
{({ messageElements }) => {
|
{({ messageElements }) => {
|
||||||
const result: React.ReactNode[] = [];
|
const result: React.ReactNode[] = [];
|
||||||
let elemIdx = 0;
|
let elemIdx = 0;
|
||||||
let citIdx = 0;
|
|
||||||
let seenToolCalls = false;
|
let seenToolCalls = false;
|
||||||
|
|
||||||
for (const msg of messages) {
|
for (const msg of messages) {
|
||||||
|
|
@ -368,19 +367,15 @@ function MessageViewWithCitations({
|
||||||
}
|
}
|
||||||
|
|
||||||
// After an assistant text response that followed tool calls,
|
// After an assistant text response that followed tool calls,
|
||||||
// inject the next citations entry (one per turn)
|
// show citations from the latest turn
|
||||||
if (msg.role === "assistant" && msg.content && seenToolCalls) {
|
if (msg.role === "assistant" && msg.content && seenToolCalls) {
|
||||||
if (citIdx < citationsHistory.length) {
|
if (latestCitations.length > 0) {
|
||||||
const citations = citationsHistory[citIdx];
|
result.push(
|
||||||
if (citations?.length) {
|
<CitationBlock
|
||||||
result.push(
|
key={`citations-${i}`}
|
||||||
<CitationBlock
|
citations={latestCitations}
|
||||||
key={`citations-${citIdx}`}
|
/>,
|
||||||
citations={citations}
|
);
|
||||||
/>,
|
|
||||||
);
|
|
||||||
}
|
|
||||||
citIdx++;
|
|
||||||
}
|
}
|
||||||
seenToolCalls = false;
|
seenToolCalls = false;
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -9,12 +9,6 @@ export interface Citation {
|
||||||
content: string;
|
content: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface QAHistoryEntry {
|
|
||||||
question: string;
|
|
||||||
answer: string;
|
|
||||||
citations: Citation[];
|
|
||||||
}
|
|
||||||
|
|
||||||
export interface DocumentInfo {
|
export interface DocumentInfo {
|
||||||
id: string;
|
id: string;
|
||||||
title: string;
|
title: string;
|
||||||
|
|
@ -22,20 +16,13 @@ export interface DocumentInfo {
|
||||||
created: string;
|
created: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface ResearchEntry {
|
|
||||||
question: string;
|
|
||||||
title: string;
|
|
||||||
executive_summary: string;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Matches RAGState from the backend skill
|
// Matches RAGState from the backend skill
|
||||||
export interface RAGState {
|
export interface RAGState {
|
||||||
citations: Citation[];
|
citation_index: Record<string, Citation>;
|
||||||
qa_history: QAHistoryEntry[];
|
citations: string[][];
|
||||||
document_filter: string | null;
|
document_filter: string | null;
|
||||||
searches: Record<string, unknown[]>;
|
searches: Record<string, unknown[]>;
|
||||||
documents: DocumentInfo[];
|
documents: DocumentInfo[];
|
||||||
reports: ResearchEntry[];
|
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface StoredMessage {
|
export interface StoredMessage {
|
||||||
|
|
@ -59,20 +46,21 @@ const ACTIVE_SESSION_KEY = "haiku.rag.activeSession";
|
||||||
|
|
||||||
export function normalizeRAGState(state?: Partial<RAGState>): RAGState {
|
export function normalizeRAGState(state?: Partial<RAGState>): RAGState {
|
||||||
return {
|
return {
|
||||||
|
citation_index: state?.citation_index ?? {},
|
||||||
citations: state?.citations ?? [],
|
citations: state?.citations ?? [],
|
||||||
qa_history: state?.qa_history ?? [],
|
|
||||||
document_filter: state?.document_filter ?? null,
|
document_filter: state?.document_filter ?? null,
|
||||||
searches: state?.searches ?? {},
|
searches: state?.searches ?? {},
|
||||||
documents: state?.documents ?? [],
|
documents: state?.documents ?? [],
|
||||||
reports: state?.reports ?? [],
|
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
// Derive per-turn citation arrays from qa_history
|
export function getLatestCitations(state: RAGState): Citation[] {
|
||||||
export function deriveCitationsHistory(state: RAGState): Citation[][] {
|
const turns = state.citations;
|
||||||
return state.qa_history
|
if (turns.length === 0) return [];
|
||||||
.filter((entry) => entry.citations?.length > 0)
|
const latestIds = turns[turns.length - 1];
|
||||||
.map((entry) => entry.citations);
|
return latestIds
|
||||||
|
.map((id) => state.citation_index[id])
|
||||||
|
.filter((c): c is Citation => c !== undefined);
|
||||||
}
|
}
|
||||||
|
|
||||||
export function getAllSessions(): StoredSession[] {
|
export function getAllSessions(): StoredSession[] {
|
||||||
|
|
|
||||||
|
|
@ -243,8 +243,7 @@ class ChatApp(App):
|
||||||
content=accumulated_text,
|
content=accumulated_text,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
# Show citations from RAG state
|
await self._show_citations_and_programs(chat_history)
|
||||||
await self._show_citations(chat_history)
|
|
||||||
elif event.type == EventType.TOOL_CALL_START:
|
elif event.type == EventType.TOOL_CALL_START:
|
||||||
assert isinstance(event, ToolCallStartEvent)
|
assert isinstance(event, ToolCallStartEvent)
|
||||||
chat_history.hide_thinking()
|
chat_history.hide_thinking()
|
||||||
|
|
@ -320,18 +319,32 @@ class ChatApp(App):
|
||||||
chat_input.disabled = False
|
chat_input.disabled = False
|
||||||
chat_input.focus()
|
chat_input.focus()
|
||||||
|
|
||||||
async def _show_citations(self, chat_history: "ChatHistory") -> None:
|
async def _show_citations_and_programs(self, chat_history: "ChatHistory") -> None:
|
||||||
"""Show citations from skill states after an agent response."""
|
"""Show citations and programs from skill states after an agent response."""
|
||||||
if not self._toolset:
|
if not self._toolset:
|
||||||
return
|
return
|
||||||
citations = []
|
citations = []
|
||||||
for namespace in (RAG_STATE_NAMESPACE, ANALYSIS_STATE_NAMESPACE):
|
for namespace in (RAG_STATE_NAMESPACE, ANALYSIS_STATE_NAMESPACE):
|
||||||
state = self._toolset.get_namespace(namespace)
|
state = self._toolset.get_namespace(namespace)
|
||||||
if state:
|
if not state:
|
||||||
citations.extend(getattr(state, "citations", []))
|
continue
|
||||||
|
citation_turns = getattr(state, "citations", [])
|
||||||
|
citation_index = getattr(state, "citation_index", {})
|
||||||
|
if citation_turns:
|
||||||
|
latest_ids = citation_turns[-1]
|
||||||
|
for cid in latest_ids:
|
||||||
|
if cid in citation_index:
|
||||||
|
citations.append(citation_index[cid])
|
||||||
if citations:
|
if citations:
|
||||||
await chat_history.add_citations(citations)
|
await chat_history.add_citations(citations)
|
||||||
|
|
||||||
|
analysis_state = self._toolset.get_namespace(ANALYSIS_STATE_NAMESPACE)
|
||||||
|
if analysis_state:
|
||||||
|
executions = getattr(analysis_state, "executions", [])
|
||||||
|
successful = [e for e in executions if e.success]
|
||||||
|
if successful:
|
||||||
|
await chat_history.add_program(successful[-1].code)
|
||||||
|
|
||||||
async def action_clear_chat(self) -> None:
|
async def action_clear_chat(self) -> None:
|
||||||
"""Clear the chat history and reset session."""
|
"""Clear the chat history and reset session."""
|
||||||
chat_history = self.query_one(ChatHistory)
|
chat_history = self.query_one(ChatHistory)
|
||||||
|
|
|
||||||
|
|
@ -8,9 +8,8 @@ AVAILABLE_TOOLS: set[str] = {
|
||||||
"list_documents",
|
"list_documents",
|
||||||
"get_document",
|
"get_document",
|
||||||
"search",
|
"search",
|
||||||
"ask",
|
"execute_code",
|
||||||
"research",
|
"cite",
|
||||||
"analyze",
|
|
||||||
}
|
}
|
||||||
|
|
||||||
DEFAULT_PREAMBLE = (
|
DEFAULT_PREAMBLE = (
|
||||||
|
|
|
||||||
|
|
@ -7,48 +7,58 @@ description: {{ description }}
|
||||||
|
|
||||||
{{ preamble }}
|
{{ preamble }}
|
||||||
|
|
||||||
## How to decide which tool to use
|
## Tools
|
||||||
{% if "ask" in tool_names %}
|
{% if "search" in tool_names %}
|
||||||
|
|
||||||
**Default rule:** If the user is asking a question, use **ask**. Only use **search** when the user explicitly wants to browse or find passages.
|
### search
|
||||||
|
Search the knowledge base using hybrid search (vector + full-text). Returns ranked results with context-expanded content. Use for answering questions, finding passages, exploring topics.
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "list_documents" in tool_names %}
|
{% if "list_documents" in tool_names %}
|
||||||
- **list_documents** — Use when the user wants to browse or see what documents are available (e.g., "what documents do you have?", "show me the documents", "list available docs").
|
|
||||||
|
### list_documents
|
||||||
|
List all documents in the knowledge base.
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "get_document" in tool_names %}
|
{% if "get_document" in tool_names %}
|
||||||
- **get_document** — Use when the user wants the full content of a specific document (e.g., "get the paper about X", "show me document Y"). Accepts a document ID, title, or URI — partial matches work.
|
|
||||||
|
### get_document
|
||||||
|
Retrieve a document by ID, title, or URI. Partial matches work.
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "search" in tool_names %}
|
{% if "execute_code" in tool_names %}
|
||||||
- **search** — Use when the user wants to browse, explore, or find specific passages across documents (e.g., "search for embeddings", "find mentions of transformers"). Returns all matching results as sources.
|
|
||||||
|
### execute_code
|
||||||
|
Execute Python code in a sandboxed interpreter. Inside the code you have access to `await search()`, `await list_documents()`, `await llm()`, and a virtual filesystem at `/documents/` with document content and structure.
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "ask" in tool_names %}
|
{% if "cite" in tool_names %}
|
||||||
- **ask** — Use for factual questions that need a synthesized answer (e.g., "what is DocLayNet?", "explain the methodology"). Searches, synthesizes, and returns only the chunks actually used as citations. Always include the citations in your response.
|
|
||||||
{% endif %}
|
### cite
|
||||||
{% if "research" in tool_names %}
|
Register chunk IDs as citations. Call after formulating your answer with chunk_id values from search results that support it. Do NOT include chunk IDs in your answer text.
|
||||||
- **research** — Deep multi-agent research that produces comprehensive reports. **Only use when the user explicitly requests deep research** (e.g., "do a deep research on X", "research this topic thoroughly"). Never call this tool on your own — it is slow and expensive.
|
|
||||||
{% endif %}
|
|
||||||
{% if "analyze" in tool_names %}
|
|
||||||
- **analyze** — Use for complex analytical questions that require computation, aggregation, or data traversal across documents (e.g., "how many pages?", "compare table 3 across documents", "calculate average word count"). Executes Python code in a sandboxed interpreter.
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "search" in tool_names %}
|
{% if "search" in tool_names %}
|
||||||
|
|
||||||
## When search returns irrelevant results
|
## How to answer questions
|
||||||
|
|
||||||
If your first search returns results that clearly don't match the question, **do not keep searching with variations**. Instead:
|
1. Call `search` with relevant keywords from the question
|
||||||
{% if "ask" in tool_names %}
|
2. Review results — they are ordered by relevance (rank 1 = best match)
|
||||||
- Use **ask** if the question is factual
|
3. If needed, search again with different keywords (up to 3-4 searches total)
|
||||||
|
4. Synthesize a concise answer based strictly on the retrieved content
|
||||||
|
{% if "cite" in tool_names %}
|
||||||
|
5. Call `cite` with the chunk IDs you referenced
|
||||||
|
{% endif %}
|
||||||
|
|
||||||
|
## Guidelines
|
||||||
|
|
||||||
|
- Base answers strictly on retrieved content — do not use external knowledge
|
||||||
|
- Be concise and direct — avoid elaboration unless asked
|
||||||
|
- If results don't match the question, report that the knowledge base lacks the information
|
||||||
|
{% if "cite" in tool_names %}
|
||||||
|
- Do NOT include chunk IDs or UUIDs in your answer text — use the `cite` tool separately
|
||||||
{% endif %}
|
{% endif %}
|
||||||
- Report that the knowledge base doesn't contain relevant information
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "get_document" in tool_names %}
|
{% if "get_document" in tool_names %}
|
||||||
|
|
||||||
## When the user mentions a specific document
|
## When the user mentions a specific document
|
||||||
|
|
||||||
If the user says "search in [doc]", "find in [doc]", or "answer from [doc]":
|
If the user says "search in [doc]", "find in [doc]", or "answer from [doc]":
|
||||||
- Extract the **topic** as the `query`/`question` parameter
|
- Use **get_document** or **list_documents** first to identify the document
|
||||||
- Use **get_document** or **list_documents** first to identify the document, then search/ask with a filter
|
- Then search for the topic
|
||||||
|
|
||||||
Examples:
|
|
||||||
- "search for embeddings in the ML paper" -> first identify "ML paper", then search for "embeddings"
|
|
||||||
- "what does the DocLayNet paper say about annotations?" -> ask with question="what are the annotation methods?"
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
|
|
|
||||||
|
|
@ -5,23 +5,17 @@ from pydantic import BaseModel, Field
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.skills.models import Skill
|
from haiku.skills.models import Skill
|
||||||
from haiku.skills.parser import parse_skill_md
|
from haiku.skills.parser import parse_skill_md
|
||||||
{% if "ask" in tool_names or "research" in tool_names %}
|
{% if "cite" in tool_names %}
|
||||||
from haiku.rag.agents.research.models import Citation
|
from haiku.rag.agents.research.models import Citation
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "search" in tool_names or "list_documents" in tool_names or "get_document" in tool_names %}
|
{% if "search" in tool_names or "list_documents" in tool_names or "get_document" in tool_names %}
|
||||||
from haiku.rag.tools.document import DocumentInfo
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "ask" in tool_names %}
|
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
{% endif %}
|
|
||||||
{% if "search" in tool_names %}
|
{% if "search" in tool_names %}
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "research" in tool_names %}
|
{% if "execute_code" in tool_names %}
|
||||||
from haiku.rag.skills._tools import ResearchEntry
|
from haiku.rag.skills._tools import CodeExecutionEntry
|
||||||
{% endif %}
|
|
||||||
{% if "analyze" in tool_names %}
|
|
||||||
from haiku.rag.skills._tools import AnalysisEntry
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
|
|
||||||
_TOOL_NAMES = {{ tool_names | tojson }}
|
_TOOL_NAMES = {{ tool_names | tojson }}
|
||||||
|
|
@ -36,11 +30,9 @@ _CONFIG_PATH = _ASSETS_DIR / "haiku.rag.yaml"
|
||||||
|
|
||||||
|
|
||||||
class SkillState(BaseModel):
|
class SkillState(BaseModel):
|
||||||
{% if "ask" in tool_names or "research" in tool_names %}
|
{% if "cite" in tool_names %}
|
||||||
citations: list[Citation] = Field(default_factory=list)
|
citation_index: dict[str, Citation] = Field(default_factory=dict)
|
||||||
{% endif %}
|
citations: list[list[str]] = Field(default_factory=list)
|
||||||
{% if "ask" in tool_names %}
|
|
||||||
qa_history: list[QAHistoryEntry] = Field(default_factory=list)
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
document_filter: str | None = None
|
document_filter: str | None = None
|
||||||
{% if "search" in tool_names %}
|
{% if "search" in tool_names %}
|
||||||
|
|
@ -49,11 +41,8 @@ class SkillState(BaseModel):
|
||||||
{% if "search" in tool_names or "list_documents" in tool_names or "get_document" in tool_names %}
|
{% if "search" in tool_names or "list_documents" in tool_names or "get_document" in tool_names %}
|
||||||
documents: list[DocumentInfo] = Field(default_factory=list)
|
documents: list[DocumentInfo] = Field(default_factory=list)
|
||||||
{% endif %}
|
{% endif %}
|
||||||
{% if "research" in tool_names %}
|
{% if "execute_code" in tool_names %}
|
||||||
reports: list[ResearchEntry] = Field(default_factory=list)
|
executions: list[CodeExecutionEntry] = Field(default_factory=list)
|
||||||
{% endif %}
|
|
||||||
{% if "analyze" in tool_names %}
|
|
||||||
analyses: list[AnalysisEntry] = Field(default_factory=list)
|
|
||||||
{% endif %}
|
{% endif %}
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -8,57 +8,14 @@ from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.tools.document import DocumentInfo
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
from haiku.skills.state import SkillRunDeps
|
from haiku.skills.state import SkillRunDeps
|
||||||
|
|
||||||
|
|
||||||
class ResearchEntry(BaseModel):
|
class CodeExecutionEntry(BaseModel):
|
||||||
question: str
|
code: str
|
||||||
title: str
|
stdout: str
|
||||||
executive_summary: str
|
stderr: str = ""
|
||||||
|
success: bool = True
|
||||||
|
|
||||||
class AnalysisEntry(BaseModel):
|
|
||||||
question: str
|
|
||||||
answer: str
|
|
||||||
program: str | None = None
|
|
||||||
|
|
||||||
|
|
||||||
async def find_relevant_prior_qa(
|
|
||||||
qa_history: list[QAHistoryEntry],
|
|
||||||
query: str,
|
|
||||||
config: AppConfig,
|
|
||||||
) -> list[QAHistoryEntry]:
|
|
||||||
from haiku.rag.embeddings import get_embedder
|
|
||||||
from haiku.rag.tools.qa import PRIOR_ANSWER_RELEVANCE_THRESHOLD
|
|
||||||
from haiku.rag.utils import cosine_similarity
|
|
||||||
|
|
||||||
if not qa_history:
|
|
||||||
return []
|
|
||||||
|
|
||||||
embedder = get_embedder(config)
|
|
||||||
query_embedding = await embedder.embed_query(query)
|
|
||||||
|
|
||||||
to_embed = []
|
|
||||||
to_embed_indices = []
|
|
||||||
for i, qa in enumerate(qa_history):
|
|
||||||
if qa.question_embedding is None:
|
|
||||||
to_embed.append(qa.question)
|
|
||||||
to_embed_indices.append(i)
|
|
||||||
|
|
||||||
if to_embed:
|
|
||||||
new_embeddings = await embedder.embed_documents(to_embed)
|
|
||||||
for i, idx in enumerate(to_embed_indices):
|
|
||||||
qa_history[idx].question_embedding = new_embeddings[i]
|
|
||||||
|
|
||||||
matches = []
|
|
||||||
for qa in qa_history:
|
|
||||||
if qa.question_embedding is not None:
|
|
||||||
similarity = cosine_similarity(query_embedding, qa.question_embedding)
|
|
||||||
if similarity >= PRIOR_ANSWER_RELEVANCE_THRESHOLD:
|
|
||||||
matches.append(qa)
|
|
||||||
|
|
||||||
return matches
|
|
||||||
|
|
||||||
|
|
||||||
async def skill_search(
|
async def skill_search(
|
||||||
|
|
@ -88,14 +45,12 @@ async def skill_search(
|
||||||
async def skill_list_documents(
|
async def skill_list_documents(
|
||||||
db_path: Path,
|
db_path: Path,
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
limit: int | None = None,
|
|
||||||
offset: int | None = None,
|
|
||||||
filter: str | None = None,
|
filter: str | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
||||||
documents = await rag.list_documents(limit, offset, filter=filter)
|
documents = await rag.list_documents(filter=filter)
|
||||||
return [
|
return [
|
||||||
{
|
{
|
||||||
"id": doc.id,
|
"id": doc.id,
|
||||||
|
|
@ -131,100 +86,6 @@ async def skill_get_document(
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def skill_ask(
|
|
||||||
db_path: Path,
|
|
||||||
config: AppConfig,
|
|
||||||
question: str,
|
|
||||||
qa_history: list[QAHistoryEntry] | None = None,
|
|
||||||
document_filter: str | None = None,
|
|
||||||
) -> tuple[str, list[Citation]]:
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
from haiku.rag.utils import format_citations
|
|
||||||
|
|
||||||
ask_question = question
|
|
||||||
if qa_history:
|
|
||||||
matches = await find_relevant_prior_qa(qa_history, question, config)
|
|
||||||
if matches:
|
|
||||||
prior_parts = []
|
|
||||||
for qa in matches:
|
|
||||||
part = f"Q: {qa.question}\nA: {qa.answer}"
|
|
||||||
if qa.citations:
|
|
||||||
part += "\n" + format_citations(qa.citations)
|
|
||||||
prior_parts.append(part)
|
|
||||||
ask_question = (
|
|
||||||
"Context from prior questions in this session:\n\n"
|
|
||||||
+ "\n\n---\n\n".join(prior_parts)
|
|
||||||
+ "\n\n---\n\nCurrent question: "
|
|
||||||
+ question
|
|
||||||
)
|
|
||||||
|
|
||||||
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
|
||||||
answer, citations = await rag.ask(
|
|
||||||
ask_question,
|
|
||||||
filter=document_filter,
|
|
||||||
)
|
|
||||||
|
|
||||||
return answer, citations
|
|
||||||
|
|
||||||
|
|
||||||
async def skill_research(
|
|
||||||
db_path: Path,
|
|
||||||
config: AppConfig,
|
|
||||||
question: str,
|
|
||||||
document_filter: str | None = None,
|
|
||||||
) -> tuple[str, str, str]:
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
|
|
||||||
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
|
||||||
report = await rag.research(question, filter=document_filter)
|
|
||||||
|
|
||||||
parts = [
|
|
||||||
f"# {report.title}",
|
|
||||||
f"\n## Executive Summary\n{report.executive_summary}",
|
|
||||||
]
|
|
||||||
if report.main_findings:
|
|
||||||
parts.append("\n## Main Findings")
|
|
||||||
for finding in report.main_findings:
|
|
||||||
parts.append(f"- {finding}")
|
|
||||||
if report.conclusions:
|
|
||||||
parts.append("\n## Conclusions")
|
|
||||||
for conclusion in report.conclusions:
|
|
||||||
parts.append(f"- {conclusion}")
|
|
||||||
if report.limitations:
|
|
||||||
parts.append("\n## Limitations")
|
|
||||||
for limitation in report.limitations:
|
|
||||||
parts.append(f"- {limitation}")
|
|
||||||
if report.recommendations:
|
|
||||||
parts.append("\n## Recommendations")
|
|
||||||
for rec in report.recommendations:
|
|
||||||
parts.append(f"- {rec}")
|
|
||||||
parts.append(f"\n## Sources\n{report.sources_summary}")
|
|
||||||
|
|
||||||
formatted = "\n".join(parts)
|
|
||||||
return formatted, report.title, report.executive_summary
|
|
||||||
|
|
||||||
|
|
||||||
async def skill_analyze(
|
|
||||||
db_path: Path,
|
|
||||||
config: AppConfig,
|
|
||||||
question: str,
|
|
||||||
document: str | None = None,
|
|
||||||
document_filter: str | None = None,
|
|
||||||
) -> tuple[str, str, str | None, "list[Citation]"]:
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
|
|
||||||
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
|
||||||
documents = [document] if document else None
|
|
||||||
result = await rag.analyze(
|
|
||||||
question, documents=documents, filter=document_filter
|
|
||||||
)
|
|
||||||
output = result.answer
|
|
||||||
if result.program:
|
|
||||||
output += f"\n\nProgram:\n{result.program}"
|
|
||||||
|
|
||||||
return output, result.answer, result.program, result.citations
|
|
||||||
|
|
||||||
|
|
||||||
def update_documents_state(
|
def update_documents_state(
|
||||||
documents_state: list[DocumentInfo],
|
documents_state: list[DocumentInfo],
|
||||||
doc_dicts: list[dict[str, Any]],
|
doc_dicts: list[dict[str, Any]],
|
||||||
|
|
@ -246,13 +107,18 @@ def _get_state(ctx: RunContext[SkillRunDeps], state_type: type[BaseModel]) -> An
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _append_citations(state: Any, citations: "list[Citation]") -> None:
|
def _register_citations(state: Any, citations: "list[Citation]") -> None:
|
||||||
"""Index and append citations to a skill state's citations list."""
|
"""Add citations to the index and record the turn's chunk IDs."""
|
||||||
next_index = len(state.citations) + 1
|
chunk_ids = []
|
||||||
|
next_index = len(state.citation_index) + 1
|
||||||
for citation in citations:
|
for citation in citations:
|
||||||
citation.index = next_index
|
cid = citation.chunk_id
|
||||||
next_index += 1
|
if cid not in state.citation_index:
|
||||||
state.citations.extend(citations)
|
citation.index = next_index
|
||||||
|
next_index += 1
|
||||||
|
state.citation_index[cid] = citation
|
||||||
|
chunk_ids.append(cid)
|
||||||
|
state.citations.append(chunk_ids)
|
||||||
|
|
||||||
|
|
||||||
def create_skill_extras(
|
def create_skill_extras(
|
||||||
|
|
@ -323,6 +189,8 @@ def create_skill_tools(
|
||||||
tools: dict[str, Any] = {}
|
tools: dict[str, Any] = {}
|
||||||
|
|
||||||
if "search" in tool_names:
|
if "search" in tool_names:
|
||||||
|
max_searches = config.qa.max_searches
|
||||||
|
search_counts: dict[str, int] = {}
|
||||||
|
|
||||||
async def search(
|
async def search(
|
||||||
ctx: RunContext[SkillRunDeps], query: str, limit: int | None = None
|
ctx: RunContext[SkillRunDeps], query: str, limit: int | None = None
|
||||||
|
|
@ -335,6 +203,14 @@ def create_skill_tools(
|
||||||
query: The search query.
|
query: The search query.
|
||||||
limit: Maximum number of results.
|
limit: Maximum number of results.
|
||||||
"""
|
"""
|
||||||
|
rid = ctx.run_id or ""
|
||||||
|
search_counts[rid] = search_counts.get(rid, 0) + 1
|
||||||
|
if search_counts[rid] > max_searches:
|
||||||
|
return (
|
||||||
|
"Search limit reached. Answer the question using "
|
||||||
|
"the results you already have."
|
||||||
|
)
|
||||||
|
|
||||||
state = _get_state(ctx, state_type)
|
state = _get_state(ctx, state_type)
|
||||||
formatted, results = await skill_search(
|
formatted, results = await skill_search(
|
||||||
db_path,
|
db_path,
|
||||||
|
|
@ -353,21 +229,12 @@ def create_skill_tools(
|
||||||
|
|
||||||
async def list_documents(
|
async def list_documents(
|
||||||
ctx: RunContext[SkillRunDeps],
|
ctx: RunContext[SkillRunDeps],
|
||||||
limit: int | None = None,
|
|
||||||
offset: int | None = None,
|
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""List documents in the knowledge base with optional pagination.
|
"""List all documents in the knowledge base."""
|
||||||
|
|
||||||
Args:
|
|
||||||
limit: Maximum number of documents to return.
|
|
||||||
offset: Number of documents to skip.
|
|
||||||
"""
|
|
||||||
state = _get_state(ctx, state_type)
|
state = _get_state(ctx, state_type)
|
||||||
result = await skill_list_documents(
|
result = await skill_list_documents(
|
||||||
db_path,
|
db_path,
|
||||||
config,
|
config,
|
||||||
limit,
|
|
||||||
offset,
|
|
||||||
filter=state.document_filter if state else None,
|
filter=state.document_filter if state else None,
|
||||||
)
|
)
|
||||||
if state:
|
if state:
|
||||||
|
|
@ -395,119 +262,84 @@ def create_skill_tools(
|
||||||
|
|
||||||
tools["get_document"] = get_document
|
tools["get_document"] = get_document
|
||||||
|
|
||||||
if "ask" in tool_names:
|
if "execute_code" in tool_names:
|
||||||
|
|
||||||
async def ask(ctx: RunContext[SkillRunDeps], question: str) -> str:
|
async def execute_code(ctx: RunContext[SkillRunDeps], code: str) -> str:
|
||||||
"""Ask a question and get an answer with citations from the knowledge base.
|
"""Execute Python code in a sandboxed interpreter.
|
||||||
|
|
||||||
|
The code has access to search(), list_documents(), llm() functions
|
||||||
|
and a virtual filesystem at /documents/ with document content and
|
||||||
|
structure (metadata.json, content.txt, items.jsonl per document).
|
||||||
|
|
||||||
|
Use print() to output results. Each call runs in a fresh
|
||||||
|
interpreter — variables do not persist between calls.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
question: The question to ask.
|
code: Python code to execute.
|
||||||
"""
|
"""
|
||||||
from haiku.rag.utils import format_citations
|
from haiku.rag.agents.analysis.dependencies import AnalysisContext
|
||||||
|
from haiku.rag.agents.analysis.sandbox import Sandbox
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
state = _get_state(ctx, state_type)
|
state = _get_state(ctx, state_type)
|
||||||
answer, citations = await skill_ask(
|
doc_filter = state.document_filter if state else None
|
||||||
db_path,
|
context = AnalysisContext(filter=doc_filter)
|
||||||
config,
|
|
||||||
question,
|
async with HaikuRAG(db_path, config=config, read_only=True) as rag:
|
||||||
qa_history=state.qa_history if state else None,
|
sandbox = Sandbox(client=rag, config=config, context=context)
|
||||||
document_filter=state.document_filter if state else None,
|
result = await sandbox.execute(code)
|
||||||
)
|
|
||||||
|
if state and sandbox._search_results:
|
||||||
|
existing = state.searches.get("_sandbox", [])
|
||||||
|
seen = {r.chunk_id for r in existing}
|
||||||
|
for sr in sandbox._search_results:
|
||||||
|
if sr.chunk_id not in seen:
|
||||||
|
existing.append(sr)
|
||||||
|
seen.add(sr.chunk_id)
|
||||||
|
state.searches["_sandbox"] = existing
|
||||||
|
|
||||||
if state:
|
if state:
|
||||||
_append_citations(state, citations)
|
state.executions.append(
|
||||||
state.qa_history.append(
|
CodeExecutionEntry(
|
||||||
QAHistoryEntry(
|
code=code,
|
||||||
question=question, answer=answer, citations=citations
|
stdout=result.stdout,
|
||||||
|
stderr=result.stderr,
|
||||||
|
success=result.success,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if result.success:
|
||||||
|
return result.stdout if result.stdout else "No output."
|
||||||
|
return f"Error: {result.stderr}\n\nOutput: {result.stdout}"
|
||||||
|
|
||||||
|
tools["execute_code"] = execute_code
|
||||||
|
|
||||||
|
if "cite" in tool_names:
|
||||||
|
|
||||||
|
async def cite(ctx: RunContext[SkillRunDeps], chunk_ids: list[str]) -> str:
|
||||||
|
"""Register chunk IDs as citations for your answer.
|
||||||
|
|
||||||
|
Call this after searching, with the chunk_id values from search
|
||||||
|
results that support your answer.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk_ids: List of chunk_id values from search results.
|
||||||
|
"""
|
||||||
|
from haiku.rag.agents.research.models import resolve_citations
|
||||||
|
|
||||||
|
state = _get_state(ctx, state_type)
|
||||||
|
if not state:
|
||||||
|
return "No state available."
|
||||||
|
|
||||||
|
all_results = []
|
||||||
|
for results_list in state.searches.values():
|
||||||
|
all_results.extend(results_list)
|
||||||
|
|
||||||
|
citations = resolve_citations(chunk_ids, all_results)
|
||||||
if citations:
|
if citations:
|
||||||
answer += "\n\n" + format_citations(citations)
|
_register_citations(state, citations)
|
||||||
|
return f"Registered {len(citations)} citation(s)."
|
||||||
|
|
||||||
return answer
|
tools["cite"] = cite
|
||||||
|
|
||||||
tools["ask"] = ask
|
|
||||||
|
|
||||||
if "research" in tool_names:
|
|
||||||
|
|
||||||
async def research(ctx: RunContext[SkillRunDeps], question: str) -> str:
|
|
||||||
"""Conduct deep multi-agent research on a question.
|
|
||||||
|
|
||||||
Iteratively searches, analyzes, and synthesizes information from the
|
|
||||||
knowledge base to produce a comprehensive research report.
|
|
||||||
Only use when the user explicitly requests deep research.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
question: The research question to investigate.
|
|
||||||
"""
|
|
||||||
state = _get_state(ctx, state_type)
|
|
||||||
formatted, title, executive_summary = await skill_research(
|
|
||||||
db_path,
|
|
||||||
config,
|
|
||||||
question,
|
|
||||||
document_filter=state.document_filter if state else None,
|
|
||||||
)
|
|
||||||
|
|
||||||
if state:
|
|
||||||
state.reports.append(
|
|
||||||
ResearchEntry(
|
|
||||||
question=question,
|
|
||||||
title=title,
|
|
||||||
executive_summary=executive_summary,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
state.qa_history.append(
|
|
||||||
QAHistoryEntry(question=question, answer=executive_summary)
|
|
||||||
)
|
|
||||||
|
|
||||||
return formatted
|
|
||||||
|
|
||||||
tools["research"] = research
|
|
||||||
|
|
||||||
if "analyze" in tool_names:
|
|
||||||
|
|
||||||
async def analyze(
|
|
||||||
ctx: RunContext[SkillRunDeps],
|
|
||||||
question: str,
|
|
||||||
document: str | None = None,
|
|
||||||
) -> str:
|
|
||||||
"""Answer complex analytical questions using code execution.
|
|
||||||
|
|
||||||
Use this for questions requiring computation, aggregation, or
|
|
||||||
data traversal across documents.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
question: The question to answer.
|
|
||||||
document: Optional document ID or title to pre-load for analysis.
|
|
||||||
"""
|
|
||||||
from haiku.rag.utils import format_citations
|
|
||||||
|
|
||||||
state = _get_state(ctx, state_type)
|
|
||||||
state_filter = state.document_filter if state else None
|
|
||||||
output, answer, program, citations = await skill_analyze(
|
|
||||||
db_path,
|
|
||||||
config,
|
|
||||||
question,
|
|
||||||
document=document,
|
|
||||||
document_filter=state_filter,
|
|
||||||
)
|
|
||||||
if state:
|
|
||||||
state.analyses.append(
|
|
||||||
AnalysisEntry(
|
|
||||||
question=question,
|
|
||||||
answer=answer,
|
|
||||||
program=program,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
if citations:
|
|
||||||
_append_citations(state, citations)
|
|
||||||
|
|
||||||
if citations:
|
|
||||||
output += "\n\n" + format_citations(citations)
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
tools["analyze"] = analyze
|
|
||||||
|
|
||||||
return tools
|
return tools
|
||||||
|
|
|
||||||
|
|
@ -6,15 +6,20 @@ from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from haiku.rag.agents.research.models import Citation
|
from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.skills._tools import AnalysisEntry
|
from haiku.rag.skills._tools import CodeExecutionEntry
|
||||||
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
from haiku.skills.models import Skill, SkillMetadata, SkillSource, StateMetadata
|
from haiku.skills.models import Skill, SkillMetadata, SkillSource, StateMetadata
|
||||||
from haiku.skills.parser import parse_skill_md
|
from haiku.skills.parser import parse_skill_md
|
||||||
|
|
||||||
|
|
||||||
class AnalysisState(BaseModel):
|
class AnalysisState(BaseModel):
|
||||||
document_filter: str | None = None
|
document_filter: str | None = None
|
||||||
analyses: list[AnalysisEntry] = []
|
executions: list[CodeExecutionEntry] = Field(default_factory=list)
|
||||||
citations: list[Citation] = Field(default_factory=list)
|
citation_index: dict[str, Citation] = Field(default_factory=dict)
|
||||||
|
citations: list[list[str]] = Field(default_factory=list)
|
||||||
|
searches: dict[str, list[SearchResult]] = Field(default_factory=dict)
|
||||||
|
documents: list[DocumentInfo] = Field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
STATE_TYPE = AnalysisState
|
STATE_TYPE = AnalysisState
|
||||||
|
|
@ -69,7 +74,12 @@ def create_skill(
|
||||||
else:
|
else:
|
||||||
db_path = config.storage.data_dir / "haiku.rag.lancedb"
|
db_path = config.storage.data_dir / "haiku.rag.lancedb"
|
||||||
|
|
||||||
tools = create_skill_tools(db_path, config, AnalysisState, ["analyze"])
|
tools = create_skill_tools(
|
||||||
|
db_path,
|
||||||
|
config,
|
||||||
|
AnalysisState,
|
||||||
|
["search", "list_documents", "execute_code", "cite"],
|
||||||
|
)
|
||||||
extras = create_skill_extras(db_path, config)
|
extras = create_skill_extras(db_path, config)
|
||||||
|
|
||||||
skill_instructions = instructions()
|
skill_instructions = instructions()
|
||||||
|
|
|
||||||
|
|
@ -10,4 +10,69 @@ description: >
|
||||||
|
|
||||||
# Analysis
|
# Analysis
|
||||||
|
|
||||||
Use the `analyze` tool for complex analytical questions. It writes and executes Python code against the knowledge base in a sandboxed Python interpreter.
|
You solve complex analytical questions by writing and executing Python code against the knowledge base.
|
||||||
|
|
||||||
|
## Tools
|
||||||
|
|
||||||
|
### execute_code
|
||||||
|
Execute Python code in a sandboxed interpreter. Each call runs in a fresh interpreter — variables do not persist between calls. Use `print()` to output results.
|
||||||
|
|
||||||
|
Inside the code, these functions are available (use `await`):
|
||||||
|
- `await search(query, limit=10)` → list of dicts with keys: chunk_id, content, document_id, document_title, document_uri, score, page_numbers, headings, doc_item_refs, labels
|
||||||
|
- `await list_documents()` → list of dicts with keys: id, title, uri, created_at
|
||||||
|
- `await llm(prompt)` → string response from an LLM (for classification, summarization, extraction)
|
||||||
|
|
||||||
|
Available modules: `json`, `re`, `math`, `pathlib`
|
||||||
|
Not supported: class definitions, generators/yield, match statements, decorators, `with` statements
|
||||||
|
|
||||||
|
### search
|
||||||
|
Search the knowledge base directly (outside code execution). Use for initial exploration before writing code.
|
||||||
|
|
||||||
|
### list_documents
|
||||||
|
List available documents. Use to discover what's in the knowledge base.
|
||||||
|
|
||||||
|
### cite
|
||||||
|
Register chunk IDs as citations. Call after your analysis with chunk_id values from search results that support your answer.
|
||||||
|
|
||||||
|
## Document Filesystem (inside execute_code)
|
||||||
|
|
||||||
|
All documents are mounted as a virtual filesystem at `/documents/`:
|
||||||
|
|
||||||
|
```
|
||||||
|
/documents/{document_id}/
|
||||||
|
metadata.json # {"id", "title", "uri", "created_at"}
|
||||||
|
content.txt # Full document text
|
||||||
|
items.jsonl # Structured items (one JSON object per line)
|
||||||
|
```
|
||||||
|
|
||||||
|
### metadata.json
|
||||||
|
Document metadata. Use `Path('/documents').iterdir()` to discover documents.
|
||||||
|
|
||||||
|
### content.txt
|
||||||
|
Full text content. Use for regex or keyword search across a whole document.
|
||||||
|
|
||||||
|
### items.jsonl
|
||||||
|
Structured document items. Each line is a JSON object with:
|
||||||
|
- `position`: sequential position in the document
|
||||||
|
- `self_ref`: item reference (e.g. "#/texts/5", "#/tables/0")
|
||||||
|
- `label`: item type — "section_header", "text", "table", "list_item", "caption", "formula", "picture", "code", "footnote"
|
||||||
|
- `text`: rendered content (tables are markdown with `|` columns)
|
||||||
|
- `page_numbers`: list of page numbers where the item appears
|
||||||
|
|
||||||
|
### Cross-referencing search results with items
|
||||||
|
Search results include `doc_item_refs` (e.g. `["#/texts/48", "#/tables/0"]`) that correspond to `self_ref` values in items.jsonl.
|
||||||
|
|
||||||
|
## Strategy
|
||||||
|
|
||||||
|
1. Use `search` tool first to understand what's in the knowledge base
|
||||||
|
2. Use `execute_code` to write analysis code
|
||||||
|
3. Iterate: run code, examine output, refine approach
|
||||||
|
4. Call `cite` with chunk IDs from search results you referenced
|
||||||
|
|
||||||
|
## Important
|
||||||
|
|
||||||
|
- Each `execute_code` call runs in a fresh interpreter (no persistent variables between calls)
|
||||||
|
- Use `print()` to output results — the output is your only feedback
|
||||||
|
- Always execute code to answer questions — don't just describe what code would do
|
||||||
|
- Use `await` for all async functions inside execute_code (search, list_documents, llm)
|
||||||
|
- Do NOT include chunk IDs or UUIDs in your answer text — your answer should read naturally. Use the `cite` tool separately to register citations.
|
||||||
|
|
|
||||||
|
|
@ -6,10 +6,8 @@ from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from haiku.rag.agents.research.models import Citation
|
from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.skills._tools import ResearchEntry
|
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.tools.document import DocumentInfo
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
from haiku.skills.models import Skill, SkillMetadata, SkillSource, StateMetadata
|
from haiku.skills.models import Skill, SkillMetadata, SkillSource, StateMetadata
|
||||||
from haiku.skills.parser import parse_skill_md
|
from haiku.skills.parser import parse_skill_md
|
||||||
|
|
||||||
|
|
@ -21,7 +19,7 @@ CRITICAL RULES:
|
||||||
3. When a skill returns citations, always include them in your response
|
3. When a skill returns citations, always include them in your response
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_RAG_TOOLS = ["search", "list_documents", "get_document", "ask", "research"]
|
_RAG_TOOLS = ["search", "list_documents", "get_document", "cite"]
|
||||||
|
|
||||||
|
|
||||||
def get_agent_preamble(config: AppConfig) -> str:
|
def get_agent_preamble(config: AppConfig) -> str:
|
||||||
|
|
@ -32,12 +30,11 @@ def get_agent_preamble(config: AppConfig) -> str:
|
||||||
|
|
||||||
|
|
||||||
class RAGState(BaseModel):
|
class RAGState(BaseModel):
|
||||||
citations: list[Citation] = Field(default_factory=list)
|
citation_index: dict[str, Citation] = Field(default_factory=dict)
|
||||||
qa_history: list[QAHistoryEntry] = Field(default_factory=list)
|
citations: list[list[str]] = Field(default_factory=list)
|
||||||
document_filter: str | None = None
|
document_filter: str | None = None
|
||||||
searches: dict[str, list[SearchResult]] = Field(default_factory=dict)
|
searches: dict[str, list[SearchResult]] = Field(default_factory=dict)
|
||||||
documents: list[DocumentInfo] = Field(default_factory=list)
|
documents: list[DocumentInfo] = Field(default_factory=list)
|
||||||
reports: list[ResearchEntry] = Field(default_factory=list)
|
|
||||||
|
|
||||||
|
|
||||||
STATE_TYPE = RAGState
|
STATE_TYPE = RAGState
|
||||||
|
|
|
||||||
|
|
@ -5,31 +5,55 @@ description: Search, retrieve and analyze documents using RAG (Retrieval Augment
|
||||||
|
|
||||||
# RAG
|
# RAG
|
||||||
|
|
||||||
You are a RAG (Retrieval Augmented Generation) assistant with access to a document knowledge base.
|
You are a RAG assistant with access to a document knowledge base.
|
||||||
Use your tools to search and answer questions. Never make up information — always use tools to get facts from the knowledge base.
|
Use your tools to search and answer questions. Never make up information — always use tools to get facts from the knowledge base.
|
||||||
|
|
||||||
## How to decide which tool to use
|
## Tools
|
||||||
|
|
||||||
**Default rule:** If the user is asking a question, use **ask**. Only use **search** when the user explicitly wants to browse or find passages.
|
### search
|
||||||
|
Search the knowledge base using hybrid search (vector + full-text). Returns ranked results with context-expanded content.
|
||||||
|
|
||||||
- **list_documents** — Use when the user wants to browse or see what documents are available (e.g., "what documents do you have?", "show me the documents", "list available docs").
|
Each result includes:
|
||||||
- **get_document** — Use when the user wants the full content of a specific document (e.g., "get the paper about X", "show me document Y"). Accepts a document ID, title, or URI — partial matches work.
|
- `chunk_id` in brackets and rank position (rank 1 = most relevant)
|
||||||
- **search** — Use when the user wants to browse, explore, or find specific passages across documents (e.g., "search for embeddings", "find mentions of transformers"). Returns all matching results as sources.
|
- Source: document title and section hierarchy
|
||||||
- **ask** — Use for factual questions that need a synthesized answer (e.g., "what is DocLayNet?", "explain the methodology"). Searches, synthesizes, and returns only the chunks actually used as citations. Always include the citations in your response.
|
- Type: content type (paragraph, table, code, list_item)
|
||||||
- **research** — Deep multi-agent research that produces comprehensive reports. **Only use when the user explicitly requests deep research** (e.g., "do a deep research on X", "research this topic thoroughly"). Never call this tool on your own — it is slow and expensive.
|
- Content: the actual text
|
||||||
|
|
||||||
## When search returns irrelevant results
|
### list_documents
|
||||||
|
List available documents in the knowledge base. Use when the user wants to browse what's available.
|
||||||
|
|
||||||
If your first search returns results that clearly don't match the question, **do not keep searching with variations**. Instead:
|
### get_document
|
||||||
- Use **ask** if the question is factual
|
Retrieve a document by ID, title, or URI. Partial matches work. Use when the user wants the full content of a specific document.
|
||||||
- Report that the knowledge base doesn't contain relevant information
|
|
||||||
|
### cite
|
||||||
|
Register chunk IDs as citations for your answer. Call this AFTER formulating your answer, with the `chunk_id` values from search results that support it.
|
||||||
|
|
||||||
|
## How to answer questions
|
||||||
|
|
||||||
|
1. Call `search` with relevant keywords from the question
|
||||||
|
2. Review the results — they are ordered by relevance (rank 1 = best match)
|
||||||
|
3. If needed, search again with different keywords (you have a limited number of searches)
|
||||||
|
4. Synthesize a concise answer based strictly on the retrieved content
|
||||||
|
5. Call `cite` with the chunk IDs of search results that informed your answer
|
||||||
|
|
||||||
|
## Guidelines
|
||||||
|
|
||||||
|
- Base answers strictly on retrieved content — do not use external knowledge
|
||||||
|
- Use the Source and Type metadata to understand context
|
||||||
|
- If multiple results are relevant, synthesize them coherently
|
||||||
|
- Be concise and direct — avoid elaboration unless asked
|
||||||
|
- If the search tool tells you the search limit is reached, stop searching and answer with what you have
|
||||||
|
- If the retrieved documents do not directly address the question, say: "I cannot find enough information in the knowledge base to answer this question." Do not guess or infer from tangentially related content.
|
||||||
|
- Do NOT include chunk IDs or UUIDs in your answer text — your answer should read naturally. Use the `cite` tool separately to register citations.
|
||||||
|
|
||||||
## When the user mentions a specific document
|
## When the user mentions a specific document
|
||||||
|
|
||||||
If the user says "search in [doc]", "find in [doc]", or "answer from [doc]":
|
If the user says "search in [doc]", "find in [doc]", or "answer from [doc]":
|
||||||
- Extract the **topic** as the `query`/`question` parameter
|
- Use `get_document` or `list_documents` first to identify the document
|
||||||
- Use **get_document** or **list_documents** first to identify the document, then search/ask with a filter
|
- Then search for the topic
|
||||||
|
|
||||||
Examples:
|
## When search returns irrelevant results
|
||||||
- "search for embeddings in the ML paper" -> first identify "ML paper", then search for "embeddings"
|
|
||||||
- "what does the DocLayNet paper say about annotations?" -> ask with question="what are the annotation methods?"
|
If your first search returns results that clearly don't match the question:
|
||||||
|
- Try one more search with different keywords
|
||||||
|
- If still irrelevant, report that the knowledge base doesn't contain relevant information
|
||||||
|
|
|
||||||
File diff suppressed because one or more lines are too long
|
|
@ -1,7 +1,3 @@
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
from haiku.rag.agents.analysis.models import AnalysisResult
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.skills.analysis import (
|
from haiku.rag.skills.analysis import (
|
||||||
STATE_NAMESPACE,
|
STATE_NAMESPACE,
|
||||||
|
|
@ -64,7 +60,7 @@ class TestAnalysisSkillCreation:
|
||||||
|
|
||||||
skill = create_skill(config=test_app_config, db_path=temp_db_path)
|
skill = create_skill(config=test_app_config, db_path=temp_db_path)
|
||||||
tool_names = {getattr(t, "__name__") for t in skill.tools if callable(t)}
|
tool_names = {getattr(t, "__name__") for t in skill.tools if callable(t)}
|
||||||
assert tool_names == {"analyze"}
|
assert tool_names == {"search", "list_documents", "execute_code", "cite"}
|
||||||
|
|
||||||
def test_create_skill_has_state(self, test_app_config, temp_db_path):
|
def test_create_skill_has_state(self, test_app_config, temp_db_path):
|
||||||
from haiku.rag.skills.analysis import AnalysisState, create_skill
|
from haiku.rag.skills.analysis import AnalysisState, create_skill
|
||||||
|
|
@ -81,8 +77,6 @@ class TestAnalysisSkillCreation:
|
||||||
assert skill.extras["db_path"] is temp_db_path
|
assert skill.extras["db_path"] is temp_db_path
|
||||||
assert "visualize_chunk" in skill.extras
|
assert "visualize_chunk" in skill.extras
|
||||||
assert "list_documents" in skill.extras
|
assert "list_documents" in skill.extras
|
||||||
assert callable(skill.extras["visualize_chunk"])
|
|
||||||
assert callable(skill.extras["list_documents"])
|
|
||||||
|
|
||||||
def test_create_skill_from_env(self, monkeypatch, temp_db_path):
|
def test_create_skill_from_env(self, monkeypatch, temp_db_path):
|
||||||
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
||||||
|
|
@ -118,80 +112,50 @@ class TestDomainPreambleInAnalysisSkillInstructions:
|
||||||
assert base_instructions in skill.instructions
|
assert base_instructions in skill.instructions
|
||||||
|
|
||||||
|
|
||||||
class TestAnalyzeTool:
|
class TestExecuteCodeTool:
|
||||||
async def test_analyze_returns_result(self, rag_db, monkeypatch):
|
async def test_execute_code_returns_output(self, rag_db):
|
||||||
from haiku.rag.skills.analysis import create_skill
|
from haiku.rag.skills.analysis import create_skill
|
||||||
|
|
||||||
monkeypatch.setattr(
|
|
||||||
HaikuRAG,
|
|
||||||
"analyze",
|
|
||||||
AsyncMock(return_value=AnalysisResult(answer="42", program="print(42)")),
|
|
||||||
)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
skill = create_skill(db_path=rag_db)
|
||||||
analyze = _get_tool(skill, "analyze")
|
execute_code = _get_tool(skill, "execute_code")
|
||||||
ctx = _make_ctx()
|
|
||||||
result = await analyze(ctx, question="How many documents?")
|
|
||||||
assert isinstance(result, str)
|
|
||||||
assert "42" in result
|
|
||||||
assert "print(42)" in result
|
|
||||||
|
|
||||||
async def test_analyze_updates_state(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.analysis import AnalysisState, create_skill
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
|
||||||
HaikuRAG,
|
|
||||||
"analyze",
|
|
||||||
AsyncMock(return_value=AnalysisResult(answer="42", program="print(42)")),
|
|
||||||
)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
analyze = _get_tool(skill, "analyze")
|
|
||||||
state = AnalysisState()
|
state = AnalysisState()
|
||||||
ctx = _make_ctx(state)
|
ctx = _make_ctx(state)
|
||||||
await analyze(ctx, question="How many documents?")
|
result = await execute_code(ctx, code="print('hello')")
|
||||||
assert len(state.analyses) == 1
|
assert "hello" in result
|
||||||
assert state.analyses[0].question == "How many documents?"
|
|
||||||
assert state.analyses[0].answer == "42"
|
|
||||||
assert state.analyses[0].program == "print(42)"
|
|
||||||
|
|
||||||
async def test_analyze_applies_document_filter_from_state(
|
async def test_execute_code_updates_state(self, rag_db):
|
||||||
self, rag_db, monkeypatch
|
|
||||||
):
|
|
||||||
from haiku.rag.skills.analysis import AnalysisState, create_skill
|
|
||||||
|
|
||||||
captured_kwargs = {}
|
|
||||||
|
|
||||||
async def mock_analyze(self, question, **kwargs):
|
|
||||||
captured_kwargs.update(kwargs)
|
|
||||||
return AnalysisResult(answer="42", program="print(42)")
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "analyze", mock_analyze)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
analyze = _get_tool(skill, "analyze")
|
|
||||||
state = AnalysisState(document_filter="title = 'AI Overview'")
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await analyze(ctx, question="How many documents?")
|
|
||||||
assert captured_kwargs.get("filter") == "title = 'AI Overview'"
|
|
||||||
|
|
||||||
async def test_analyze_with_document(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.analysis import create_skill
|
from haiku.rag.skills.analysis import create_skill
|
||||||
|
|
||||||
captured_kwargs = {}
|
skill = create_skill(db_path=rag_db)
|
||||||
|
execute_code = _get_tool(skill, "execute_code")
|
||||||
|
state = AnalysisState()
|
||||||
|
ctx = _make_ctx(state)
|
||||||
|
await execute_code(ctx, code="print('hello')")
|
||||||
|
assert len(state.executions) == 1
|
||||||
|
assert state.executions[0].code == "print('hello')"
|
||||||
|
assert state.executions[0].success is True
|
||||||
|
assert "hello" in state.executions[0].stdout
|
||||||
|
|
||||||
async def mock_analyze(self, question, **kwargs):
|
async def test_execute_code_reports_errors(self, rag_db):
|
||||||
captured_kwargs.update(kwargs)
|
from haiku.rag.skills.analysis import create_skill
|
||||||
return AnalysisResult(answer="Result", program="code()")
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "analyze", mock_analyze)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
skill = create_skill(db_path=rag_db)
|
||||||
analyze = _get_tool(skill, "analyze")
|
execute_code = _get_tool(skill, "execute_code")
|
||||||
ctx = _make_ctx()
|
state = AnalysisState()
|
||||||
await analyze(
|
ctx = _make_ctx(state)
|
||||||
ctx,
|
result = await execute_code(ctx, code="x = 1/0")
|
||||||
question="Count pages",
|
assert "Error" in result
|
||||||
document="AI Overview",
|
assert "ZeroDivisionError" in result
|
||||||
|
assert state.executions[0].success is False
|
||||||
|
|
||||||
|
async def test_execute_code_applies_document_filter(self, rag_db):
|
||||||
|
from haiku.rag.skills.analysis import create_skill
|
||||||
|
|
||||||
|
skill = create_skill(db_path=rag_db)
|
||||||
|
execute_code = _get_tool(skill, "execute_code")
|
||||||
|
state = AnalysisState(document_filter="title = 'AI Overview'")
|
||||||
|
ctx = _make_ctx(state)
|
||||||
|
result = await execute_code(
|
||||||
|
ctx, code="docs = await list_documents()\nprint(len(docs))"
|
||||||
)
|
)
|
||||||
assert captured_kwargs.get("documents") == ["AI Overview"]
|
assert "1" in result
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,3 @@
|
||||||
from unittest.mock import AsyncMock
|
|
||||||
|
|
||||||
from haiku.rag.agents.research.models import Citation, ResearchReport
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.skills.rag import (
|
from haiku.rag.skills.rag import (
|
||||||
STATE_NAMESPACE,
|
STATE_NAMESPACE,
|
||||||
|
|
@ -13,7 +9,6 @@ from haiku.rag.skills.rag import (
|
||||||
)
|
)
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.tools.document import DocumentInfo
|
from haiku.rag.tools.document import DocumentInfo
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
from haiku.skills.models import SkillMetadata, StateMetadata
|
from haiku.skills.models import SkillMetadata, StateMetadata
|
||||||
|
|
||||||
from .conftest import _get_tool, _make_ctx
|
from .conftest import _get_tool, _make_ctx
|
||||||
|
|
@ -116,13 +111,7 @@ class TestRAGSkillCreation:
|
||||||
|
|
||||||
skill = create_skill(config=test_app_config, db_path=temp_db_path)
|
skill = create_skill(config=test_app_config, db_path=temp_db_path)
|
||||||
tool_names = {getattr(t, "__name__") for t in skill.tools if callable(t)}
|
tool_names = {getattr(t, "__name__") for t in skill.tools if callable(t)}
|
||||||
assert tool_names == {
|
assert tool_names == {"search", "list_documents", "get_document", "cite"}
|
||||||
"search",
|
|
||||||
"list_documents",
|
|
||||||
"get_document",
|
|
||||||
"ask",
|
|
||||||
"research",
|
|
||||||
}
|
|
||||||
|
|
||||||
def test_create_skill_has_state(self, test_app_config, temp_db_path):
|
def test_create_skill_has_state(self, test_app_config, temp_db_path):
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
@ -139,8 +128,6 @@ class TestRAGSkillCreation:
|
||||||
assert skill.extras["db_path"] is temp_db_path
|
assert skill.extras["db_path"] is temp_db_path
|
||||||
assert "visualize_chunk" in skill.extras
|
assert "visualize_chunk" in skill.extras
|
||||||
assert "list_documents" in skill.extras
|
assert "list_documents" in skill.extras
|
||||||
assert callable(skill.extras["visualize_chunk"])
|
|
||||||
assert callable(skill.extras["list_documents"])
|
|
||||||
|
|
||||||
def test_create_skill_from_env(self, monkeypatch, temp_db_path):
|
def test_create_skill_from_env(self, monkeypatch, temp_db_path):
|
||||||
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
||||||
|
|
@ -169,44 +156,6 @@ class TestSkillExtras:
|
||||||
assert len(results) == 1
|
assert len(results) == 1
|
||||||
assert results[0]["title"] == "AI Overview"
|
assert results[0]["title"] == "AI Overview"
|
||||||
|
|
||||||
async def test_visualize_chunk_unknown_returns_empty(
|
|
||||||
self,
|
|
||||||
test_app_config,
|
|
||||||
rag_db,
|
|
||||||
):
|
|
||||||
from haiku.rag.skills.rag import create_skill
|
|
||||||
|
|
||||||
skill = create_skill(config=test_app_config, db_path=rag_db)
|
|
||||||
visualize = skill.extras["visualize_chunk"]
|
|
||||||
result = await visualize("nonexistent-chunk-id")
|
|
||||||
assert result == []
|
|
||||||
|
|
||||||
async def test_visualize_chunk_returns_images(
|
|
||||||
self,
|
|
||||||
test_app_config,
|
|
||||||
rag_db,
|
|
||||||
monkeypatch,
|
|
||||||
):
|
|
||||||
from haiku.rag.client import HaikuRAG
|
|
||||||
from haiku.rag.skills.rag import create_skill
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
|
||||||
HaikuRAG, "visualize_chunk", AsyncMock(return_value=["img1"])
|
|
||||||
)
|
|
||||||
|
|
||||||
skill = create_skill(config=test_app_config, db_path=rag_db)
|
|
||||||
visualize = skill.extras["visualize_chunk"]
|
|
||||||
|
|
||||||
# Get a real chunk_id from the db
|
|
||||||
async with HaikuRAG(rag_db, read_only=True) as rag:
|
|
||||||
docs = await rag.list_documents()
|
|
||||||
doc = await rag.get_document_by_id(docs[0].id)
|
|
||||||
chunks = await rag.chunk_repository.get_by_document_id(doc.id)
|
|
||||||
chunk_id = str(chunks[0].id)
|
|
||||||
|
|
||||||
result = await visualize(chunk_id)
|
|
||||||
assert result == ["img1"]
|
|
||||||
|
|
||||||
|
|
||||||
class TestSearchTool:
|
class TestSearchTool:
|
||||||
async def test_search_returns_formatted_string(self, rag_db):
|
async def test_search_returns_formatted_string(self, rag_db):
|
||||||
|
|
@ -274,7 +223,6 @@ class TestListDocumentsTool:
|
||||||
await list_docs(ctx)
|
await list_docs(ctx)
|
||||||
assert len(state.documents) == 2
|
assert len(state.documents) == 2
|
||||||
assert isinstance(state.documents[0], DocumentInfo)
|
assert isinstance(state.documents[0], DocumentInfo)
|
||||||
assert state.documents[0].id is not None
|
|
||||||
|
|
||||||
async def test_list_documents_applies_document_filter_from_state(self, rag_db):
|
async def test_list_documents_applies_document_filter_from_state(self, rag_db):
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
@ -287,17 +235,6 @@ class TestListDocumentsTool:
|
||||||
assert len(results) == 1
|
assert len(results) == 1
|
||||||
assert results[0]["title"] == "AI Overview"
|
assert results[0]["title"] == "AI Overview"
|
||||||
|
|
||||||
async def test_list_documents_no_duplicates_in_state(self, rag_db):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
list_docs = _get_tool(skill, "list_documents")
|
|
||||||
state = RAGState()
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await list_docs(ctx)
|
|
||||||
await list_docs(ctx)
|
|
||||||
assert len(state.documents) == 2
|
|
||||||
|
|
||||||
|
|
||||||
class TestGetDocumentTool:
|
class TestGetDocumentTool:
|
||||||
async def test_get_document_by_title(self, rag_db):
|
async def test_get_document_by_title(self, rag_db):
|
||||||
|
|
@ -310,18 +247,6 @@ class TestGetDocumentTool:
|
||||||
assert result is not None
|
assert result is not None
|
||||||
assert result["title"] == "AI Overview"
|
assert result["title"] == "AI Overview"
|
||||||
|
|
||||||
async def test_get_document_updates_state(self, rag_db):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
get_doc = _get_tool(skill, "get_document")
|
|
||||||
state = RAGState()
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await get_doc(ctx, query="AI Overview")
|
|
||||||
assert len(state.documents) == 1
|
|
||||||
assert isinstance(state.documents[0], DocumentInfo)
|
|
||||||
assert state.documents[0].title == "AI Overview"
|
|
||||||
|
|
||||||
async def test_get_document_not_found(self, rag_db):
|
async def test_get_document_not_found(self, rag_db):
|
||||||
from haiku.rag.skills.rag import create_skill
|
from haiku.rag.skills.rag import create_skill
|
||||||
|
|
||||||
|
|
@ -332,343 +257,57 @@ class TestGetDocumentTool:
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
class TestAskTool:
|
class TestCiteTool:
|
||||||
async def test_ask_returns_answer_with_citations(self, rag_db, monkeypatch):
|
async def test_cite_registers_citations(self, rag_db):
|
||||||
from haiku.rag.skills.rag import create_skill
|
|
||||||
|
|
||||||
citations = [
|
|
||||||
Citation(
|
|
||||||
document_id="d1",
|
|
||||||
chunk_id="c1",
|
|
||||||
document_uri="test://ai-overview",
|
|
||||||
document_title="AI Overview",
|
|
||||||
content="AI is transforming industries.",
|
|
||||||
)
|
|
||||||
]
|
|
||||||
monkeypatch.setattr(
|
|
||||||
HaikuRAG,
|
|
||||||
"ask",
|
|
||||||
AsyncMock(return_value=("AI transforms industries worldwide.", citations)),
|
|
||||||
)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
ask = _get_tool(skill, "ask")
|
|
||||||
ctx = _make_ctx()
|
|
||||||
result = await ask(ctx, question="What is AI?")
|
|
||||||
assert isinstance(result, str)
|
|
||||||
assert "AI transforms industries" in result
|
|
||||||
|
|
||||||
async def test_ask_updates_state(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
||||||
citations = [
|
|
||||||
Citation(
|
|
||||||
document_id="d1",
|
|
||||||
chunk_id="c1",
|
|
||||||
document_uri="test://ai-overview",
|
|
||||||
content="AI content",
|
|
||||||
)
|
|
||||||
]
|
|
||||||
monkeypatch.setattr(
|
|
||||||
HaikuRAG,
|
|
||||||
"ask",
|
|
||||||
AsyncMock(return_value=("AI transforms industries.", citations)),
|
|
||||||
)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
skill = create_skill(db_path=rag_db)
|
||||||
ask = _get_tool(skill, "ask")
|
search = _get_tool(skill, "search")
|
||||||
|
cite = _get_tool(skill, "cite")
|
||||||
state = RAGState()
|
state = RAGState()
|
||||||
ctx = _make_ctx(state)
|
ctx = _make_ctx(state)
|
||||||
await ask(ctx, question="What is AI?")
|
|
||||||
|
await search(ctx, query="artificial intelligence")
|
||||||
|
chunk_ids = [
|
||||||
|
sr.chunk_id
|
||||||
|
for results in state.searches.values()
|
||||||
|
for sr in results
|
||||||
|
if sr.chunk_id
|
||||||
|
][:2]
|
||||||
|
|
||||||
|
result = await cite(ctx, chunk_ids=chunk_ids)
|
||||||
|
assert "Registered" in result
|
||||||
assert len(state.citations) == 1
|
assert len(state.citations) == 1
|
||||||
assert len(state.qa_history) == 1
|
assert len(state.citations[0]) == 2
|
||||||
assert isinstance(state.qa_history[0], QAHistoryEntry)
|
assert all(cid in state.citation_index for cid in chunk_ids)
|
||||||
assert state.qa_history[0].question == "What is AI?"
|
|
||||||
|
|
||||||
async def test_ask_assigns_citation_indices(self, rag_db, monkeypatch):
|
async def test_cite_deduplicates_in_index(self, rag_db):
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
from haiku.rag.skills.rag import RAGState, create_skill
|
||||||
|
|
||||||
first_citations = [
|
|
||||||
Citation(
|
|
||||||
document_id="d1",
|
|
||||||
chunk_id="c1",
|
|
||||||
document_uri="test://doc1",
|
|
||||||
content="First.",
|
|
||||||
),
|
|
||||||
Citation(
|
|
||||||
document_id="d2",
|
|
||||||
chunk_id="c2",
|
|
||||||
document_uri="test://doc2",
|
|
||||||
content="Second.",
|
|
||||||
),
|
|
||||||
]
|
|
||||||
second_citations = [
|
|
||||||
Citation(
|
|
||||||
document_id="d3",
|
|
||||||
chunk_id="c3",
|
|
||||||
document_uri="test://doc3",
|
|
||||||
content="Third.",
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
call_count = 0
|
|
||||||
|
|
||||||
async def mock_ask(self, question, **kwargs):
|
|
||||||
nonlocal call_count
|
|
||||||
call_count += 1
|
|
||||||
if call_count == 1:
|
|
||||||
return ("Answer 1", first_citations)
|
|
||||||
return ("Answer 2", second_citations)
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "ask", mock_ask)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
skill = create_skill(db_path=rag_db)
|
||||||
ask = _get_tool(skill, "ask")
|
search = _get_tool(skill, "search")
|
||||||
|
cite = _get_tool(skill, "cite")
|
||||||
state = RAGState()
|
state = RAGState()
|
||||||
ctx = _make_ctx(state)
|
ctx = _make_ctx(state)
|
||||||
|
|
||||||
await ask(ctx, question="First question")
|
await search(ctx, query="artificial intelligence")
|
||||||
assert state.citations[0].index == 1
|
chunk_ids = [
|
||||||
assert state.citations[1].index == 2
|
sr.chunk_id
|
||||||
|
for results in state.searches.values()
|
||||||
|
for sr in results
|
||||||
|
if sr.chunk_id
|
||||||
|
][:1]
|
||||||
|
|
||||||
await ask(ctx, question="Second question")
|
await cite(ctx, chunk_ids=chunk_ids)
|
||||||
assert state.citations[2].index == 3
|
await cite(ctx, chunk_ids=chunk_ids)
|
||||||
|
assert len(state.citation_index) == 1
|
||||||
|
assert len(state.citations) == 2
|
||||||
|
|
||||||
async def test_ask_applies_document_filter_from_state(self, rag_db, monkeypatch):
|
async def test_cite_without_state(self, rag_db):
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
|
|
||||||
captured_kwargs = {}
|
|
||||||
|
|
||||||
async def mock_ask(self, question, **kwargs):
|
|
||||||
captured_kwargs.update(kwargs)
|
|
||||||
return ("Answer.", [])
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "ask", mock_ask)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
ask = _get_tool(skill, "ask")
|
|
||||||
state = RAGState(document_filter="title = 'AI Overview'")
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await ask(ctx, question="What is AI?")
|
|
||||||
assert captured_kwargs.get("filter") == "title = 'AI Overview'"
|
|
||||||
|
|
||||||
async def test_ask_includes_prior_qa_context(self, rag_db, monkeypatch):
|
|
||||||
import random
|
|
||||||
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
from tests.skills.conftest import VECTOR_DIM
|
|
||||||
|
|
||||||
captured_questions = []
|
|
||||||
|
|
||||||
async def mock_ask(self, question, **kwargs):
|
|
||||||
captured_questions.append(question)
|
|
||||||
return ("Answer about AI.", [])
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "ask", mock_ask)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
ask = _get_tool(skill, "ask")
|
|
||||||
|
|
||||||
# Pre-compute the embedding the fake embedder will produce for "Tell me about AI"
|
|
||||||
query_text = "Tell me about AI"
|
|
||||||
random.seed(hash(query_text) % (2**32))
|
|
||||||
query_embedding = [random.random() for _ in range(VECTOR_DIM)]
|
|
||||||
|
|
||||||
prior_citations = [
|
|
||||||
Citation(
|
|
||||||
document_id="d1",
|
|
||||||
chunk_id="c1",
|
|
||||||
document_uri="test://ai-overview",
|
|
||||||
document_title="AI Overview",
|
|
||||||
content="AI content from source.",
|
|
||||||
)
|
|
||||||
]
|
|
||||||
state = RAGState(
|
|
||||||
qa_history=[
|
|
||||||
QAHistoryEntry(
|
|
||||||
question="What is artificial intelligence?",
|
|
||||||
answer="AI is the simulation of human intelligence by machines.",
|
|
||||||
question_embedding=query_embedding,
|
|
||||||
citations=prior_citations,
|
|
||||||
),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await ask(ctx, question=query_text)
|
|
||||||
|
|
||||||
# rag.ask() should receive augmented question with prior context
|
|
||||||
assert len(captured_questions) == 1
|
|
||||||
augmented = captured_questions[0]
|
|
||||||
assert "Context from prior questions" in augmented
|
|
||||||
assert "What is artificial intelligence?" in augmented
|
|
||||||
assert "AI is the simulation" in augmented
|
|
||||||
assert "AI Overview" in augmented
|
|
||||||
assert query_text in augmented
|
|
||||||
|
|
||||||
# State should store the original question, not the augmented one
|
|
||||||
assert state.qa_history[-1].question == query_text
|
|
||||||
|
|
||||||
async def test_ask_embeds_prior_qa_on_demand(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
from tests.skills.conftest import VECTOR_DIM
|
|
||||||
|
|
||||||
captured_questions = []
|
|
||||||
|
|
||||||
async def mock_ask(self, question, **kwargs):
|
|
||||||
captured_questions.append(question)
|
|
||||||
return ("Answer about AI.", [])
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "ask", mock_ask)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
ask = _get_tool(skill, "ask")
|
|
||||||
|
|
||||||
# Use the same question text for the prior QA entry and query so
|
|
||||||
# their fake embeddings are identical (cosine similarity = 1.0).
|
|
||||||
prior_question = "Tell me about AI"
|
|
||||||
query_text = prior_question
|
|
||||||
|
|
||||||
# Leave question_embedding=None to exercise the lazy embedding path
|
|
||||||
state = RAGState(
|
|
||||||
qa_history=[
|
|
||||||
QAHistoryEntry(
|
|
||||||
question=prior_question,
|
|
||||||
answer="AI is the simulation of human intelligence by machines.",
|
|
||||||
question_embedding=None,
|
|
||||||
),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await ask(ctx, question=query_text)
|
|
||||||
|
|
||||||
# The lazy embedding should have populated question_embedding
|
|
||||||
assert state.qa_history[0].question_embedding is not None
|
|
||||||
assert len(state.qa_history[0].question_embedding) == VECTOR_DIM
|
|
||||||
|
|
||||||
# The augmented question should include prior context
|
|
||||||
assert len(captured_questions) == 1
|
|
||||||
assert "Context from prior questions" in captured_questions[0]
|
|
||||||
assert prior_question in captured_questions[0]
|
|
||||||
|
|
||||||
async def test_ask_no_prior_qa_context_when_irrelevant(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
from tests.skills.conftest import VECTOR_DIM
|
|
||||||
|
|
||||||
captured_questions = []
|
|
||||||
|
|
||||||
async def mock_ask(self, question, **kwargs):
|
|
||||||
captured_questions.append(question)
|
|
||||||
return ("Answer.", [])
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "ask", mock_ask)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
ask = _get_tool(skill, "ask")
|
|
||||||
|
|
||||||
# Use orthogonal embedding — won't match the fake embedder's output
|
|
||||||
orthogonal = [1.0 if i % 2 == 0 else -1.0 for i in range(VECTOR_DIM)]
|
|
||||||
state = RAGState(
|
|
||||||
qa_history=[
|
|
||||||
QAHistoryEntry(
|
|
||||||
question="What is the weather?",
|
|
||||||
answer="It is sunny today.",
|
|
||||||
question_embedding=orthogonal,
|
|
||||||
),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await ask(ctx, question="Explain quantum computing")
|
|
||||||
|
|
||||||
# rag.ask() should receive the original question unchanged
|
|
||||||
assert len(captured_questions) == 1
|
|
||||||
assert captured_questions[0] == "Explain quantum computing"
|
|
||||||
|
|
||||||
|
|
||||||
class TestResearchTool:
|
|
||||||
async def test_research_returns_report(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import create_skill
|
from haiku.rag.skills.rag import create_skill
|
||||||
|
|
||||||
report = ResearchReport(
|
|
||||||
title="AI Research",
|
|
||||||
executive_summary="AI is transforming industries.",
|
|
||||||
main_findings=["Finding 1"],
|
|
||||||
conclusions=["Conclusion 1"],
|
|
||||||
sources_summary="Multiple sources consulted.",
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(HaikuRAG, "research", AsyncMock(return_value=report))
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
skill = create_skill(db_path=rag_db)
|
||||||
research = _get_tool(skill, "research")
|
cite = _get_tool(skill, "cite")
|
||||||
ctx = _make_ctx()
|
|
||||||
result = await research(ctx, question="What is AI?")
|
|
||||||
assert isinstance(result, str)
|
|
||||||
assert "AI Research" in result
|
|
||||||
|
|
||||||
async def test_research_updates_state(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
|
|
||||||
report = ResearchReport(
|
|
||||||
title="AI Research",
|
|
||||||
executive_summary="AI is transforming industries.",
|
|
||||||
main_findings=["Finding 1"],
|
|
||||||
conclusions=["Conclusion 1"],
|
|
||||||
sources_summary="Multiple sources consulted.",
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(HaikuRAG, "research", AsyncMock(return_value=report))
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
research = _get_tool(skill, "research")
|
|
||||||
state = RAGState()
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await research(ctx, question="What is AI?")
|
|
||||||
assert len(state.reports) == 1
|
|
||||||
assert state.reports[0].question == "What is AI?"
|
|
||||||
assert len(state.qa_history) == 1
|
|
||||||
assert state.qa_history[0].question == "What is AI?"
|
|
||||||
assert state.qa_history[0].answer == "AI is transforming industries."
|
|
||||||
|
|
||||||
async def test_research_applies_document_filter_from_state(
|
|
||||||
self, rag_db, monkeypatch
|
|
||||||
):
|
|
||||||
from haiku.rag.skills.rag import RAGState, create_skill
|
|
||||||
|
|
||||||
captured_kwargs = {}
|
|
||||||
|
|
||||||
report = ResearchReport(
|
|
||||||
title="AI Research",
|
|
||||||
executive_summary="Summary.",
|
|
||||||
main_findings=["Finding"],
|
|
||||||
conclusions=["Conclusion"],
|
|
||||||
sources_summary="Sources.",
|
|
||||||
)
|
|
||||||
|
|
||||||
async def mock_research(self, question, **kwargs):
|
|
||||||
captured_kwargs.update(kwargs)
|
|
||||||
return report
|
|
||||||
|
|
||||||
monkeypatch.setattr(HaikuRAG, "research", mock_research)
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
research = _get_tool(skill, "research")
|
|
||||||
state = RAGState(document_filter="title = 'AI Overview'")
|
|
||||||
ctx = _make_ctx(state)
|
|
||||||
await research(ctx, question="What is AI?")
|
|
||||||
assert captured_kwargs.get("filter") == "title = 'AI Overview'"
|
|
||||||
|
|
||||||
async def test_research_without_state(self, rag_db, monkeypatch):
|
|
||||||
from haiku.rag.skills.rag import create_skill
|
|
||||||
|
|
||||||
report = ResearchReport(
|
|
||||||
title="AI Research",
|
|
||||||
executive_summary="Summary.",
|
|
||||||
main_findings=["Finding"],
|
|
||||||
conclusions=["Conclusion"],
|
|
||||||
sources_summary="Sources.",
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(HaikuRAG, "research", AsyncMock(return_value=report))
|
|
||||||
|
|
||||||
skill = create_skill(db_path=rag_db)
|
|
||||||
research = _get_tool(skill, "research")
|
|
||||||
ctx = _make_ctx(state=None)
|
ctx = _make_ctx(state=None)
|
||||||
result = await research(ctx, question="What is AI?")
|
result = await cite(ctx, chunk_ids=["nonexistent"])
|
||||||
assert isinstance(result, str)
|
assert "No state" in result
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue