Replace ask_question and analyze with execute_code

In Claude Code the client is the model, so the server no longer runs one.
execute_code runs a Python program per call in the analysis sandbox over
the selected documents and returns what it printed; the sandbox is created
and closed per call so Monty's cumulative budget and a frozen mount never
outlive a program. --no-agents goes with the two tools, and format_citations
in haiku.rag.utils goes with its only caller.

The sandbox exposes chunk metadata to code: chunk_meta on search results,
metadata on list_documents rows and in metadata.json, and chunks.jsonl per
document. A host-side failure inside a program, a document read or an
in-code search raising, reaches the program by exception type only and is
logged with its traceback. recovery_hint moves to haiku.rag.sandbox.

Closes #604.
This commit is contained in:
Yiorgis Gozadinos 2026-09-07 10:02:42 +03:00
parent 0026142e7d
commit b374d5eb83
No known key found for this signature in database
21 changed files with 631 additions and 570 deletions

View file

@ -7,8 +7,13 @@
- Claude Code plugin under `claude-plugin/`: the server configuration and the
`haiku-rag` skill. `claude plugin marketplace add ggozad/haiku.rag`, then
`claude plugin install haiku-rag`.
- `haiku-rag mcp --no-agents` leaves `ask_question` and `analyze`
unregistered. `create_mcp_server(agents=)`, `HaikuRAGApp.run_mcp(agents=)`.
- MCP tool `execute_code(code, filter, sources)`: runs a program in the
analysis sandbox over the selected documents and returns what it printed;
one sandbox per call.
- In the analysis sandbox, `search()` results carry `chunk_meta`,
`list_documents()` rows and `metadata.json` carry the document `metadata`,
and `/documents/{id}/chunks.jsonl` lists chunk ids with their metadata.
`recovery_hint` in `haiku.rag.sandbox`.
- MCP tools `get_document_outline` (heading tree with page numbers) and
`get_document_section` (one section's text, subsections included), built
on `document_items`. `build_toc` in `haiku.rag.context`.
@ -40,23 +45,20 @@
`SearchResult.format_for_agent(include_document_id=, include_chunk_meta=)`;
`collect_pictures` in `haiku.rag.tools.search`.
- MCP tools raise on failure; an empty result no longer doubles as an error.
Unknown document, unknown collection, invalid filter and invalid base64
carry a message; `ask_question` and `analyze` failures name the exception
type. Anything else is masked (`mask_error_details=True`) and logged
server-side.
Unknown document, unknown collection, invalid filter, invalid base64 and a
failing program carry a message. Anything else is masked
(`mask_error_details=True`) and logged server-side.
- A host-side failure inside the analysis sandbox (a document read or an
in-code `search()` raising) reaches the program as
`RuntimeError("<call> failed: <ExceptionType>")`; the traceback is logged.
- `haiku-rag mcp` covers the configured `lancedb.databases` set. `sources` on
`search_documents`, `search_documents_by_image`, `ask_question` and
`analyze`; `source` on `get_document`; an unknown name is a tool error.
`DocumentInfo.source`; citations name their database when the server
covers several. `format_citations(citations, include_source=False)`.
### Fixed
- MCP citations no longer repeat the URI of an untitled document.
`search_documents`, `search_documents_by_image` and `execute_code`; `source`
on `get_document`; an unknown name is a tool error. `DocumentInfo.source`.
### Removed
- `cite` on the MCP `ask_question` tool; citations are always appended.
- MCP tools `ask_question` and `analyze`.
- `format_citations` in `haiku.rag.utils`; `format_citations_rich` stays.
- MCP write tools `add_document_from_file`, `add_document_from_url`,
`add_document_from_text` and `delete_document`. The server opens the
database read-only; ingest with `haiku-rag add`, `add-src`, `delete` or

View file

@ -16,7 +16,7 @@ Built on [LanceDB](https://lancedb.com/), [Pydantic AI](https://ai.pydantic.dev/
- **Hybrid search** — Vector + full-text with Reciprocal Rank Fusion
- **Multimodal & cross-modal search** — Multimodal embedders (vLLM, VoyageAI, Cohere) put picture vectors in the same space as text; supports text-as-query → figure hits and image-as-query
- **Question answering** — RAG capability with citations (page numbers, section headings)
- **Vision QA** — Vision-capable models receive figure bytes alongside chunk text; attach your own images to questions in `ask`, `analyze`, MCP, and the chat TUI
- **Vision QA** — Vision-capable models receive figure bytes alongside chunk text; attach your own images to questions in `ask`, `analyze` and the chat TUI
- **Reranking** — local cross-encoders, Cohere, Zero Entropy, or vLLM
- **Analysis capability** — Complex analytical tasks via sandboxed Python code execution (aggregation, computation, multi-document analysis)
- **Evidence compaction** — Optional capability that replaces earlier questions' search results on the request with the evidence they cited, so long conversations stop resending everything they retrieved

View file

@ -1,6 +1,6 @@
---
name: haiku-rag
description: Search, read and question the user's haiku.rag knowledge base
description: Search, read and compute over the user's haiku.rag knowledge base
through the haiku-rag MCP tools. Use whenever a request could be answered
from the user's ingested documents, when asked to find, look up, check or
cite something in their documents or knowledge base, or when the question is
@ -12,8 +12,7 @@ allowed-tools:
- mcp__plugin_haiku-rag_haiku-rag__get_document_outline
- mcp__plugin_haiku-rag_haiku-rag__get_document_section
- mcp__plugin_haiku-rag_haiku-rag__list_documents
- mcp__plugin_haiku-rag_haiku-rag__ask_question
- mcp__plugin_haiku-rag_haiku-rag__analyze
- mcp__plugin_haiku-rag_haiku-rag__execute_code
---
# Working with the knowledge base
@ -38,13 +37,17 @@ whole text in reading order. For a long one, `get_document_outline` gives the
heading tree with page numbers and `get_document_section` the text of one
section, subsections included.
## Answer or compute
## Compute
`ask_question` runs the RAG agent on the server and returns an answer with
citations; use it when the user wants an answer rather than material.
`analyze` runs code in a sandbox over the documents; use it for counting,
aggregation, comparison across many documents or computation over tables. Both
cost a model call and are slower than a search.
`execute_code` runs a Python program on the server over the same documents.
Under `/documents/{id}/` each has `metadata.json`, `content.txt`, `items.jsonl`,
`chunks.jsonl` and `toc.json`, and the program can `await search(query)` and
`await list_documents()`. Write code when the answer is a count, an aggregate, a
comparison across many documents, a lookup by document or chunk metadata, or a
pattern over whole documents: whatever search cannot rank. Each call is one
program and variables do not carry over, so gather, compute and `print` a
compact result in the same program. `filter` and `sources` select the documents
it sees. Answer and cite from what it printed.
## Explore

View file

@ -16,7 +16,7 @@ When `qa.max_searches` or `analysis.max_executions` runs out, the exhausted tool
| `analysis_execute_code(code)` | Run Python against the virtual document filesystem. |
| `analysis_cite(chunk_ids)` | Register retrieved or filesystem-derived chunk IDs. |
The sandbox exposes documents under `/documents/{document_id}/` with `metadata.json`, `content.txt`, `items.jsonl`, and `toc.json`.
The sandbox exposes documents under `/documents/{document_id}/` with `metadata.json`, `content.txt`, `items.jsonl`, `chunks.jsonl` (chunk ids with their metadata) and `toc.json`. In code, `await search()` results carry `chunk_meta` and `await list_documents()` rows carry `metadata`.
## Compose an agent

View file

@ -477,9 +477,6 @@ haiku-rag mcp --port 9000
# Bind to all interfaces (containers, trusted LAN)
haiku-rag mcp --host 0.0.0.0
# Without the ask_question and analyze tools
haiku-rag mcp --no-agents
```
See [MCP](mcp.md) for details. For continuous document ingestion

View file

@ -20,7 +20,7 @@ Context expansion is automatic and section-aware. For structured documents (with
## Question Answering Configuration
Configure the RAG capability (used by `client.ask`, `haiku-rag ask`, and the MCP `ask_question` tool):
Configure the RAG capability (used by `client.ask` and `haiku-rag ask`):
```yaml
qa:

View file

@ -19,8 +19,6 @@ haiku-rag mcp --host 0.0.0.0 --port 8001
# stdio transport (for Claude Desktop)
haiku-rag mcp --stdio
# Without ask_question and analyze, which run a model on the server
haiku-rag mcp --stdio --no-agents
```
`--host` defaults to `127.0.0.1` (loopback only). Bind to `0.0.0.0` only
@ -59,6 +57,10 @@ plugin:
claude mcp add haiku-rag -- haiku-rag mcp --stdio
```
The skill works with that registration too: copy `claude-plugin/skills/haiku-rag`
into `~/.claude/skills/` and change the tool prefix in its `allowed-tools` from
`mcp__plugin_haiku-rag_haiku-rag__` to `mcp__haiku-rag__`.
## Claude Desktop Integration
Add to your Claude Desktop configuration (`claude_desktop_config.json`):
@ -103,8 +105,7 @@ repeating it.
| `get_document_outline` | always | `document_id`, `source` |
| `get_document_section` | always | `document_id`, `section_id`, `source` |
| `list_documents` | always | `limit`, `offset`, `filter` |
| `ask_question` | unless `--no-agents` | `question`, `images_base64`, `sources` |
| `analyze` | unless `--no-agents` | `question`, `filter`, `images_base64`, `sources` |
| `execute_code` | always | `code`, `filter`, `sources` |
`search_documents` runs hybrid search, vector and full-text. Its text content
is the rendering the in-process agents read: results best first, each with its
@ -126,12 +127,18 @@ node's `id` in the outline is the `section_id`. A document without headings
has an empty outline. `list_documents` returns titles, URIs and metadata,
which is how a client learns what a filter can match.
`ask_question` runs the RAG agent on the server and returns an answer
followed by its citations. `analyze` writes and runs Python in a sandbox
over the documents, for counting, aggregation and computation across
documents. Both cost a model call. Claude Code moves a call still running
after about two minutes to a background task, which a slow local model can
trigger; `--no-agents` leaves both tools out.
`execute_code` runs a Python program in the sandbox of the
[analysis capability](capabilities/analysis.md), over the documents `filter`
and `sources` select, and returns what it printed. The program reads
`/documents/{document_id}/` (`metadata.json`, `content.txt`, `items.jsonl`,
`chunks.jsonl`, `toc.json`) and can `await search()` and
`await list_documents()`; the tool description spells out the fields and the
interpreter's limits. Each call is one program: nothing carries over between
calls, and the sandbox is created and closed per call. A failing program is a
tool error carrying the interpreter's message and any output printed before
it. `analysis.code_timeout` bounds a call and `analysis.max_output_chars` its
output; no model runs on the server. Claude Code moves a call still running
after about two minutes to a background task.
### Filters
@ -150,8 +157,10 @@ title = 'Q3 report'
A failure is an MCP error, never an empty result. Expected failures carry a
message: a document or section id that matches nothing, a collection the
server does not cover, a filter the query engine rejects (with its message),
invalid base64,
and an `ask_question` or `analyze` failure naming only the exception type.
invalid base64, and a program that fails in `execute_code`. A failure on the
server inside a program, a database read or an in-code search raising, reaches
the program and the client as its exception type only; the traceback goes to
the server log.
Anything else reaches the client as `Error calling tool 'name'` and its
traceback goes to the server log.

View file

@ -925,7 +925,6 @@ class HaikuRAGApp:
transport: str | None = None,
host: str = "127.0.0.1",
port: int = 8001,
agents: bool = True,
):
"""Run the MCP server until interrupted.
@ -935,7 +934,7 @@ class HaikuRAGApp:
# The resolved scope: a path overrides a configured URI, and a derived
# single-database configuration drops the name results and citations
# carry.
server = _mcp_server_covering(self.scope, self.config, agents=agents)
server = _mcp_server_covering(self.scope, self.config)
try:
if transport == "stdio":
await server.run_stdio_async()

View file

@ -19,7 +19,7 @@ from haiku.rag.capabilities._base import (
)
from haiku.rag.capabilities._tools import merge_results
from haiku.rag.config.models import AppConfig
from haiku.rag.sandbox import AnalysisContext, Sandbox
from haiku.rag.sandbox import AnalysisContext, Sandbox, recovery_hint
STATE_NAMESPACE = "analysis"
_CAPABILITY_ID = "haiku-rag-analysis"
@ -49,21 +49,6 @@ def multiple_collections_instructions() -> str:
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."""
@ -139,7 +124,7 @@ class AnalysisCapability(RAGCapabilityBase[AnalysisState]):
)
if not result.success:
raise ToolFailed(
f"{result.stderr}{_recovery_hint(result.stderr)}"
f"{result.stderr}{recovery_hint(result.stderr)}"
f"\n\nOutput: {result.stdout}"
)
return result.stdout or "No output."

View file

@ -13,8 +13,8 @@ You can mix the two. The rule: always call `analysis_cite` before answering —
Execute Python code in a sandboxed interpreter. Variables persist between calls — you can build state incrementally. 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, picture_refs (subset of doc_item_refs labeled `picture`)
- `await list_documents()` → list of dicts with keys: id, title, uri, created_at
- `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, picture_refs (subset of doc_item_refs labeled `picture`), chunk_meta (the matched chunk's stored metadata, custom keys included)
- `await list_documents()` → list of dicts with keys: id, title, uri, created_at, metadata
Available modules: `json`, `re`, `math`, `pathlib`
Not supported: class inheritance and metaclasses, generators/yield, match statements, decorators, `collections`, iterating a file object (`for line in f`)
@ -39,9 +39,10 @@ All documents are mounted as a virtual filesystem at `/documents/`:
```
/documents/{document_id}/
metadata.json # {"id", "title", "uri", "created_at"}
metadata.json # {"id", "title", "uri", "created_at", "metadata"}
content.txt # Full document text
items.jsonl # Structured items (one JSON object per line)
chunks.jsonl # Chunks in order with their metadata (one JSON object per line)
toc.json # Section tree derived from heading_level
```
@ -70,7 +71,7 @@ for line in Path(f'/documents/{doc_id}/items.jsonl').read_text().strip().split("
```
### metadata.json
Document metadata: `id`, `title`, `uri`, `created_at`.
Document metadata: `id`, `title`, `uri`, `created_at`, and `metadata`, the keys stored with the document.
### content.txt
Full text content. Use for regex or keyword search across a whole document.
@ -86,6 +87,9 @@ Each row carries:
- `chunk_ids`: chunks that contain this item — pass to `analysis_cite()` to ground an answer that read this item directly
- `heading_level`: H-level for `section_header` rows; `0` on non-header rows
### chunks.jsonl
The document's chunks in order, one JSON object per line: `chunk_id` and `metadata`, the chunk's stored metadata (`doc_item_refs`, `headings`, `labels`, `page_numbers`, and any custom keys such as paragraph or footnote numbers). To read by chunk metadata, keep the matching rows and take the `items.jsonl` rows whose `chunk_ids` name them.
### toc.json
Section tree derived from `heading_level`: `{"doc_id", "title", "tree": [...]}` where each node has `{self_ref, level, title, page_numbers, item_range: [start, end_exclusive], chunk_ids, children}`. `item_range` is a line slice into `items.jsonl``items[start:end]`. `chunk_ids` aggregates the citable chunks across all items in the section — pass directly to `analysis_cite()` to ground a section-scoped answer without a corpus-wide `search()` call. `tree: []` for docs with no headers.

View file

@ -884,20 +884,13 @@ def mcp(
"--port",
help="Port to bind MCP server to (ignored with --stdio)",
),
no_agents: bool = typer.Option(
False,
"--no-agents",
help="Do not register ask_question and analyze, which run a model",
),
) -> None:
"""Run the MCP server."""
app = create_app(db, covers_set=True)
transport = "stdio" if stdio else None
asyncio.run(
app.run_mcp(transport=transport, host=host, port=port, agents=not no_agents)
)
asyncio.run(app.run_mcp(transport=transport, host=host, port=port))
if __name__ == "__main__":

View file

@ -16,13 +16,13 @@ from pydantic import Field
from haiku.rag.client import HaikuRAG
from haiku.rag.config import AppConfig, get_config
from haiku.rag.context import build_toc
from haiku.rag.sandbox import AnalysisContext, Sandbox, recovery_hint
from haiku.rag.store.exceptions import UnknownDatabaseError
from haiku.rag.store.models import Document, SearchResult
from haiku.rag.store.models.document_item import DocumentItem
from haiku.rag.store.schema import DocumentMetaRecord
from haiku.rag.tools.document import DocumentInfo, DocumentSection, OutlineNode
from haiku.rag.tools.search import collect_pictures
from haiku.rag.utils import format_citations
if TYPE_CHECKING:
from typing import Any
@ -65,12 +65,6 @@ def _decode_image(image_base64: str) -> bytes:
raise ToolError("Invalid base64 image") from e
def _decode_images(images_base64: list[str] | None) -> list[bytes] | None:
if not images_base64:
return None
return [_decode_image(b64) for b64 in images_base64]
async def _check_filter(
rag: HaikuRAG, filter: str | None, sources: list[str] | None = None
) -> None:
@ -98,18 +92,14 @@ async def _check_filter(
) from e
def _instructions(scope: "DatabaseScope", config: AppConfig, agents: bool) -> str:
def _instructions(scope: "DatabaseScope", config: AppConfig) -> str:
"""What the server is for, naming no tools: the client has every tool's
description from the listing."""
lines = [
"haiku-rag is the user's knowledge base: documents they ingested, "
"searchable by meaning and keyword, readable whole or section by section."
"searchable by meaning and keyword, readable whole or section by section, "
"or computed across with code."
]
if agents:
lines.append(
"Questions can be answered from them with citations, or computed "
"across them with code."
)
lines.append(
"Use it whenever a question could be answered from those documents, "
"before answering from memory, and say when it had nothing relevant."
@ -185,9 +175,7 @@ def _find(toc: list["dict[str, Any]"], section_id: str) -> "dict[str, Any] | Non
def create_mcp_server(
db_path: Path | None = None,
config: AppConfig | None = None,
agents: bool = True,
db_path: Path | None = None, config: AppConfig | None = None
) -> FastMCP:
"""Create an MCP server over the databases the configuration places.
@ -196,20 +184,14 @@ def create_mcp_server(
None to serve the databases the configuration places. Beside
`lancedb.databases` a path raises `AmbiguousDatabaseError`.
config: Configuration to use.
agents: Register `ask_question` and `analyze`, which run a model on
the server.
"""
from haiku.rag.client.scope import DatabaseScope
config = config if config is not None else get_config()
return _covering(
DatabaseScope.resolve(config, database_path=db_path), config, agents
)
return _covering(DatabaseScope.resolve(config, database_path=db_path), config)
def _covering(
scope: "DatabaseScope", config: AppConfig, agents: bool = True
) -> FastMCP:
def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
"""An MCP server over databases someone already resolved.
Internal, as ``HaikuRAG._covering`` is: the public factory takes a path and
@ -255,7 +237,7 @@ def _covering(
# the traceback goes to the server log. A ToolError reaches the client as is.
mcp = FastMCP(
"haiku-rag",
instructions=_instructions(scope, config, agents),
instructions=_instructions(scope, config),
version=metadata.version("haiku.rag-slim"),
lifespan=lifespan,
mask_error_details=True,
@ -469,72 +451,51 @@ def _covering(
for doc in documents
]
if agents:
@mcp.tool(annotations=_read_only("Run code over the documents"))
async def execute_code(
code: str, filter: Filter = None, sources: Sources = None
) -> str:
"""Run a Python program over the documents and return what it printed.
@mcp.tool(annotations=_read_only("Ask a question"))
async def ask_question(
question: str,
images_base64: list[str] | None = None,
sources: Sources = None,
) -> str:
"""Answer a question from the documents with a retrieval agent.
Use this when the answer is a count, an aggregate, a comparison across
many documents, a lookup by document or chunk metadata, or a pattern
over whole documents: whatever a search cannot rank. The program runs
in a sandboxed interpreter on the server. Each call is one program,
nothing carries over between calls, and `print` is the only output.
Use this when the user wants an answer rather than material to read.
It runs a model on the server and is slower than a search. Returns
the answer, followed by citations to the passages it rests on.
Inside the program, `/documents/{document_id}/` holds `metadata.json`
(id, title, uri, created_at, metadata), `content.txt` (the whole text),
`items.jsonl` (one item per line: self_ref, label, text, page_numbers,
heading_level, chunk_ids), `chunks.jsonl` (one chunk per line: chunk_id,
metadata) and `toc.json` (the section tree, each node with an item_range
slice into items.jsonl). Read files with `Path.read_text()` or `open()`;
a file object cannot be iterated, use `.readlines()`.
`await search(query, limit=10)` returns dicts with chunk_id, content,
document_id, document_title, document_uri, source, score, page_numbers,
headings, doc_item_refs, labels and chunk_meta. `await list_documents()`
returns dicts with id, title, uri, created_at, source and metadata.
Modules: json, re, math, pathlib. Not available: generators, class
inheritance, match statements, decorators, collections.
Args:
question: The question, in natural language.
images_base64: Images to attach to the question, PNG or JPEG
bytes as base64. Needs a vision-capable model on the server.
"""
images = _decode_images(images_base64)
rag = await _client()
try:
answer, citations = await rag.ask(
question, images=images, sources=sources
)
except UnknownDatabaseError as e:
raise ToolError(str(e)) from e
except Exception as e:
logger.exception("ask_question failed")
raise ToolError(f"ask_question failed: {type(e).__name__}") from e
if citations:
answer += "\n\n" + format_citations(
citations, include_source=rag.covers_multiple
)
return answer
@mcp.tool(annotations=_read_only("Analyze documents"))
async def analyze(
question: str,
filter: Filter = None,
images_base64: list[str] | None = None,
sources: Sources = None,
) -> str:
"""Compute an answer across documents with code.
Use this for counting, aggregation, comparison across many documents
or arithmetic over tables, where reading passages is not enough. A
model writes and runs Python in a sandbox over the selected documents.
It is the slowest tool. Returns the answer as text.
Args:
question: The question, in natural language.
images_base64: Images to attach to the question, PNG or JPEG
bytes as base64. Needs a vision-capable model on the server.
"""
images = _decode_images(images_base64)
rag = await _client()
try:
result = await rag.analyze(
question, filter=filter, images=images, sources=sources
)
except UnknownDatabaseError as e:
raise ToolError(str(e)) from e
except Exception as e:
logger.exception("analyze failed")
raise ToolError(f"analyze failed: {type(e).__name__}") from e
return result.answer
Args:
code: The program. Use `await` on search and list_documents.
"""
rag = await _client()
sandbox = Sandbox._covering(
scope, config, AnalysisContext(filter=filter, sources=sources), rag=rag
)
try:
await _check_filter(rag, filter, sources)
result = await sandbox.execute(code)
except UnknownDatabaseError as e:
raise ToolError(str(e)) from e
finally:
await sandbox.close()
if not result.success:
raise ToolError(
f"{result.stderr}{recovery_hint(result.stderr)}"
f"\n\nOutput: {result.stdout}"
)
return result.stdout or "No output."
return mcp

View file

@ -1,10 +1,11 @@
from haiku.rag.sandbox.dependencies import AnalysisContext
from haiku.rag.sandbox.models import AnalysisResult
from haiku.rag.sandbox.sandbox import Sandbox, SandboxResult
from haiku.rag.sandbox.sandbox import Sandbox, SandboxResult, recovery_hint
__all__ = [
"AnalysisContext",
"AnalysisResult",
"Sandbox",
"SandboxResult",
"recovery_hint",
]

View file

@ -1,5 +1,6 @@
import asyncio
import json
import logging
import os
from collections.abc import AsyncIterator, Callable, Coroutine
from contextlib import asynccontextmanager, suppress
@ -19,7 +20,7 @@ from pydantic_monty import (
from haiku.rag.config.models import AppConfig
from haiku.rag.context import build_toc
from haiku.rag.sandbox.dependencies import AnalysisContext
from haiku.rag.store.models.chunk import SearchResult
from haiku.rag.store.models.chunk import Chunk, SearchResult
from haiku.rag.store.models.document_item import PICTURE_REF_PREFIX, DocumentItem
from haiku.rag.utils import gather_all
@ -30,6 +31,19 @@ if TYPE_CHECKING:
from haiku.rag.client.scope import DatabaseScope
logger = logging.getLogger(__name__)
def _host_failure(where: str, e: Exception) -> RuntimeError:
"""The error a program gets for a failure on the host side of a call.
The message and traceback go to the log. The program, and through the MCP
server its client, learn the exception type only.
"""
logger.exception("%s failed inside the sandbox", where)
return RuntimeError(f"{where} failed: {type(e).__name__}")
@dataclass
class SandboxResult:
"""Result of executing code in the sandbox."""
@ -39,6 +53,21 @@ class SandboxResult:
success: bool
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 ""
class Sandbox:
"""Execute code in a sandboxed Python interpreter.
@ -46,7 +75,8 @@ class Sandbox:
The interpreter runs in a subprocess worker checked out of an ``AsyncMonty``
pool. External functions (search, list_documents) are called by Monty code
using ``await`` and resolved asynchronously on the host. Documents are
exposed via a virtual filesystem at ``/documents/{id}/``.
exposed via a virtual filesystem at ``/documents/{id}/``: ``metadata.json``,
``content.txt``, ``items.jsonl``, ``chunks.jsonl`` and ``toc.json``.
The session persists across ``execute()`` calls within the same Sandbox
instance variables carry over. Call ``close()`` to return the worker to
@ -76,6 +106,7 @@ class Sandbox:
_doc_items: dict[str, list["DocumentItem"]]
_doc_chunk_index: dict[str, dict[str, list[str]]]
_items_jsonl_cache: dict[str, str]
_chunks_jsonl_cache: dict[str, str]
_toc_json_cache: dict[str, str]
_opened: "HaikuRAG | None"
_pool: AsyncMonty | None
@ -142,6 +173,7 @@ class Sandbox:
self._doc_items = {}
self._doc_chunk_index = {}
self._items_jsonl_cache = {}
self._chunks_jsonl_cache = {}
self._toc_json_cache = {}
self._pool = None
self._session = None
@ -249,7 +281,7 @@ class Sandbox:
loop overruns it by however long the outstanding reads take. Raising from
inside the callback answers the worker's suspension, which keeps the
session usable cancelling ``feed_run`` from outside does not, and wedges
the protocol.
the protocol. A failed read reaches the program by type only.
"""
assert self._loop is not None, (
"VFS reads happen during execute(); the loop must be captured first."
@ -260,7 +292,10 @@ class Sandbox:
"time limit exceeded: no further document reads after "
f"{self._config.analysis.code_timeout}s"
)
return asyncio.run_coroutine_threadsafe(coro, self._loop).result()
try:
return asyncio.run_coroutine_threadsafe(coro, self._loop).result()
except Exception as e:
raise _host_failure("document read", e) from None
async def _discard_session(self) -> None:
"""Drop a session whose worker is gone.
@ -330,6 +365,7 @@ class Sandbox:
"doc_item_refs": r.doc_item_refs,
"labels": r.labels,
"picture_refs": picture_refs,
"chunk_meta": r.chunk_meta,
}
)
return out
@ -343,15 +379,28 @@ class Sandbox:
"uri": d.uri,
"created_at": str(d.created_at),
"source": d.source,
"metadata": d.metadata,
}
for d in docs
]
return {
"search": search,
"list_documents": list_documents,
"search": self._guarded("search()", search),
"list_documents": self._guarded("list_documents()", list_documents),
}
@staticmethod
def _guarded(
where: str, fn: Callable[..., Coroutine[Any, Any, Any]]
) -> Callable[..., Coroutine[Any, Any, Any]]:
async def call(*args: Any, **kwargs: Any) -> Any:
try:
return await fn(*args, **kwargs)
except Exception as e:
raise _host_failure(where, e) from None
return call
async def _build_vfs(self) -> OSAccess:
"""Build the virtual filesystem with document data.
@ -359,6 +408,7 @@ class Sandbox:
- metadata.json: CallbackFile (eager, small)
- content.txt: CallbackFile (lazy, can be large)
- items.jsonl: CallbackFile (lazy, bulk-cached)
- chunks.jsonl: CallbackFile (lazy, bulk-cached)
- toc.json: CallbackFile (lazy, bulk-cached)
"""
files: list[CallbackFile] = []
@ -433,6 +483,31 @@ class Sandbox:
return read_items
def _make_chunks_reader(
did: str,
) -> Callable[["PurePosixPath"], str]:
def read_chunks(_path: "PurePosixPath") -> str:
cached = sandbox._chunks_jsonl_cache.get(did)
if cached is not None:
return cached
async def _fetch() -> list[Chunk]:
async with sandbox._connection(sandbox._owners.get(did)) as rag:
return await rag.chunk_repository.get_by_document_id(did)
chunks = sandbox._run_on_loop(_fetch())
jsonl = "\n".join(
json.dumps(
{"chunk_id": chunk.id, "metadata": chunk.metadata},
ensure_ascii=False,
)
for chunk in chunks
)
sandbox._chunks_jsonl_cache[did] = jsonl
return jsonl
return read_chunks
def _make_toc_reader(
did: str,
) -> Callable[["PurePosixPath"], str]:
@ -467,6 +542,7 @@ class Sandbox:
"title": doc.title,
"uri": doc.uri,
"created_at": str(doc.created_at),
"metadata": doc.metadata,
},
ensure_ascii=False,
)
@ -508,6 +584,13 @@ class Sandbox:
write=_deny_write,
)
)
files.append(
CallbackFile(
f"{doc_dir}/chunks.jsonl",
read=_make_chunks_reader(doc_id),
write=_deny_write,
)
)
# HAIKU_RAG_DISABLE_TOC is an evaluation-time toggle for measuring
# whether toc.json's outline view earns its place in the VFS.
# Production callers should leave it unset.

View file

@ -393,48 +393,6 @@ def _citation_label(c: "Citation") -> str:
return c.document_title or c.document_uri
def format_citations(citations: "list[Citation]", include_source: bool = False) -> str:
"""Format citations as plain text with preserved formatting.
Used by things like the MCP server where Rich renderables are not available.
Pictures referenced by the chunk are surfaced as ``[Figure: <ref>]`` markers.
``include_source`` names each citation's database, for a client covering
several.
"""
if not citations:
return ""
lines = ["## Citations\n"]
for i, c in enumerate(citations):
idx = c.index if c.index is not None else (i + 1)
title = c.document_title or c.document_uri
header = f"[{idx}] {title}"
location_parts = []
if include_source and c.source:
location_parts.append(f"Collection: {c.source}")
pages = _citation_pages(c)
if pages:
location_parts.append(pages)
section = _citation_section(c)
if section:
location_parts.append(f"Section: {section}")
# The URI is the header when there is no title; do not repeat it.
line = f"{header} {c.document_uri}" if c.document_title else header
if location_parts:
line += f" - {', '.join(location_parts)}"
lines.append(line)
for ref in c.picture_refs:
lines.append(f"[Figure: {ref}]")
lines.append(c.content)
lines.append("")
return "\n".join(lines)
def truncated(text: str, limit: int) -> str:
"""The first `limit` characters of `text`, with `…` appended when anything
was dropped. A cut result is `limit` characters plus the mark."""

View file

@ -1,4 +1,5 @@
import asyncio
import logging
import threading
from pathlib import Path
@ -112,6 +113,41 @@ class TestSandboxListDocuments:
assert "Test Document" in result.stdout
assert temp_db_path.stem in result.stdout
@pytest.mark.asyncio
async def test_list_documents_carries_metadata(self, temp_db_path):
"""Rows carry the document's metadata, so a corpus-wide pass over it is
one call rather than a file read per document."""
from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel
config = AppConfig()
docling = DoclingDocument(name="d")
docling.add_text(label=DocItemLabel.TEXT, text="Test content")
async with HaikuRAG(temp_db_path, create=True) as client:
await client.import_document(
docling,
[
Chunk(
content="Test content",
embedding=[0.1] * config.embeddings.model.vector_dim,
order=0,
)
],
uri="test://doc1",
title="Test Document",
metadata={"author": "Ada"},
)
sb = Sandbox(db_path=temp_db_path, config=config, context=AnalysisContext())
try:
result = await sb.execute(
"docs = await list_documents()\nprint(docs[0]['metadata']['author'])"
)
finally:
await sb.close()
assert result.success, result.stderr
assert "Ada" in result.stdout
class TestSandboxSearch:
"""Test search function in sandbox."""
@ -188,6 +224,51 @@ class TestSandboxSearch:
assert "str" in result.stdout
assert "True" in result.stdout
@pytest.mark.asyncio
async def test_search_returns_the_matched_chunks_metadata(
self, temp_db_path, monkeypatch
):
"""Results carry the stored metadata of the chunk that matched, custom
keys included."""
from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel
from haiku.rag.embeddings import EmbedderWrapper
config = AppConfig()
dim = config.embeddings.model.vector_dim
async def embed_query(self, text):
return [0.1] * dim
monkeypatch.setattr(EmbedderWrapper, "embed_query", embed_query)
docling = DoclingDocument(name="d")
docling.add_text(label=DocItemLabel.TEXT, text="Paragraph fourteen.")
async with HaikuRAG(temp_db_path, create=True) as client:
await client.import_document(
docling,
[
Chunk(
content="Paragraph fourteen.",
embedding=[0.1] * dim,
order=0,
metadata={"para_no": "14"},
)
],
uri="test://paras",
)
sb = Sandbox(db_path=temp_db_path, config=config, context=AnalysisContext())
try:
result = await sb.execute(
"results = await search('fourteen', limit=1)\n"
"print(results[0]['chunk_meta']['para_no'])"
)
finally:
await sb.close()
assert result.success, result.stderr
assert "14" in result.stdout
class TestSandboxExternalFunctionEdgeCases:
"""Test edge cases in external function dispatch."""
@ -240,6 +321,71 @@ class TestSandboxExternalFunctionEdgeCases:
assert not result.success
assert "external error" in result.stderr
@pytest.mark.asyncio
async def test_a_failing_search_reaches_the_program_by_type_only(
self, sandbox, monkeypatch, caplog
):
"""A host-side failure inside search() names its exception type to
the program; the message and traceback go to the log."""
async def boom(self, *args, **kwargs):
raise ValueError("failed at /secret/path")
monkeypatch.setattr(HaikuRAG, "search", boom)
with caplog.at_level(logging.ERROR, logger="haiku.rag.sandbox.sandbox"):
result = await sandbox.execute("await search('hello')")
assert not result.success
assert "search() failed: ValueError" in result.stderr
assert "/secret/path" not in result.stderr
assert any(
r.exc_info and "failed at /secret/path" in str(r.exc_info[1])
for r in caplog.records
)
@pytest.mark.asyncio
async def test_a_failing_document_read_reaches_the_program_by_type_only(
self, temp_db_path, monkeypatch, caplog
):
"""A program can catch a failed file read, and what it catches names
the exception type only."""
from haiku.rag.store.models.document import Document
async with HaikuRAG(temp_db_path, create=True) as client:
doc = await client.document_repository.create(
Document(content="x", uri="test://read", title="Read")
)
repository = type(client.document_repository)
async def boom(self, *args, **kwargs):
raise ValueError("failed at /secret/path")
monkeypatch.setattr(repository, "get_content", boom)
sb = Sandbox(
db_path=temp_db_path, config=AppConfig(), context=AnalysisContext()
)
try:
with caplog.at_level(logging.ERROR, logger="haiku.rag.sandbox.sandbox"):
result = await sb.execute(
"from pathlib import Path\n"
"try:\n"
f" Path('/documents/{doc.id}/content.txt').read_text()\n"
"except Exception as e:\n"
" print('caught:', e)"
)
finally:
await sb.close()
assert result.success, result.stderr
assert "caught:" in result.stdout
assert "ValueError" in result.stdout
assert "/secret/path" not in result.stdout
assert any(
r.exc_info and "failed at /secret/path" in str(r.exc_info[1])
for r in caplog.records
)
class TestSandboxOutputTruncation:
"""Test output truncation behavior."""
@ -312,13 +458,14 @@ class TestSandboxVFS:
@pytest.mark.asyncio
@pytest.mark.vcr()
async def test_metadata_json(self, temp_db_path):
"""metadata.json contains document title and uri."""
"""metadata.json contains document title, uri and stored metadata."""
config = AppConfig()
async with HaikuRAG(temp_db_path, create=True) as client:
doc = await client.create_document(
content="Test content",
uri="test://doc1",
title="Test Document",
metadata={"author": "Ada"},
)
context = AnalysisContext()
@ -328,11 +475,13 @@ class TestSandboxVFS:
"import json\n"
f"meta = json.loads(Path('/documents/{doc.id}/metadata.json').read_text())\n"
"print(meta['title'])\n"
"print(meta['uri'])"
"print(meta['uri'])\n"
"print(meta['metadata']['author'])"
)
assert result.success
assert result.success, result.stderr
assert "Test Document" in result.stdout
assert "test://doc1" in result.stdout
assert "Ada" in result.stdout
@pytest.mark.asyncio
@pytest.mark.vcr()
@ -386,6 +535,59 @@ class TestSandboxVFS:
assert result.success
assert result.stdout.count("True") == 6
@pytest.mark.asyncio
async def test_chunks_jsonl(self, temp_db_path):
"""chunks.jsonl lists a document's chunks in order with their stored
metadata; a chunk found by its metadata leads to its items through
their chunk_ids."""
from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel
config = AppConfig()
dim = config.embeddings.model.vector_dim
docling = DoclingDocument(name="d")
docling.add_text(label=DocItemLabel.TEXT, text="Paragraph thirteen.")
docling.add_text(label=DocItemLabel.TEXT, text="Paragraph fourteen.")
async with HaikuRAG(temp_db_path, create=True) as client:
doc = await client.import_document(
docling,
[
Chunk(
content="Paragraph thirteen.",
embedding=[0.1] * dim,
order=0,
metadata={"para_no": "13", "doc_item_refs": ["#/texts/0"]},
),
Chunk(
content="Paragraph fourteen.",
embedding=[0.1] * dim,
order=1,
metadata={"para_no": "14", "doc_item_refs": ["#/texts/1"]},
),
],
uri="test://paras",
)
sb = Sandbox(db_path=temp_db_path, config=config, context=AnalysisContext())
try:
result = await sb.execute(
"from pathlib import Path\n"
"import json\n"
f"root = Path('/documents/{doc.id}')\n"
"def rows(name):\n"
" return [json.loads(l) for l in (root / name).read_text().strip().split('\\n')]\n"
"chunks = rows('chunks.jsonl')\n"
"print(len(chunks))\n"
"hit = [c for c in chunks if c['metadata'].get('para_no') == '14']\n"
"print(len(hit))\n"
"items = rows('items.jsonl')\n"
"print([i['text'] for i in items if hit[0]['chunk_id'] in i['chunk_ids']])"
)
finally:
await sb.close()
assert result.success, result.stderr
assert result.stdout.splitlines() == ["2", "1", "['Paragraph fourteen.']"]
@pytest.mark.asyncio
@pytest.mark.vcr()
async def test_open_read(self, temp_db_path):
@ -433,7 +635,8 @@ class TestSandboxVFS:
@pytest.mark.asyncio
@pytest.mark.parametrize(
"filename", ["content.txt", "items.jsonl", "toc.json", "metadata.json"]
"filename",
["content.txt", "items.jsonl", "chunks.jsonl", "toc.json", "metadata.json"],
)
async def test_write_denied_for_every_document_file(self, temp_db_path, filename):
"""Every file in the document VFS is read-only, metadata.json included."""
@ -797,6 +1000,26 @@ class TestSandboxReadDeadline:
cannot check its duration budget while one is in flight. The sandbox
enforces the budget itself, before each read."""
@pytest.mark.asyncio
async def test_a_failed_read_reaches_the_program_by_type_only(
self, sandbox, caplog
):
"""The bridged read hands the program the exception type, not the
message, and logs the traceback."""
sandbox._loop = asyncio.get_running_loop()
async def failing_read():
raise ValueError("failed at /secret/path")
with caplog.at_level(logging.ERROR, logger="haiku.rag.sandbox.sandbox"):
with pytest.raises(RuntimeError, match="document read failed: ValueError"):
await asyncio.to_thread(sandbox._run_on_loop, failing_read())
assert any(
r.exc_info and "failed at /secret/path" in str(r.exc_info[1])
for r in caplog.records
)
@pytest.mark.asyncio
async def test_read_after_deadline_raises_without_scheduling(self, sandbox):
"""A read attempted past the deadline fails instead of querying."""

View file

@ -16,6 +16,7 @@ import pytest
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.sandbox import AnalysisContext, Sandbox
from haiku.rag.store.models.chunk import Chunk
from haiku.rag.store.models.document import Document
from haiku.rag.store.models.document_item import DocumentItem
@ -434,6 +435,38 @@ class TestVfsReadPaths:
"nope",
)
async def test_chunks_jsonl_lists_chunks_in_order_with_their_metadata(
self, temp_db_path
):
"""One row per chunk, in chunk order, carrying the stored metadata as
is; the second read of a document is served from the sandbox's cache."""
config = AppConfig()
dim = config.embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client:
doc_id = await _empty_doc(client, uri="test://paras", title="Paras")
for order, para_no in enumerate(["13", "14"]):
await client.chunk_repository.create(
Chunk(
document_id=doc_id,
content=f"Paragraph {para_no}.",
embedding=[0.1] * dim,
order=order,
metadata={"para_no": para_no, "doc_item_refs": []},
)
)
sandbox = Sandbox(temp_db_path, config, AnalysisContext())
first = await _read_vfs_text(sandbox, f"/documents/{doc_id}/chunks.jsonl")
rows = [json.loads(line) for line in first.split("\n")]
assert [row["metadata"]["para_no"] for row in rows] == ["13", "14"]
assert all(set(row) == {"chunk_id", "metadata"} for row in rows)
assert rows[0]["metadata"] == {"para_no": "13", "doc_item_refs": []}
assert sandbox._chunks_jsonl_cache[doc_id] == first
assert (
await _read_vfs_text(sandbox, f"/documents/{doc_id}/chunks.jsonl") == first
)
async def test_toc_skips_gaps_in_item_positions(self, temp_db_path):
"""Positions need not be contiguous — a heading's span may cover
positions that carry no item."""

View file

@ -705,20 +705,6 @@ async def test_run_mcp_http(app, client, monkeypatch):
)
async def test_run_mcp_hands_the_server_the_agents_switch(app, client, monkeypatch):
seen = {}
def fake_covering(scope, config, agents=True):
seen["agents"] = agents
return AsyncMock()
monkeypatch.setattr("haiku.rag.app._mcp_server_covering", fake_covering)
await app.run_mcp(transport="stdio", agents=False)
assert seen["agents"] is False
async def test_run_mcp_survives_interruption(app, client, monkeypatch):
server = AsyncMock()
server.run_stdio_async.side_effect = KeyboardInterrupt

View file

@ -1041,14 +1041,6 @@ def test_mcp_stdio_selects_the_transport(app_stub):
app_stub.run_mcp.assert_called_once()
kwargs = app_stub.run_mcp.call_args.kwargs
assert kwargs["transport"] == "stdio"
assert kwargs["agents"] is True
def test_mcp_no_agents_leaves_the_agent_tools_out(app_stub):
result = runner.invoke(cli, ["mcp", "--no-agents"] + DB_ARGS)
assert result.exit_code == 0, result.output
assert app_stub.run_mcp.call_args.kwargs["agents"] is False
def test_mcp_without_stdio_leaves_the_transport_unset(app_stub):

View file

@ -1,7 +1,6 @@
import logging
import re
from pathlib import Path
from types import SimpleNamespace
import pytest
from fastmcp.exceptions import ToolError
@ -280,32 +279,6 @@ class TestMCPReadTools:
]
assert overview["metadata"] == {"author": "Ada"}
@pytest.mark.asyncio
async def test_ask_question_appends_the_citations(self, mcp_db, monkeypatch):
from haiku.rag.store.models.citation import Citation
citation = Citation(
chunk_id="c1",
document_id="d1",
content="cited text",
document_uri="test://ai-overview",
document_title="AI Overview",
source="alpha",
)
async def fake_ask(self, question, filter=None, images=None, sources=None):
return ("the answer", [citation])
monkeypatch.setattr(HaikuRAG, "ask", fake_ask)
mcp = create_mcp_server(mcp_db)
ask = await _get_tool(mcp, "ask_question")
answer = await ask(question="q")
assert answer.startswith("the answer")
assert "AI Overview" in answer
# One database: its name adds nothing.
assert "alpha" not in answer
@pytest.fixture
async def outlined_db(temp_db_path):
@ -666,18 +639,6 @@ class TestMCPDescribesItself:
assert "beta" in covering_both
assert "beta" not in covering_one
@pytest.mark.asyncio
async def test_instructions_without_agents_drop_only_their_clause(self, mcp_db):
from fastmcp import Client
async with Client(create_mcp_server(mcp_db)) as client:
full = client.instructions.splitlines()
async with Client(create_mcp_server(mcp_db, agents=False)) as client:
without = client.instructions.splitlines()
assert set(without) < set(full)
assert len(without) == len(full) - 1
@pytest.mark.asyncio
async def test_instructions_carry_the_domain_preamble(self, mcp_db):
from fastmcp import Client
@ -702,7 +663,7 @@ class TestMCPDescribesItself:
async with Client(create_mcp_server(mcp_db)) as client:
tools = await client.list_tools()
assert len(tools) == 8
assert len(tools) == 7
for tool in tools:
assert tool.annotations is not None, tool.name
assert tool.annotations.read_only_hint is True, tool.name
@ -722,7 +683,7 @@ class TestMCPDescribesItself:
for name, schema in tool.input_schema.get("properties", {}).items()
if not schema.get("description")
]
assert len(tools) == 8
assert len(tools) == 7
assert undescribed == []
@ -737,21 +698,138 @@ class TestMCPToolSet:
"get_document_outline",
"get_document_section",
"list_documents",
"ask_question",
"analyze",
"execute_code",
}
_COUNT_DOCUMENTS = (
"from pathlib import Path\n"
"n = 0\n"
"for d in Path('/documents').iterdir():\n"
" n += 1\n"
"print(n)"
)
class TestMCPExecuteCode:
"""`execute_code` runs one program per call in the analysis sandbox over
the documents the filter and sources select, and returns what it printed."""
@pytest.mark.asyncio
async def test_without_agents_the_agent_tools_are_not_registered(self, mcp_db):
mcp = create_mcp_server(mcp_db, agents=False)
async def test_a_program_reads_the_documents_and_returns_what_it_printed(
self, mcp_db
):
result = await _call(
create_mcp_server(mcp_db), "execute_code", code=_COUNT_DOCUMENTS
)
assert {t.name for t in await mcp.list_tools()} == {
"search_documents",
"get_document",
"get_document_outline",
"get_document_section",
"list_documents",
}
assert not result.is_error
assert result.content[0].text.strip() == "2"
@pytest.mark.asyncio
async def test_a_silent_program_says_so(self, mcp_db):
result = await _call(create_mcp_server(mcp_db), "execute_code", code="x = 1")
assert not result.is_error
assert result.content[0].text == "No output."
@pytest.mark.asyncio
async def test_filter_narrows_the_documents_a_program_sees(self, mcp_db):
result = await _call(
create_mcp_server(mcp_db),
"execute_code",
code=_COUNT_DOCUMENTS,
filter="title = 'AI Overview'",
)
assert result.content[0].text.strip() == "1"
@pytest.mark.asyncio
async def test_sources_narrows_the_documents_a_program_sees(self, two_dbs):
mcp = _covering_all(two_dbs)
both = await _call(mcp, "execute_code", code=_COUNT_DOCUMENTS)
beta = await _call(mcp, "execute_code", code=_COUNT_DOCUMENTS, sources=["beta"])
assert both.content[0].text.strip() == "2"
assert beta.content[0].text.strip() == "1"
@pytest.mark.asyncio
async def test_a_failing_program_is_an_error_carrying_the_cause_and_its_output(
self, mcp_db
):
code = (
"from pathlib import Path\n"
"print('before')\n"
"for d in Path('/documents').iterdir():\n"
" for line in open(d / 'items.jsonl'):\n"
" pass"
)
result = await _call(create_mcp_server(mcp_db), "execute_code", code=code)
assert result.is_error
text = result.content[0].text
assert "not iterable" in text
assert ".readlines()" in text
assert "Output: before" in text
@pytest.mark.asyncio
async def test_calls_share_no_state(self, mcp_db):
mcp = create_mcp_server(mcp_db)
first = await _call(mcp, "execute_code", code="x = 1\nprint(x)")
second = await _call(mcp, "execute_code", code="print(x)")
assert first.content[0].text.strip() == "1"
assert second.is_error
assert "NameError" in second.content[0].text
@pytest.mark.asyncio
async def test_every_call_closes_its_sandbox(self, mcp_db, monkeypatch):
from haiku.rag.sandbox import Sandbox
closed = []
close = Sandbox.close
async def closing(self):
closed.append(self)
await close(self)
monkeypatch.setattr(Sandbox, "close", closing)
mcp = create_mcp_server(mcp_db)
await _call(mcp, "execute_code", code="print(1)")
await _call(mcp, "execute_code", code="raise ValueError('x')")
assert len(closed) == 2
assert closed[0] is not closed[1]
@pytest.mark.asyncio
async def test_a_program_reaches_chunk_metadata(self, mcp_db):
async with HaikuRAG(mcp_db, create=True) as rag:
doc = await rag.get_document_by_uri("test://ai-overview")
embedding = (await rag.embedder.embed_documents(["x"]))[0]
await rag.chunk_repository.create(
Chunk(
document_id=doc.id,
content="Paragraph fourteen.",
metadata={"para_no": "14"},
embedding=embedding,
)
)
code = (
"from pathlib import Path\n"
"import json\n"
f"text = Path('/documents/{doc.id}/chunks.jsonl').read_text()\n"
"rows = [json.loads(line) for line in text.strip().split('\\n')]\n"
"print(len([r for r in rows if r['metadata'].get('para_no') == '14']))"
)
result = await _call(create_mcp_server(mcp_db), "execute_code", code=code)
assert not result.is_error, result.content[0].text
assert result.content[0].text.strip() == "1"
class TestMCPCoversTheConfiguredSet:
@ -784,8 +862,7 @@ class TestMCPCoversTheConfiguredSet:
{"image_base64": "AAAA", "sources": ["nope"]},
),
("get_document", {"document_id": "x", "source": "nope"}),
("ask_question", {"question": "q", "sources": ["nope"]}),
("analyze", {"question": "q", "sources": ["nope"]}),
("execute_code", {"code": "print(1)", "sources": ["nope"]}),
],
)
async def test_an_unknown_database_is_an_error_not_an_empty_result(
@ -845,59 +922,6 @@ class TestMCPCoversTheConfiguredSet:
assert {_line(block, "Collection") for block in blocks} == {"alpha", "beta"}
@pytest.mark.asyncio
async def test_ask_question_names_each_citations_database(
self, two_dbs, monkeypatch
):
from haiku.rag.store.models.citation import Citation
def cited(source):
return Citation(
chunk_id="c1",
document_id="d1",
content="cited text",
document_uri="test://cats",
document_title="Cats",
source=source,
)
async def fake_ask(self, question, filter=None, images=None, sources=None):
return ("the answer", [cited("alpha"), cited("beta")])
monkeypatch.setattr(HaikuRAG, "ask", fake_ask)
mcp = _covering_all(two_dbs)
ask = await _get_tool(mcp, "ask_question")
answer = await ask(question="q")
assert "alpha" in answer
assert "beta" in answer
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_name,client_method,returns",
[
("ask_question", "ask", ("answer", [])),
("analyze", "analyze", SimpleNamespace(answer="answer")),
],
)
async def test_agents_search_the_selected_databases(
self, two_dbs, monkeypatch, tool_name, client_method, returns
):
seen = {}
async def fake(self, question, filter=None, images=None, sources=None):
seen["sources"] = sources
return returns
monkeypatch.setattr(HaikuRAG, client_method, fake)
mcp = _covering_all(two_dbs)
tool = await _get_tool(mcp, tool_name)
await tool(question="q", sources=["beta"])
assert seen["sources"] == ["beta"]
class TestMCPImageQuery:
"""search_documents_by_image is registered only when the embedder is multimodal."""
@ -963,63 +987,6 @@ class TestMCPImageQuery:
assert not searched
class TestMCPImageInput:
@pytest.mark.asyncio
async def test_ask_question_decodes_images(self, mcp_db, monkeypatch):
from base64 import b64encode
captured = {}
async def fake_ask(self, question, filter=None, images=None, sources=None):
captured["images"] = images
return ("answer", [])
monkeypatch.setattr(HaikuRAG, "ask", fake_ask)
mcp = create_mcp_server(mcp_db)
ask = await _get_tool(mcp, "ask_question")
png = b"fake image bytes"
result = await ask(question="q", images_base64=[b64encode(png).decode()])
assert result == "answer"
assert captured["images"] == [png]
@pytest.mark.asyncio
async def test_analyze_decodes_images(self, mcp_db, monkeypatch):
from base64 import b64encode
from types import SimpleNamespace
captured = {}
async def fake_analyze(self, question, filter=None, images=None, sources=None):
captured["images"] = images
return SimpleNamespace(answer="answer")
monkeypatch.setattr(HaikuRAG, "analyze", fake_analyze)
mcp = create_mcp_server(mcp_db)
analyze = await _get_tool(mcp, "analyze")
jpeg = b"fake jpeg bytes"
result = await analyze(question="q", images_base64=[b64encode(jpeg).decode()])
assert result == "answer"
assert captured["images"] == [jpeg]
@pytest.mark.asyncio
async def test_ask_question_without_images_passes_none(self, mcp_db, monkeypatch):
captured = {}
async def fake_ask(self, question, filter=None, images=None, sources=None):
captured["images"] = images
return ("answer", [])
monkeypatch.setattr(HaikuRAG, "ask", fake_ask)
mcp = create_mcp_server(mcp_db)
ask = await _get_tool(mcp, "ask_question")
result = await ask(question="q")
assert result == "answer"
assert captured["images"] is None
@pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning")
class TestMCPErrorContract:
"""A failure is an error on the wire, never an empty result. Expected
@ -1038,7 +1005,11 @@ class TestMCPErrorContract:
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_name,kwargs",
[("search_documents", {"query": "x"}), ("list_documents", {})],
[
("search_documents", {"query": "x"}),
("list_documents", {}),
("execute_code", {"code": "print(1)"}),
],
)
async def test_an_invalid_filter_is_an_error_naming_the_filter(
self, mcp_db, tool_name, kwargs
@ -1076,39 +1047,28 @@ class TestMCPErrorContract:
@pytest.mark.parametrize(
"payload", ["!!! not base64 !!!", "é"], ids=["outside_alphabet", "non_ascii"]
)
@pytest.mark.parametrize(
"tool_name,image_param,many",
[
("search_documents_by_image", "image_base64", False),
("ask_question", "images_base64", True),
("analyze", "images_base64", True),
],
)
async def test_invalid_base64_is_an_error(
self, mcp_db, multimodal_embedder, tool_name, image_param, many, payload
self, mcp_db, multimodal_embedder, payload
):
kwargs: dict[str, object] = {"question": "q"} if many else {}
kwargs[image_param] = [payload] if many else payload
result = await _call(create_mcp_server(mcp_db), tool_name, **kwargs)
result = await _call(
create_mcp_server(mcp_db), "search_documents_by_image", image_base64=payload
)
assert result.is_error
assert "base64" in result.content[0].text
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client_method,tool_name",
[("ask", "ask_question"), ("analyze", "analyze")],
)
async def test_an_agent_failure_names_only_its_type(
self, mcp_db, monkeypatch, caplog, client_method, tool_name
async def test_a_host_failure_inside_a_program_names_only_its_type(
self, mcp_db, monkeypatch, caplog
):
async def boom(self, question, filter=None, images=None, sources=None):
async def boom(self, *args, **kwargs):
raise RuntimeError("boom at /secret/path")
monkeypatch.setattr(HaikuRAG, client_method, boom)
with caplog.at_level(logging.ERROR, logger="haiku.rag.mcp"):
result = await _call(create_mcp_server(mcp_db), tool_name, question="q")
monkeypatch.setattr(HaikuRAG, "search", boom)
with caplog.at_level(logging.ERROR, logger="haiku.rag.sandbox.sandbox"):
result = await _call(
create_mcp_server(mcp_db), "execute_code", code="await search('x')"
)
assert result.is_error
assert "RuntimeError" in result.content[0].text

View file

@ -662,150 +662,6 @@ def test_format_bytes():
assert format_bytes(1125899906842624) == "1.0 PB"
# --- format_citations tests ---
def test_format_citations_empty():
from haiku.rag.utils import format_citations
assert format_citations([]) == ""
def test_format_citations_with_citation():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
document_title="Test Doc",
content="Some content",
page_numbers=[1],
headings=["Intro"],
)
result = format_citations([citation])
assert "[1] Test Doc" in result
assert "doc1" not in result
assert "chunk1" not in result
assert "test://doc" in result
assert "p. 1" in result
assert "Section: Intro" in result
assert "Some content" in result
def test_format_citations_multiple_pages():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
content="Content",
page_numbers=[1, 2, 3],
)
result = format_citations([citation])
assert "[1] test://doc" in result
assert "pp. 1-3" in result
# No title: the URI stands in, and the document id never leaks.
assert "doc1" not in result
def test_format_citations_with_index():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
index=5,
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
document_title="Test Doc",
content="Content",
)
result = format_citations([citation])
assert "[5] Test Doc" in result
def test_format_citations_sequential_indices():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citations = [
Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc1",
document_title="First",
content="Content 1",
),
Citation(
document_id="doc2",
chunk_id="chunk2",
document_uri="test://doc2",
document_title="Second",
content="Content 2",
),
]
result = format_citations(citations)
assert "[1] First" in result
assert "[2] Second" in result
def test_format_citations_names_the_source_when_asked():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
document_title="Test Doc",
content="Content",
source="papers",
)
assert "papers" in format_citations([citation], include_source=True)
assert "papers" not in format_citations([citation])
def test_format_citations_names_an_untitled_document_once():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
content="Content",
page_numbers=[3],
)
result = format_citations([citation])
assert result.count("test://doc") == 1
assert "[1] test://doc - p. 3" in result
# --- format_citations tests (pictures) ---
def test_format_citations_picture_refs_render_as_markers():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
document_title="Test Doc",
content="text body",
picture_refs=["#/pictures/0", "#/pictures/3"],
)
result = format_citations([citation])
assert "[Figure: #/pictures/0]" in result
assert "[Figure: #/pictures/3]" in result
# --- format_citations_rich tests ---
@ -846,6 +702,22 @@ async def test_format_citations_rich_header_and_footer():
assert "chunk: chunk-uuid-1" in output
async def test_format_citations_rich_names_a_single_page():
from haiku.rag.store.models.citation import Citation
from haiku.rag.utils import format_citations_rich
citation = Citation(
document_id="doc1",
chunk_id="chunk1",
document_uri="test://doc",
content="Body",
page_numbers=[3],
)
output = _render_rich(await format_citations_rich([citation]))
assert "p. 3" in output
assert "pp." not in output
async def test_format_citations_rich_names_the_database_when_federating():
"""Across databases, a citation has to say which one it came from."""
from unittest.mock import AsyncMock