Properly expand context in the case of docling documents and merge chunks and meta
This commit is contained in:
parent
ac439b6c69
commit
16a97cc140
4 changed files with 243 additions and 201 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Reference in a new issue