From 3858ab905a2e8a0b3fce62e44f25ab3c4d8c4429 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 20 May 2026 11:37:35 +0300 Subject: [PATCH] toc.json nodes carry chunk_ids; fix cite to accept DB-resolvable chunk_ids --- haiku_rag_slim/haiku/rag/sandbox/sandbox.py | 79 +++++++++--- haiku_rag_slim/haiku/rag/skills/_tools.py | 58 ++++++--- .../haiku/rag/skills/rag-analysis/SKILL.md | 5 +- tests/sandbox/test_sandbox_toc.py | 121 ++++++++++++++++-- tests/skills/test_rag.py | 31 +++++ 5 files changed, 243 insertions(+), 51 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/sandbox/sandbox.py b/haiku_rag_slim/haiku/rag/sandbox/sandbox.py index 72b1a959..0bb9dc62 100644 --- a/haiku_rag_slim/haiku/rag/sandbox/sandbox.py +++ b/haiku_rag_slim/haiku/rag/sandbox/sandbox.py @@ -2,6 +2,7 @@ import asyncio import atexit import concurrent.futures import json +import os from collections.abc import Callable from dataclasses import dataclass from pathlib import Path @@ -37,7 +38,10 @@ def _run_async(coro: Any) -> Any: return _executor.submit(asyncio.run, coro).result() -def _build_toc(items: list["DocumentItem"]) -> list[dict[str, Any]]: +def _build_toc( + items: list["DocumentItem"], + chunk_index: dict[str, list[str]], +) -> list[dict[str, Any]]: """Build a nested section tree from items in position order. Each ``section_header`` with ``heading_level > 0`` becomes a node. Nesting @@ -49,6 +53,10 @@ def _build_toc(items: list["DocumentItem"]) -> list[dict[str, Any]]: the next sibling or ancestor that ends this section), or the total item count if no such header exists. + ``chunk_ids`` aggregates the chunks covered by all items in the section's + ``item_range`` (deduped, order preserved). Pass directly to ``cite()`` to + ground a section-scoped answer without a corpus-wide ``search()`` call. + Items without a section_header label (or with ``heading_level == 0``) are skipped. When all section_headers carry the same level the output is a flat sibling list (see docling-project/docling#2121 for an upstream case @@ -61,6 +69,7 @@ def _build_toc(items: list["DocumentItem"]) -> list[dict[str, Any]]: return [] total = max((i.position for i in items), default=-1) + 1 + items_by_position: dict[int, DocumentItem] = {i.position: i for i in items} ends: list[int] = [] for idx, h in enumerate(headers): @@ -74,13 +83,23 @@ def _build_toc(items: list["DocumentItem"]) -> list[dict[str, Any]]: roots: list[dict[str, Any]] = [] stack: list[tuple[int, dict[str, Any]]] = [] for h, end in zip(headers, ends, strict=True): + seen: set[str] = set() + chunk_ids: list[str] = [] + for pos in range(h.position, end): + item = items_by_position.get(pos) + if item is None: + continue + for cid in chunk_index.get(item.self_ref, []): + if cid not in seen: + seen.add(cid) + chunk_ids.append(cid) node: dict[str, Any] = { "self_ref": h.self_ref, "level": h.heading_level, "title": h.text, - "position": h.position, "page_numbers": list(h.page_numbers), "item_range": [h.position, end], + "chunk_ids": chunk_ids, "children": [], } while stack and stack[-1][0] >= h.heading_level: @@ -111,6 +130,7 @@ class Sandbox: _context: AnalysisContext _search_results: "list[SearchResult]" _doc_items: dict[str, list["DocumentItem"]] + _doc_chunk_index: dict[str, dict[str, list[str]]] _items_jsonl_cache: dict[str, str] _toc_json_cache: dict[str, str] _repl: MontyRepl | None @@ -127,6 +147,7 @@ class Sandbox: self._context = context self._search_results = [] self._doc_items = {} + self._doc_chunk_index = {} self._items_jsonl_cache = {} self._toc_json_cache = {} self._repl = None @@ -233,6 +254,27 @@ class Sandbox: sandbox._doc_items[did] = items return items + def _get_chunk_index(did: str) -> dict[str, list[str]]: + """Fetch the self_ref → chunk_ids index for one doc, cached.""" + cached = sandbox._doc_chunk_index.get(did) + if cached is not None: + return cached + + async def _fetch() -> dict[str, list[str]]: + from haiku.rag.client import HaikuRAG + + async with HaikuRAG(db_path, config=config, read_only=True) as rag: + index = ( + await rag.chunk_repository.get_chunk_ids_by_self_ref_grouped( + [did] + ) + ) + return index.get(did, {}) + + chunk_index = _run_async(_fetch()) + sandbox._doc_chunk_index[did] = chunk_index + return chunk_index + def _make_items_reader( did: str, ) -> Callable[["PurePosixPath"], str]: @@ -241,17 +283,7 @@ class Sandbox: if cached is not None: return cached items = _get_items(did) - - async def _fetch_chunks() -> dict[str, list[str]]: - from haiku.rag.client import HaikuRAG - - async with HaikuRAG(db_path, config=config, read_only=True) as rag: - index = await rag.chunk_repository.get_chunk_ids_by_self_ref_grouped( - [did] - ) - return index.get(did, {}) - - doc_chunk_index = _run_async(_fetch_chunks()) + chunk_index = _get_chunk_index(did) jsonl = "\n".join( json.dumps( { @@ -260,7 +292,7 @@ class Sandbox: "text": item.text, "page_numbers": item.page_numbers, "heading_level": item.heading_level, - "chunk_ids": doc_chunk_index.get(item.self_ref, []), + "chunk_ids": chunk_index.get(item.self_ref, []), }, ensure_ascii=False, ) @@ -279,11 +311,12 @@ class Sandbox: if cached is not None: return cached items = _get_items(did) + chunk_index = _get_chunk_index(did) toc = json.dumps( { "doc_id": did, "title": doc_titles.get(did), - "tree": _build_toc(items), + "tree": _build_toc(items, chunk_index), }, ensure_ascii=False, ) @@ -340,13 +373,17 @@ class Sandbox: write=_deny_write, ) ) - files.append( - CallbackFile( - f"{doc_dir}/toc.json", - read=_make_toc_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. + if not os.environ.get("HAIKU_RAG_DISABLE_TOC"): + files.append( + CallbackFile( + f"{doc_dir}/toc.json", + read=_make_toc_reader(doc_id), + write=_deny_write, + ) ) - ) return OSAccess(files) diff --git a/haiku_rag_slim/haiku/rag/skills/_tools.py b/haiku_rag_slim/haiku/rag/skills/_tools.py index 51d438b3..89155d18 100644 --- a/haiku_rag_slim/haiku/rag/skills/_tools.py +++ b/haiku_rag_slim/haiku/rag/skills/_tools.py @@ -295,42 +295,64 @@ def create_skill_tools( async def cite(ctx: RunContext[RAGRunDeps], chunk_ids: list[str]) -> str: """Register chunk IDs as citations for your answer. - Call this after searching, with the chunk_id values from search - results that support your answer. + Accepts chunk_ids from search results AND from direct file reads + (items.jsonl, toc.json). Verbatim copies only — chunk_ids that + don't exist in the database trigger a retry. Args: - chunk_ids: List of chunk_id values from search results. + chunk_ids: List of chunk_id values from search results or VFS reads. """ from haiku.rag.agents.research.models import resolve_citations + from haiku.rag.store.models.chunk import SearchResult state = _get_state(ctx, state_type) if not state: return "No state available." - all_results = [] + if not chunk_ids: + return "Registered 0 citations (empty chunk_ids)." + + all_results: list[SearchResult] = [] for results_list in state.searches.values(): all_results.extend(results_list) citations = resolve_citations(chunk_ids, all_results) + resolved_ids = {c.chunk_id for c in citations} + missing = [ + cid.strip("[]") + for cid in chunk_ids + if cid.strip("[]") not in resolved_ids + ] + + rag = ctx.deps.rag if ctx.deps else None + if missing and rag is not None: + synthetic: list[SearchResult] = [] + doc_cache: dict[str, Any] = {} + for cid in missing: + chunk = await rag.get_chunk_by_id(cid) + if chunk is None or not chunk.document_id: + continue + did = chunk.document_id + if did in doc_cache: + doc = doc_cache[did] + else: + doc = await rag.get_document_by_id(did) + doc_cache[did] = doc + chunk.document_uri = doc.uri if doc else None + chunk.document_title = doc.title if doc else None + synthetic.append(SearchResult.from_chunk(chunk, score=1.0)) + if synthetic: + citations.extend(resolve_citations(missing, synthetic)) + if citations: _register_citations(state, citations) return f"Registered {len(citations)} citation(s)." - if not chunk_ids: - return "Registered 0 citations (empty chunk_ids)." - - if not any(r.chunk_id for r in all_results): - raise ModelRetry( - f"None of the supplied chunk_ids {list(chunk_ids)} can be " - "resolved: no search results have been recorded in this " - "session yet. Call `search` first, then cite chunk_ids " - "from its response." - ) raise ModelRetry( - f"None of the supplied chunk_ids {list(chunk_ids)} match a " - "chunk_id from search results. Copy chunk_ids verbatim from " - "the search response — never reconstruct, abbreviate, or " - "paraphrase them." + f"None of the supplied chunk_ids {list(chunk_ids)} could be " + "resolved. Copy chunk_ids verbatim from `search` results or " + "from the `chunk_ids` field on items.jsonl / toc.json rows — " + "never reconstruct, abbreviate, or paraphrase them." ) tools["cite"] = cite diff --git a/haiku_rag_slim/haiku/rag/skills/rag-analysis/SKILL.md b/haiku_rag_slim/haiku/rag/skills/rag-analysis/SKILL.md index aaba38b0..f024098a 100644 --- a/haiku_rag_slim/haiku/rag/skills/rag-analysis/SKILL.md +++ b/haiku_rag_slim/haiku/rag/skills/rag-analysis/SKILL.md @@ -92,7 +92,7 @@ Each row carries: - `heading_level`: H-level for `section_header` rows; `0` on non-header rows ### toc.json -Section tree derived from `heading_level`: `{"doc_id", "title", "tree": [...]}` where each node has `{self_ref, level, title, position, page_numbers, item_range: [start, end_exclusive], children}`. `item_range` is a line slice into `items.jsonl` — `items[start:end]`. `tree: []` for docs with no headers. +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 `cite()` to ground a section-scoped answer without a corpus-wide `search()` call. `tree: []` for docs with no headers. ### Cross-referencing search results with items Search results include `doc_item_refs` (e.g. `["#/texts/48", "#/tables/0"]`) that correspond to `self_ref` values in `items.jsonl`. To find which section a hit lives in: locate the item by `self_ref`, take its line index, and walk `toc.json` to find the deepest node whose `item_range` contains that index. @@ -102,7 +102,8 @@ Search results include `doc_item_refs` (e.g. `["#/texts/48", "#/tables/0"]`) tha 1. Search first. 2. If the top results contain the answer, call `cite` with the supporting chunk_ids and write a concise answer. 3. Reach for `execute_code` when search results are insufficient or when the task requires computation, aggregation, traversal across documents, or section-scoped reading. From inside code you can search again with different terms, or read `items.jsonl` / `toc.json` / `content.txt` directly from the document filesystem. -4. Call `cite` with the chunk_ids that ground your answer before writing the final response. +4. For questions about a *known document's* structure ("which section contains X", "list the sections of doc Y", "summarise section Z"), read `/documents/{id}/toc.json` first. Each node carries `item_range` (a slice into `items.jsonl`) and `chunk_ids` (citable). Prefer this over `search()` for in-document navigation — `search()` ranks across the whole corpus and can return chunks from unrelated documents. +5. Call `cite` with the chunk_ids that ground your answer before writing the final response. You MUST call `cite` with at least one chunk ID before producing your final answer, **unless** you are refusing for lack of information. Answers without citations are considered ungrounded. In a refusal case do **not** call `cite` — there is nothing to cite. diff --git a/tests/sandbox/test_sandbox_toc.py b/tests/sandbox/test_sandbox_toc.py index 494e8b6f..27a4e2a4 100644 --- a/tests/sandbox/test_sandbox_toc.py +++ b/tests/sandbox/test_sandbox_toc.py @@ -186,12 +186,101 @@ class TestTocShape: assert titles == ["Real H1"] assert toc["tree"][0]["item_range"] == [0, 4] + async def test_node_shape_has_chunk_ids_not_position(self, temp_db_path): + async with HaikuRAG(temp_db_path, create=True) as client: + doc_id = await _empty_doc(client, uri="test://shape", title="Shape") + await client.document_item_repository.create_items( + doc_id, [_header(doc_id, 0, 1, "Only"), _para(doc_id, 1)] + ) + + sandbox = Sandbox(temp_db_path, AppConfig(), AnalysisContext()) + toc = await _read_toc(sandbox, doc_id) + node = toc["tree"][0] + expected = { + "self_ref", + "level", + "title", + "page_numbers", + "item_range", + "chunk_ids", + "children", + } + assert expected <= set(node) + assert "position" not in node + assert node["chunk_ids"] == [] + + +@pytest.mark.asyncio +class TestTocChunkIdsAggregation: + """toc.json nodes carry the union of chunk_ids covered by their item_range.""" + + async def test_chunk_ids_union_over_item_range(self, temp_db_path, monkeypatch): + async with HaikuRAG(temp_db_path, create=True) as client: + doc_id = await _empty_doc(client, uri="test://chunks", title="Chunks") + await client.document_item_repository.create_items( + doc_id, + [ + _header(doc_id, 0, 1, "Intro"), + _para(doc_id, 1), + _para(doc_id, 2), + _header(doc_id, 3, 2, "Background"), + _para(doc_id, 4), + _header(doc_id, 5, 1, "Methods"), + _para(doc_id, 6), + ], + ) + + from haiku.rag.store.repositories.chunk import ChunkRepository + + # self_ref → list[chunk_id]. Intro covers #/texts/0..2, Background + # covers #/texts/3..4, Methods covers #/texts/5..6. Item at #/texts/2 + # belongs to two chunks (cA + cB) — verifying dedup-preserving-order. + fake_index = { + "#/texts/1": ["cA"], + "#/texts/2": ["cA", "cB"], + "#/texts/3": ["cB"], + "#/texts/4": ["cC"], + "#/texts/5": ["cD"], + "#/texts/6": ["cD"], + } + + async def fake_grouped(self, document_ids): + return {doc_id: fake_index} + + monkeypatch.setattr( + ChunkRepository, + "get_chunk_ids_by_self_ref_grouped", + fake_grouped, + ) + + sandbox = Sandbox(temp_db_path, AppConfig(), AnalysisContext()) + toc = await _read_toc(sandbox, doc_id) + + intro = toc["tree"][0] + assert intro["title"] == "Intro" + # Intro spans positions 0..2 (header) plus its child Background's + # range — item_range is [0, 5]. + assert intro["item_range"] == [0, 5] + assert intro["chunk_ids"] == ["cA", "cB", "cC"] + + background = intro["children"][0] + assert background["title"] == "Background" + assert background["item_range"] == [3, 5] + assert background["chunk_ids"] == ["cB", "cC"] + + methods = toc["tree"][1] + assert methods["title"] == "Methods" + assert methods["item_range"] == [5, 7] + assert methods["chunk_ids"] == ["cD"] + @pytest.mark.asyncio class TestTocCaching: - """items + toc reads for a doc share one items fetch; repeat reads hit cache.""" + """items + toc reads for a doc share one items fetch and one chunk-index fetch.""" - async def test_items_fetched_once_per_doc(self, temp_db_path, monkeypatch): + async def test_items_and_chunk_index_fetched_once_per_doc( + self, temp_db_path, monkeypatch + ): async with HaikuRAG(temp_db_path, create=True) as client: doc_id = await _empty_doc(client, uri="test://cache", title="Cache") await client.document_item_repository.create_items( @@ -201,24 +290,36 @@ class TestTocCaching: sandbox = Sandbox(temp_db_path, AppConfig(), AnalysisContext()) + from haiku.rag.store.repositories.chunk import ChunkRepository from haiku.rag.store.repositories.document_item import DocumentItemRepository - call_count = {"n": 0} - original = DocumentItemRepository.get_all_items + items_calls = {"n": 0} + chunk_calls = {"n": 0} + original_items = DocumentItemRepository.get_all_items + original_chunks = ChunkRepository.get_chunk_ids_by_self_ref_grouped - async def counting(self, document_id): - call_count["n"] += 1 - return await original(self, document_id) + async def counting_items(self, document_id): + items_calls["n"] += 1 + return await original_items(self, document_id) - monkeypatch.setattr(DocumentItemRepository, "get_all_items", counting) + async def counting_chunks(self, document_ids): + chunk_calls["n"] += 1 + return await original_chunks(self, document_ids) - # Items + toc share `_doc_items`. Four reads → one items fetch. + monkeypatch.setattr(DocumentItemRepository, "get_all_items", counting_items) + monkeypatch.setattr( + ChunkRepository, "get_chunk_ids_by_self_ref_grouped", counting_chunks + ) + + # Items + toc share `_doc_items` and `_doc_chunk_index`. Four reads → + # one items fetch + one chunk-index fetch. _ = await _read_toc(sandbox, doc_id) _ = await _read_items_jsonl(sandbox, doc_id) _ = await _read_toc(sandbox, doc_id) _ = await _read_items_jsonl(sandbox, doc_id) - assert call_count["n"] == 1 + assert items_calls["n"] == 1 + assert chunk_calls["n"] == 1 @pytest.mark.asyncio diff --git a/tests/skills/test_rag.py b/tests/skills/test_rag.py index 6009bfd6..73a9c7f7 100644 --- a/tests/skills/test_rag.py +++ b/tests/skills/test_rag.py @@ -373,6 +373,37 @@ class TestCiteTool: result = await cite(ctx, chunk_ids=[]) assert "0" in result + async def test_cite_accepts_chunk_id_from_db_without_prior_search( + self, rag_db, rag_client + ): + """chunk_ids sourced from items.jsonl / toc.json are valid citations. + + The skill calls cite directly with chunk_ids it read from the VFS; + no search() has been recorded in state.searches. cite must look the + chunk up in the DB and build a Citation with full document context. + """ + from haiku.rag.skills.rag import RAGState, create_skill + + docs = await rag_client.list_documents(limit=1) + assert docs, "fixture should have at least one document" + doc_id = docs[0].id + chunks = await rag_client.chunk_repository.get_by_document_id(doc_id) + assert chunks, "fixture document should have chunks" + chunk_id = chunks[0].id + + skill = create_skill(db_path=rag_db) + cite = _get_tool(skill, "cite") + state = RAGState() + ctx = _make_ctx(state, rag=rag_client) + assert not state.searches, "this test exercises the no-prior-search path" + + result = await cite(ctx, chunk_ids=[chunk_id]) + assert "Registered 1 citation(s)." == result + assert chunk_id in state.citations + registered = state.citation_index[chunk_id] + assert registered.document_id == doc_id + assert registered.document_uri # uri must be populated from doc lookup + class TestLifespan: async def test_opens_one_client_per_invocation(self, rag_db):