From b374d5eb83e7ffa54925afcca62f0b8a8f8c5d50 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 7 Sep 2026 10:02:42 +0300 Subject: [PATCH] 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. --- CHANGELOG.md | 32 +- README.md | 2 +- claude-plugin/skills/haiku-rag/SKILL.md | 21 +- docs/capabilities/analysis.md | 2 +- docs/cli.md | 3 - docs/configuration/qa.md | 2 +- docs/mcp.md | 33 +- haiku_rag_slim/haiku/rag/app.py | 3 +- .../haiku/rag/capabilities/analysis.py | 19 +- .../rag/capabilities/instructions/analysis.md | 12 +- haiku_rag_slim/haiku/rag/cli.py | 9 +- haiku_rag_slim/haiku/rag/mcp.py | 141 +++----- haiku_rag_slim/haiku/rag/sandbox/__init__.py | 3 +- haiku_rag_slim/haiku/rag/sandbox/sandbox.py | 95 ++++- haiku_rag_slim/haiku/rag/utils.py | 42 --- tests/sandbox/test_sandbox.py | 231 +++++++++++- tests/sandbox/test_sandbox_toc.py | 33 ++ tests/test_app.py | 14 - tests/test_cli.py | 8 - tests/test_mcp.py | 336 ++++++++---------- tests/test_utils.py | 160 +-------- 21 files changed, 631 insertions(+), 570 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 64343190..fc815f0c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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(" failed: ")`; 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 diff --git a/README.md b/README.md index fb3c8d5b..24d832e3 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/claude-plugin/skills/haiku-rag/SKILL.md b/claude-plugin/skills/haiku-rag/SKILL.md index a0a0d57b..981baf43 100644 --- a/claude-plugin/skills/haiku-rag/SKILL.md +++ b/claude-plugin/skills/haiku-rag/SKILL.md @@ -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 diff --git a/docs/capabilities/analysis.md b/docs/capabilities/analysis.md index b3a58138..604eb3d3 100644 --- a/docs/capabilities/analysis.md +++ b/docs/capabilities/analysis.md @@ -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 diff --git a/docs/cli.md b/docs/cli.md index b54cb8fe..5e0efae8 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -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 diff --git a/docs/configuration/qa.md b/docs/configuration/qa.md index 7b059bd8..32c4c34b 100644 --- a/docs/configuration/qa.md +++ b/docs/configuration/qa.md @@ -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: diff --git a/docs/mcp.md b/docs/mcp.md index df7b4448..9fef2c7d 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -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. diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index d2dbedad..8bc934ae 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -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() diff --git a/haiku_rag_slim/haiku/rag/capabilities/analysis.py b/haiku_rag_slim/haiku/rag/capabilities/analysis.py index fb88a081..2dafbbe8 100644 --- a/haiku_rag_slim/haiku/rag/capabilities/analysis.py +++ b/haiku_rag_slim/haiku/rag/capabilities/analysis.py @@ -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." diff --git a/haiku_rag_slim/haiku/rag/capabilities/instructions/analysis.md b/haiku_rag_slim/haiku/rag/capabilities/instructions/analysis.md index cd0a6bea..37dc8807 100644 --- a/haiku_rag_slim/haiku/rag/capabilities/instructions/analysis.md +++ b/haiku_rag_slim/haiku/rag/capabilities/instructions/analysis.md @@ -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. diff --git a/haiku_rag_slim/haiku/rag/cli.py b/haiku_rag_slim/haiku/rag/cli.py index 06647670..c47d3971 100644 --- a/haiku_rag_slim/haiku/rag/cli.py +++ b/haiku_rag_slim/haiku/rag/cli.py @@ -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__": diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py index 112a156c..ebfc533b 100644 --- a/haiku_rag_slim/haiku/rag/mcp.py +++ b/haiku_rag_slim/haiku/rag/mcp.py @@ -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 diff --git a/haiku_rag_slim/haiku/rag/sandbox/__init__.py b/haiku_rag_slim/haiku/rag/sandbox/__init__.py index 1ca4b1de..0e7576f3 100644 --- a/haiku_rag_slim/haiku/rag/sandbox/__init__.py +++ b/haiku_rag_slim/haiku/rag/sandbox/__init__.py @@ -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", ] diff --git a/haiku_rag_slim/haiku/rag/sandbox/sandbox.py b/haiku_rag_slim/haiku/rag/sandbox/sandbox.py index cddbdd46..47149f1e 100644 --- a/haiku_rag_slim/haiku/rag/sandbox/sandbox.py +++ b/haiku_rag_slim/haiku/rag/sandbox/sandbox.py @@ -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. diff --git a/haiku_rag_slim/haiku/rag/utils.py b/haiku_rag_slim/haiku/rag/utils.py index 2d71b47c..6327b808 100644 --- a/haiku_rag_slim/haiku/rag/utils.py +++ b/haiku_rag_slim/haiku/rag/utils.py @@ -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: ]`` 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.""" diff --git a/tests/sandbox/test_sandbox.py b/tests/sandbox/test_sandbox.py index ed4a8a2d..cb6a1946 100644 --- a/tests/sandbox/test_sandbox.py +++ b/tests/sandbox/test_sandbox.py @@ -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.""" diff --git a/tests/sandbox/test_sandbox_toc.py b/tests/sandbox/test_sandbox_toc.py index 2d359212..d08b9340 100644 --- a/tests/sandbox/test_sandbox_toc.py +++ b/tests/sandbox/test_sandbox_toc.py @@ -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.""" diff --git a/tests/test_app.py b/tests/test_app.py index 589c2eb1..cf7104da 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -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 diff --git a/tests/test_cli.py b/tests/test_cli.py index 7d4af634..50d389a0 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -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): diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 0fada9a3..3954625a 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -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 diff --git a/tests/test_utils.py b/tests/test_utils.py index 9d4cba45..f63bf07e 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -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