Properly expand context in the case of docling documents and merge chunks and meta

This commit is contained in:
Yiorgis Gozadinos 2025-12-01 14:53:11 +02:00
parent ac439b6c69
commit 16a97cc140
No known key found for this signature in database
4 changed files with 243 additions and 201 deletions

View file

@ -679,6 +679,32 @@ class HaikuRAG:
return expanded_results
def _merge_ranges(
self, ranges: list[tuple[int, int, SearchResult]]
) -> list[tuple[int, int, list[SearchResult]]]:
"""Merge overlapping or adjacent ranges."""
if not ranges:
return []
sorted_ranges = sorted(ranges, key=lambda x: x[0])
merged: list[tuple[int, int, list[SearchResult]]] = []
cur_min, cur_max, cur_results = (
sorted_ranges[0][0],
sorted_ranges[0][1],
[sorted_ranges[0][2]],
)
for min_idx, max_idx, result in sorted_ranges[1:]:
if cur_max >= min_idx - 1: # Overlapping or adjacent
cur_max = max(cur_max, max_idx)
cur_results.append(result)
else:
merged.append((cur_min, cur_max, cur_results))
cur_min, cur_max, cur_results = min_idx, max_idx, [result]
merged.append((cur_min, cur_max, cur_results))
return merged
async def _expand_with_docling(
self,
results: list[SearchResult],
@ -688,75 +714,51 @@ class HaikuRAG:
"""Expand results using DoclingDocument structure."""
from haiku.rag.store.models.chunk import BoundingBox
# Build index of all DocItems for expansion
all_items = list(docling_doc.iterate_items())
ref_to_index: dict[str, int] = {}
ref_to_index = {
getattr(item, "self_ref", None): i
for i, (item, _) in enumerate(all_items)
if getattr(item, "self_ref", None)
}
# Map refs to indices
for i, (_, item) in enumerate(all_items):
self_ref = getattr(item, "self_ref", None)
if self_ref:
ref_to_index[self_ref] = i
expanded_results = []
# Compute expanded ranges
ranges: list[tuple[int, int, SearchResult]] = []
passthrough: list[SearchResult] = []
for result in results:
if not result.doc_item_refs:
expanded_results.append(result)
continue
# Find indices of all refs in this result
indices = []
for ref in result.doc_item_refs:
if ref in ref_to_index:
indices.append(ref_to_index[ref])
indices = [
ref_to_index[r] for r in result.doc_item_refs if r in ref_to_index
]
if not indices:
expanded_results.append(result)
passthrough.append(result)
continue
# Expand range
min_idx = max(0, min(indices) - radius)
max_idx = min(len(all_items) - 1, max(indices) + radius)
ranges.append((min_idx, max_idx, result))
# Collect expanded DocItems
expanded_content_parts = []
expanded_refs = []
expanded_page_numbers: set[int] = set()
expanded_labels: set[str] = set()
expanded_bboxes: list[BoundingBox] = []
# Merge overlapping ranges
merged = self._merge_ranges(ranges)
final_results: list[SearchResult] = []
for min_idx, max_idx, original_results in merged:
content_parts, refs, pages, labels, bboxes = [], [], set(), set(), []
for i in range(min_idx, max_idx + 1):
_, item = all_items[i]
# Get content
text = getattr(item, "text", None)
if text:
expanded_content_parts.append(text)
# Get self_ref
self_ref = getattr(item, "self_ref", None)
if self_ref:
expanded_refs.append(self_ref)
# Get label
label = getattr(item, "label", None)
if label:
expanded_labels.add(
item, _ = all_items[i]
if text := getattr(item, "text", None):
content_parts.append(text)
if self_ref := getattr(item, "self_ref", None):
refs.append(self_ref)
if label := getattr(item, "label", None):
labels.add(
str(label.value) if hasattr(label, "value") else str(label)
)
# Get provenance (page numbers and bounding boxes)
prov = getattr(item, "prov", None)
if prov:
for prov_item in prov:
page_no = getattr(prov_item, "page_no", None)
if page_no is not None:
expanded_page_numbers.add(page_no)
bbox = getattr(prov_item, "bbox", None)
if bbox is not None:
expanded_bboxes.append(
if prov := getattr(item, "prov", None):
for p in prov:
if (page_no := getattr(p, "page_no", None)) is not None:
pages.add(page_no)
if bbox := getattr(p, "bbox", None):
bboxes.append(
BoundingBox(
page_no=page_no or 0,
left=bbox.l,
@ -766,23 +768,30 @@ class HaikuRAG:
)
)
expanded_results.append(
# Merge headings preserving order
all_headings: list[str] = []
for r in original_results:
if r.headings:
all_headings.extend(h for h in r.headings if h not in all_headings)
first = original_results[0]
final_results.append(
SearchResult(
content="\n\n".join(expanded_content_parts),
score=result.score,
chunk_id=result.chunk_id,
document_id=result.document_id,
document_uri=result.document_uri,
document_title=result.document_title,
doc_item_refs=expanded_refs,
page_numbers=sorted(expanded_page_numbers),
headings=result.headings,
labels=sorted(expanded_labels),
bounding_boxes=expanded_bboxes if expanded_bboxes else None,
content="\n\n".join(content_parts),
score=max(r.score for r in original_results),
chunk_id=first.chunk_id,
document_id=first.document_id,
document_uri=first.document_uri,
document_title=first.document_title,
doc_item_refs=refs,
page_numbers=sorted(pages),
headings=all_headings or None,
labels=sorted(labels),
bounding_boxes=bboxes or None,
)
)
return expanded_results
return final_results + passthrough
async def _expand_with_chunks(
self,
@ -791,154 +800,54 @@ class HaikuRAG:
radius: int,
) -> list[SearchResult]:
"""Expand results using chunk-based adjacency."""
# Fetch all chunks for this document
all_chunks = await self.chunk_repository.get_by_document_id(doc_id)
if not all_chunks:
return results
# Build content -> chunk mapping and order -> chunk mapping
content_to_chunk = {c.content: c for c in all_chunks}
chunk_by_order = {c.order: c for c in all_chunks}
max_order = max(chunk_by_order.keys())
min_order = min(chunk_by_order.keys())
min_order, max_order = min(chunk_by_order.keys()), max(chunk_by_order.keys())
# Build expanded ranges for merging
expanded_ranges = []
# Build ranges
ranges: list[tuple[int, int, SearchResult]] = []
passthrough: list[SearchResult] = []
for result in results:
# Find matching chunk by content
matching_chunk = content_to_chunk.get(result.content)
if matching_chunk is None:
expanded_ranges.append(
{
"original_result": result,
"score": result.score,
"min_order": -1,
"max_order": -1,
"chunks": [],
}
)
chunk = content_to_chunk.get(result.content)
if chunk is None:
passthrough.append(result)
continue
start = max(min_order, chunk.order - radius)
end = min(max_order, chunk.order + radius)
ranges.append((start, end, result))
# Calculate range
start_order = max(min_order, matching_chunk.order - radius)
end_order = min(max_order, matching_chunk.order + radius)
range_chunks = [
# Merge and build results
final_results: list[SearchResult] = []
for min_idx, max_idx, original_results in self._merge_ranges(ranges):
# Collect chunks in order
chunks_in_range = [
chunk_by_order[o]
for o in range(start_order, end_order + 1)
for o in range(min_idx, max_idx + 1)
if o in chunk_by_order
]
expanded_ranges.append(
{
"original_result": result,
"score": result.score,
"min_order": start_order,
"max_order": end_order,
"chunks": range_chunks,
}
)
# Merge overlapping ranges
merged_ranges = self._merge_chunk_ranges(expanded_ranges)
# Convert to SearchResults
expanded_results = []
for merged in merged_ranges:
if not merged["chunks"]:
# No chunks found, return original
expanded_results.append(merged["original_results"][0])
continue
combined_content = "".join(c.content for c in merged["chunks"])
original = merged["original_results"][0]
best_score = max(merged["scores"])
expanded_results.append(
first = original_results[0]
final_results.append(
SearchResult(
content=combined_content,
score=best_score,
chunk_id=original.chunk_id,
document_id=original.document_id,
document_uri=original.document_uri,
document_title=original.document_title,
doc_item_refs=original.doc_item_refs,
page_numbers=original.page_numbers,
headings=original.headings,
labels=original.labels,
bounding_boxes=original.bounding_boxes,
content="".join(c.content for c in chunks_in_range),
score=max(r.score for r in original_results),
chunk_id=first.chunk_id,
document_id=first.document_id,
document_uri=first.document_uri,
document_title=first.document_title,
doc_item_refs=first.doc_item_refs,
page_numbers=first.page_numbers,
headings=first.headings,
labels=first.labels,
bounding_boxes=first.bounding_boxes,
)
)
return expanded_results
def _merge_chunk_ranges(self, expanded_ranges: list[dict]) -> list[dict]:
"""Merge overlapping or adjacent chunk ranges."""
# Filter out ranges without chunks
valid_ranges = [r for r in expanded_ranges if r["chunks"]]
invalid_ranges = [r for r in expanded_ranges if not r["chunks"]]
if not valid_ranges:
return [
{
"original_results": [r["original_result"]],
"scores": [r["score"]],
"chunks": [],
}
for r in invalid_ranges
]
# Sort by min_order
sorted_ranges = sorted(valid_ranges, key=lambda x: x["min_order"])
merged = []
current = {
"min_order": sorted_ranges[0]["min_order"],
"max_order": sorted_ranges[0]["max_order"],
"original_results": [sorted_ranges[0]["original_result"]],
"scores": [sorted_ranges[0]["score"]],
"chunks": sorted_ranges[0]["chunks"],
}
for range_info in sorted_ranges[1:]:
# Check if ranges overlap or are adjacent
if current["max_order"] >= range_info["min_order"] - 1:
# Merge ranges
current["max_order"] = max(
current["max_order"], range_info["max_order"]
)
current["original_results"].append(range_info["original_result"])
current["scores"].append(range_info["score"])
# Merge chunks and deduplicate by order
chunks_dict = {c.order: c for c in current["chunks"]}
for chunk in range_info["chunks"]:
chunks_dict[chunk.order] = chunk
current["chunks"] = [chunks_dict[o] for o in sorted(chunks_dict.keys())]
else:
merged.append(current)
current = {
"min_order": range_info["min_order"],
"max_order": range_info["max_order"],
"original_results": [range_info["original_result"]],
"scores": [range_info["score"]],
"chunks": range_info["chunks"],
}
merged.append(current)
# Add back invalid ranges
for r in invalid_ranges:
merged.append(
{
"original_results": [r["original_result"]],
"scores": [r["score"]],
"chunks": [],
}
)
return merged
return final_results + passthrough
async def ask(
self, question: str, cite: bool = False, system_prompt: str | None = None

View file

@ -63,7 +63,6 @@ class QuestionAnswerAgent:
"""Search the knowledge base for relevant documents."""
results = await ctx.deps.client.search(query, limit=limit)
results = await ctx.deps.client.expand_context(results)
return [
ToolSearchResult(
content=r.content,

View file

@ -48,7 +48,7 @@ After your answer, include a "Citations:" section that lists:
- The document title if available (from the document_title field)
- Page number(s) if available (from the page_numbers field)
- Section heading if available (from the headings field)
- A VERBATIM excerpt (copy-paste exact text, first 50-100 characters) from the content field - do NOT summarize or paraphrase
- A VERBATIM excerpt (100-200 characters, copy-paste exact text) from the content field - do NOT summarize, truncate, or paraphrase
Example response format:
[Your answer here]

View file

@ -1027,6 +1027,140 @@ async def test_client_expand_context_keeps_separate_non_overlapping(temp_db_path
assert chunk5_expanded.score == 0.7
@pytest.mark.asyncio
async def test_client_expand_context_with_docling_merges_overlapping(temp_db_path):
"""Test that expand_context with DoclingDocument merges overlapping results."""
from haiku.rag.store.models import SearchResult
# Create a document with structured content that will have doc_item_refs
markdown_content = """# Chapter 1
This is paragraph one about topic A.
This is paragraph two about topic A continued.
This is paragraph three about topic B.
# Chapter 2
This is paragraph four about topic C.
"""
async with HaikuRAG(temp_db_path) as client:
doc = await client.create_document(
content=markdown_content,
uri="test://structured",
)
assert doc.id is not None
assert doc.docling_document_json is not None
# Get chunks which should have doc_item_refs
chunks = await client.chunk_repository.get_by_document_id(doc.id)
assert len(chunks) >= 1
# Find chunks that have doc_item_refs (from docling chunking)
chunks_with_refs = [c for c in chunks if c.get_chunk_metadata().doc_item_refs]
if len(chunks_with_refs) >= 2:
# Create search results from adjacent chunks
search_results = [
SearchResult.from_chunk(chunks_with_refs[0], 0.9),
SearchResult.from_chunk(chunks_with_refs[1], 0.8),
]
# Expand with radius that should cause overlap
expanded = await client.expand_context(search_results, radius=3)
# If chunks were adjacent, they should be merged
# The expanded results should have merged metadata
assert len(expanded) >= 1
# Check that expanded result has page_numbers populated
for r in expanded:
# Should have doc_item_refs from expansion
assert r.doc_item_refs is not None
@pytest.mark.asyncio
async def test_client_expand_context_docling_merges_metadata(temp_db_path):
"""Test that expand_context properly merges metadata from multiple results."""
from haiku.rag.store.models import SearchResult
markdown_content = """# Introduction
First paragraph of introduction.
Second paragraph of introduction.
# Methods
First paragraph of methods section.
Second paragraph of methods section.
# Results
First paragraph of results.
"""
async with HaikuRAG(temp_db_path) as client:
doc = await client.create_document(
content=markdown_content,
uri="test://metadata-merge",
)
assert doc.id is not None
chunks = await client.chunk_repository.get_by_document_id(doc.id)
chunks_with_refs = [c for c in chunks if c.get_chunk_metadata().doc_item_refs]
if len(chunks_with_refs) >= 2:
# Get chunks with different headings if possible
chunk1 = chunks_with_refs[0]
chunk2 = chunks_with_refs[-1] # Last chunk likely has different heading
search_results = [
SearchResult.from_chunk(chunk1, 0.9),
SearchResult.from_chunk(chunk2, 0.8),
]
# Expand with large radius to potentially merge
expanded = await client.expand_context(search_results, radius=10)
# Check that results have proper structure
for r in expanded:
# Content should be non-empty
assert len(r.content) > 0
# Score should be preserved (best score)
assert r.score in [0.9, 0.8]
@pytest.mark.asyncio
async def test_client_expand_context_docling_preserves_bounding_boxes(temp_db_path):
"""Test that expand_context preserves bounding boxes from DoclingDocument."""
from haiku.rag.store.models import SearchResult
async with HaikuRAG(temp_db_path) as client:
doc = await client.create_document(
content="# Test\n\nSome content here.",
uri="test://bboxes",
)
assert doc.id is not None
chunks = await client.chunk_repository.get_by_document_id(doc.id)
if chunks:
search_results = [SearchResult.from_chunk(chunks[0], 0.9)]
expanded = await client.expand_context(search_results, radius=2)
# Expanded results should exist
assert len(expanded) == 1
# Bounding boxes may or may not be present depending on document
# but the field should be accessible
_ = expanded[0].bounding_boxes
@pytest.mark.asyncio
async def test_client_create_document_stores_docling_json(temp_db_path):
"""Test that create_document stores DoclingDocument JSON."""