toc.json nodes carry chunk_ids; fix cite to accept DB-resolvable chunk_ids
This commit is contained in:
parent
04eeec77d2
commit
3858ab905a
5 changed files with 243 additions and 51 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Reference in a new issue