haiku.rag/haiku_rag_slim/haiku/rag/capabilities/analysis.py
Yiorgis Gozadinos 8a7cfb949f
Match capability instructions to the run's collection scope
The collection block was chosen when the capability was built, so a run narrowed
through `state.sources` to one collection was still told how to attribute across
collections it could not reach, while its results correctly carried no
`Collection:` line.

`get_instructions` composes it per run instead, from `state.sources` where the
question narrowed the conversation and from the lent client or the scope
otherwise. Order is preamble, base instructions, collection block.
2026-08-27 12:41:15 +03:00

242 lines
8.3 KiB
Python

from dataclasses import dataclass, field
from functools import cache
from pathlib import Path
from typing import TYPE_CHECKING, Any
from pydantic import Field
from pydantic_ai import RunContext, ToolFailed
from pydantic_ai.messages import ToolReturn
from pydantic_ai.toolsets import FunctionToolset
if TYPE_CHECKING:
from haiku.rag.client import HaikuRAG
from haiku.rag.capabilities._base import (
CodeExecutionEntry,
EvidenceState,
RAGCapabilityBase,
resolve_scope,
)
from haiku.rag.capabilities._tools import merge_results
from haiku.rag.config.models import AppConfig
from haiku.rag.sandbox import AnalysisContext, Sandbox
STATE_NAMESPACE = "analysis"
_CAPABILITY_ID = "haiku-rag-analysis"
_TOOL_NAMES = frozenset({"analysis_search", "analysis_execute_code", "analysis_cite"})
_instructions_path = Path(__file__).parent / "instructions" / "analysis.md"
_multiple_collections_path = (
Path(__file__).parent / "instructions" / "analysis_multiple_collections.md"
)
class AnalysisState(EvidenceState):
executions: list[CodeExecutionEntry] = Field(default_factory=list)
def begin_invocation(self) -> None:
super().begin_invocation()
self.executions.clear()
@cache
def instructions() -> str:
return _instructions_path.read_text().strip()
@cache
def multiple_collections_instructions() -> str:
"""Appended for a run that spans more than one collection."""
return _multiple_collections_path.read_text().rstrip()
def _recovery_hint(stderr: str) -> str:
"""Name the workaround for sandbox limits models trip over repeatedly.
The instructions already say file objects are not iterable, and models write
``for line in open(...)`` regardless. Carrying the fix in the error gives
them something to act on for the retry.
"""
if "TextIOWrapper" in stderr and "not iterable" in stderr:
return (
"\n\nHint: file objects cannot be iterated here. Read lines with "
'.readlines() or .read().split("\\n").'
)
return ""
@dataclass
class AnalysisCapability(RAGCapabilityBase[AnalysisState]):
"""Deferred capability for sandboxed computation over a RAG corpus."""
sandbox: Sandbox | None = field(default=None, repr=False)
execute_count: int = field(default=0, repr=False)
async def for_run(self, ctx: RunContext[Any]) -> "AnalysisCapability":
capability = await super().for_run(ctx)
assert isinstance(capability, AnalysisCapability)
capability.sandbox = None
capability.execute_count = 0
return capability
async def _ensure_sandbox(self) -> Sandbox:
if self.sandbox is None:
rag = await self._ensure_rag()
assert self.state is not None
self.sandbox = Sandbox._covering(
scope=self.scope,
config=self.config,
context=AnalysisContext(
filter=self.state.document_filter,
sources=self.state.sources,
),
rag=rag,
lock=self.rag_lock,
)
return self.sandbox
async def _close(self) -> None:
if self.sandbox is not None:
await self.sandbox.close()
self.sandbox = None
await super()._close()
def evidence_tool_names(self) -> set[str]:
# Searching from inside the sandbox does not count against
# `qa.max_searches`, so code execution outlives a spent search budget
# as a way to reach new evidence.
return super().evidence_tool_names() | {"analysis_execute_code"}
def _spent_tool_names(self) -> set[str]:
spent = super()._spent_tool_names()
if self.execute_count >= self.config.analysis.max_executions:
spent.add("analysis_execute_code")
return spent
async def _execute_code(self, code: str) -> str:
assert self.state is not None
self.execute_count += 1
if self.execute_count > self.config.analysis.max_executions:
raise ToolFailed(
"Code-execution limit reached. Give your final answer now from what "
"you already have; do not call analysis_execute_code again."
)
sandbox = await self._ensure_sandbox()
result = await sandbox.execute(code)
if result.success or result.stdout:
self._note_evidence()
if sandbox._search_results:
merge_results(
self.state.searches.setdefault("_sandbox", []),
sandbox._search_results,
)
self.state.executions.append(
CodeExecutionEntry(
code=code,
stdout=result.stdout,
stderr=result.stderr,
success=result.success,
)
)
if not result.success:
raise ToolFailed(
f"{result.stderr}{_recovery_hint(result.stderr)}"
f"\n\nOutput: {result.stdout}"
)
return result.stdout or "No output."
@classmethod
def from_spec(
cls,
db_path: Path | None = None,
config: AppConfig | None = None,
*,
defer_loading: bool = True,
request_limit: int | None = 30,
vision: bool | None = None,
) -> "AnalysisCapability":
"""Build from an agent spec, mirroring the factory's serializable arguments.
A live ``HaikuRAG`` client cannot be written in a spec, so ``rag`` is
absent here. ``config`` arrives as a mapping and is validated.
"""
return create_capability(
db_path,
AppConfig.model_validate(config) if config is not None else None,
defer_loading=defer_loading,
request_limit=request_limit,
vision=vision,
)
def get_toolset(self) -> FunctionToolset[Any]:
async def analysis_search(
ctx: RunContext[Any], query: str, limit: int | None = None
) -> str | ToolReturn:
"""Search the knowledge base for evidence to analyze."""
return await self._with_state(self._search(query, limit))
async def analysis_execute_code(ctx: RunContext[Any], code: str) -> Any:
"""Execute Python against the sandboxed document filesystem."""
return await self._with_state(self._execute_code(code))
async def analysis_cite(ctx: RunContext[Any], chunk_ids: list[str]) -> Any:
"""Register exact retrieved chunk IDs as citations for the answer."""
return await self._with_state(self._cite(chunk_ids))
return FunctionToolset(
[analysis_search, analysis_execute_code, analysis_cite],
id=_CAPABILITY_ID,
max_retries=3,
sequential=True,
)
def create_capability(
db_path: Path | str | None = None,
config: AppConfig | None = None,
*,
defer_loading: bool = True,
rag: "HaikuRAG | None" = None,
request_limit: int | None = 30,
vision: bool | None = None,
) -> AnalysisCapability:
"""Create a native Pydantic AI analysis capability.
``vision`` gates whether picture chunks are attached to search results as
images, and should reflect the model the hosting agent actually runs.
Defaults to ``config.analysis.model.vision`` (falling back to
``config.qa.model.vision``).
"""
if config is None:
from haiku.rag.config import get_config
config = get_config()
analysis_model = config.analysis.model or config.qa.model
scope = resolve_scope(db_path, config)
return AnalysisCapability(
scope=scope,
config=config,
borrowed_rag=rag,
state_type=AnalysisState,
state_namespace=STATE_NAMESPACE,
instruction_text=instructions(),
collection_instructions=multiple_collections_instructions(),
vision=analysis_model.vision if vision is None else vision,
tool_names=_TOOL_NAMES,
request_limit=request_limit,
id=_CAPABILITY_ID,
description=(
"Analyze the haiku.rag corpus with search and sandboxed Python code. "
"Use for counting, aggregation, statistics, data traversal, comparison "
"across documents, and other tasks best solved by writing Python code."
),
defer_loading=defer_loading,
)
__all__ = [
"AnalysisCapability",
"AnalysisState",
"STATE_NAMESPACE",
"create_capability",
"instructions",
]