Extract search, expand_context, visualize_chunk into client/search.py
This commit is contained in:
parent
d4ecebf896
commit
12939f1444
2 changed files with 202 additions and 182 deletions
|
|
@ -862,89 +862,17 @@ class HaikuRAG:
|
||||||
search_type: str = "hybrid",
|
search_type: str = "hybrid",
|
||||||
filter: str | None = None,
|
filter: str | None = None,
|
||||||
) -> list[SearchResult]:
|
) -> list[SearchResult]:
|
||||||
"""Search for relevant chunks using the specified search method with optional reranking.
|
from haiku.rag.client.search import search
|
||||||
|
|
||||||
Args:
|
return await search(self, query, limit, search_type, filter)
|
||||||
query: The search query string.
|
|
||||||
limit: Maximum number of results to return. Defaults to config.search.default_limit.
|
|
||||||
search_type: Type of search - "vector", "fts", or "hybrid" (default).
|
|
||||||
filter: Optional SQL WHERE clause to filter documents before searching chunks.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of SearchResult objects ordered by relevance.
|
|
||||||
"""
|
|
||||||
if limit is None:
|
|
||||||
limit = self._config.search.limit
|
|
||||||
|
|
||||||
reranker = get_reranker(config=self._config)
|
|
||||||
|
|
||||||
if reranker is None:
|
|
||||||
chunk_results = await self.chunk_repository.search(
|
|
||||||
query, limit, search_type, filter
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
search_limit = limit * 10
|
|
||||||
raw_results = await self.chunk_repository.search(
|
|
||||||
query, search_limit, search_type, filter
|
|
||||||
)
|
|
||||||
chunks = [chunk for chunk, _ in raw_results]
|
|
||||||
chunk_results = await reranker.rerank(query, chunks, top_n=limit)
|
|
||||||
|
|
||||||
return [SearchResult.from_chunk(chunk, score) for chunk, score in chunk_results]
|
|
||||||
|
|
||||||
async def expand_context(
|
async def expand_context(
|
||||||
self,
|
self,
|
||||||
search_results: list[SearchResult],
|
search_results: list[SearchResult],
|
||||||
) -> list[SearchResult]:
|
) -> list[SearchResult]:
|
||||||
"""Expand search results with surrounding content from the document.
|
from haiku.rag.client.search import expand_context
|
||||||
|
|
||||||
Uses the document_items table for section-bounded expansion.
|
return await expand_context(self, search_results)
|
||||||
See haiku.rag.context for the algorithm description.
|
|
||||||
|
|
||||||
Results without doc_item_refs pass through unexpanded. This happens
|
|
||||||
when chunks were created without docling metadata (e.g., custom chunks
|
|
||||||
passed to import_document).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
search_results: List of SearchResult objects from search.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of SearchResult objects with expanded content.
|
|
||||||
"""
|
|
||||||
from haiku.rag.context import expand_with_items
|
|
||||||
|
|
||||||
max_chars = self._config.search.max_context_chars
|
|
||||||
|
|
||||||
# Group by document_id for efficient processing
|
|
||||||
document_groups: dict[str | None, list[SearchResult]] = {}
|
|
||||||
for result in search_results:
|
|
||||||
doc_id = result.document_id
|
|
||||||
if doc_id not in document_groups:
|
|
||||||
document_groups[doc_id] = []
|
|
||||||
document_groups[doc_id].append(result)
|
|
||||||
|
|
||||||
expanded_results = []
|
|
||||||
|
|
||||||
for doc_id, doc_results in document_groups.items():
|
|
||||||
if doc_id is None:
|
|
||||||
expanded_results.extend(doc_results)
|
|
||||||
continue
|
|
||||||
|
|
||||||
has_refs = any(r.doc_item_refs for r in doc_results)
|
|
||||||
if not has_refs:
|
|
||||||
expanded_results.extend(doc_results)
|
|
||||||
continue
|
|
||||||
|
|
||||||
expanded = await expand_with_items(
|
|
||||||
self.document_item_repository,
|
|
||||||
doc_id,
|
|
||||||
doc_results,
|
|
||||||
max_chars,
|
|
||||||
)
|
|
||||||
expanded_results.extend(expanded)
|
|
||||||
|
|
||||||
expanded_results.sort(key=lambda r: r.score, reverse=True)
|
|
||||||
return expanded_results
|
|
||||||
|
|
||||||
async def ask(
|
async def ask(
|
||||||
self,
|
self,
|
||||||
|
|
@ -1076,113 +1004,9 @@ class HaikuRAG:
|
||||||
)
|
)
|
||||||
|
|
||||||
async def visualize_chunk(self, chunk: Chunk) -> list:
|
async def visualize_chunk(self, chunk: Chunk) -> list:
|
||||||
"""Render page images with bounding box highlights for a chunk.
|
from haiku.rag.client.search import visualize_chunk
|
||||||
|
|
||||||
Expands the chunk's context to find the full section, then resolves
|
return await visualize_chunk(self, chunk)
|
||||||
bounding boxes from all items in the expanded range. This ensures
|
|
||||||
visualization covers all pages the expanded content spans.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
chunk: The chunk to visualize.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of PIL Image objects, one per page with bounding boxes.
|
|
||||||
Empty list if no bounding boxes or page images available.
|
|
||||||
"""
|
|
||||||
from copy import deepcopy
|
|
||||||
|
|
||||||
from PIL import ImageDraw
|
|
||||||
|
|
||||||
from haiku.rag.store.models.chunk import ChunkMetadata
|
|
||||||
|
|
||||||
# Get the document structure (from cache if available)
|
|
||||||
if not chunk.document_id:
|
|
||||||
return []
|
|
||||||
|
|
||||||
doc = await self.document_repository.get_docling_data(chunk.document_id)
|
|
||||||
if not doc:
|
|
||||||
return []
|
|
||||||
|
|
||||||
docling_doc = doc.get_docling_document()
|
|
||||||
if not docling_doc:
|
|
||||||
return []
|
|
||||||
|
|
||||||
# Expand context to get all doc_item_refs in the section
|
|
||||||
chunk_meta = chunk.get_chunk_metadata()
|
|
||||||
if chunk_meta.doc_item_refs:
|
|
||||||
search_result = SearchResult(
|
|
||||||
content=chunk.content,
|
|
||||||
score=1.0,
|
|
||||||
chunk_id=chunk.id,
|
|
||||||
document_id=chunk.document_id,
|
|
||||||
doc_item_refs=chunk_meta.doc_item_refs,
|
|
||||||
page_numbers=chunk_meta.page_numbers,
|
|
||||||
)
|
|
||||||
expanded = await self.expand_context([search_result])
|
|
||||||
refs = expanded[0].doc_item_refs if expanded else chunk_meta.doc_item_refs
|
|
||||||
meta = ChunkMetadata(doc_item_refs=refs)
|
|
||||||
else:
|
|
||||||
meta = chunk_meta
|
|
||||||
bounding_boxes = meta.resolve_bounding_boxes(docling_doc)
|
|
||||||
if not bounding_boxes:
|
|
||||||
return []
|
|
||||||
|
|
||||||
# Group bounding boxes by page
|
|
||||||
boxes_by_page: dict[int, list] = {}
|
|
||||||
for bbox in bounding_boxes:
|
|
||||||
if bbox.page_no not in boxes_by_page:
|
|
||||||
boxes_by_page[bbox.page_no] = []
|
|
||||||
boxes_by_page[bbox.page_no].append(bbox)
|
|
||||||
|
|
||||||
# Load only the needed page images
|
|
||||||
pages_doc = await self.document_repository.get_pages_data(chunk.document_id)
|
|
||||||
if not pages_doc:
|
|
||||||
return []
|
|
||||||
page_images = pages_doc.get_page_images(list(boxes_by_page.keys()))
|
|
||||||
|
|
||||||
# Render each page with its bounding boxes
|
|
||||||
images = []
|
|
||||||
for page_no in sorted(boxes_by_page.keys()):
|
|
||||||
if page_no not in page_images:
|
|
||||||
continue
|
|
||||||
|
|
||||||
page = page_images[page_no]
|
|
||||||
if page.image is None or page.image.pil_image is None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
pil_image = page.image.pil_image
|
|
||||||
page_height = page.size.height
|
|
||||||
|
|
||||||
# Calculate scale factor (image pixels vs document coordinates)
|
|
||||||
scale_x = pil_image.width / page.size.width
|
|
||||||
scale_y = pil_image.height / page.size.height
|
|
||||||
|
|
||||||
# Draw bounding boxes
|
|
||||||
image = deepcopy(pil_image)
|
|
||||||
draw = ImageDraw.Draw(image, "RGBA")
|
|
||||||
|
|
||||||
for bbox in boxes_by_page[page_no]:
|
|
||||||
# Convert from document coordinates to image coordinates
|
|
||||||
# Document coords are bottom-left origin, PIL uses top-left
|
|
||||||
x0 = bbox.left * scale_x
|
|
||||||
y0 = (page_height - bbox.top) * scale_y
|
|
||||||
x1 = bbox.right * scale_x
|
|
||||||
y1 = (page_height - bbox.bottom) * scale_y
|
|
||||||
|
|
||||||
# Ensure proper ordering (y0 should be less than y1 for PIL)
|
|
||||||
if y0 > y1:
|
|
||||||
y0, y1 = y1, y0
|
|
||||||
|
|
||||||
# Draw filled rectangle with transparency
|
|
||||||
fill_color = (255, 255, 0, 40) # Yellow with transparency
|
|
||||||
outline_color = (255, 165, 0, 100) # Orange outline
|
|
||||||
|
|
||||||
draw.rectangle([(x0, y0), (x1, y1)], fill=fill_color, outline=None)
|
|
||||||
draw.rectangle([(x0, y0), (x1, y1)], outline=outline_color, width=1)
|
|
||||||
|
|
||||||
images.append(image)
|
|
||||||
|
|
||||||
return images
|
|
||||||
|
|
||||||
async def rebuild_database(
|
async def rebuild_database(
|
||||||
self, mode: RebuildMode = RebuildMode.FULL
|
self, mode: RebuildMode = RebuildMode.FULL
|
||||||
|
|
|
||||||
196
haiku_rag_slim/haiku/rag/client/search.py
Normal file
196
haiku_rag_slim/haiku/rag/client/search.py
Normal file
|
|
@ -0,0 +1,196 @@
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from haiku.rag.reranking import get_reranker
|
||||||
|
from haiku.rag.store.models.chunk import Chunk, SearchResult
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
|
|
||||||
|
async def search(
|
||||||
|
client: "HaikuRAG",
|
||||||
|
query: str,
|
||||||
|
limit: int | None = None,
|
||||||
|
search_type: str = "hybrid",
|
||||||
|
filter: str | None = None,
|
||||||
|
) -> list[SearchResult]:
|
||||||
|
"""Search for relevant chunks with optional reranking.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
client: The HaikuRAG client (provides config + chunk repository).
|
||||||
|
query: The search query string.
|
||||||
|
limit: Maximum number of results to return. Defaults to config.search.limit.
|
||||||
|
search_type: Type of search - "vector", "fts", or "hybrid" (default).
|
||||||
|
filter: Optional SQL WHERE clause to filter documents before searching chunks.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of SearchResult objects ordered by relevance.
|
||||||
|
"""
|
||||||
|
if limit is None:
|
||||||
|
limit = client._config.search.limit
|
||||||
|
|
||||||
|
reranker = get_reranker(config=client._config)
|
||||||
|
|
||||||
|
if reranker is None:
|
||||||
|
chunk_results = await client.chunk_repository.search(
|
||||||
|
query, limit, search_type, filter
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
search_limit = limit * 10
|
||||||
|
raw_results = await client.chunk_repository.search(
|
||||||
|
query, search_limit, search_type, filter
|
||||||
|
)
|
||||||
|
chunks = [chunk for chunk, _ in raw_results]
|
||||||
|
chunk_results = await reranker.rerank(query, chunks, top_n=limit)
|
||||||
|
|
||||||
|
return [SearchResult.from_chunk(chunk, score) for chunk, score in chunk_results]
|
||||||
|
|
||||||
|
|
||||||
|
async def expand_context(
|
||||||
|
client: "HaikuRAG",
|
||||||
|
search_results: list[SearchResult],
|
||||||
|
) -> list[SearchResult]:
|
||||||
|
"""Expand search results with surrounding content from the document.
|
||||||
|
|
||||||
|
Uses the document_items table for section-bounded expansion.
|
||||||
|
See haiku.rag.context for the algorithm description.
|
||||||
|
|
||||||
|
Results without doc_item_refs pass through unexpanded. This happens when
|
||||||
|
chunks were created without docling metadata (e.g., custom chunks passed
|
||||||
|
to import_document).
|
||||||
|
"""
|
||||||
|
from haiku.rag.context import expand_with_items
|
||||||
|
|
||||||
|
max_chars = client._config.search.max_context_chars
|
||||||
|
|
||||||
|
# Group by document_id for efficient processing
|
||||||
|
document_groups: dict[str | None, list[SearchResult]] = {}
|
||||||
|
for result in search_results:
|
||||||
|
doc_id = result.document_id
|
||||||
|
if doc_id not in document_groups:
|
||||||
|
document_groups[doc_id] = []
|
||||||
|
document_groups[doc_id].append(result)
|
||||||
|
|
||||||
|
expanded_results = []
|
||||||
|
|
||||||
|
for doc_id, doc_results in document_groups.items():
|
||||||
|
if doc_id is None:
|
||||||
|
expanded_results.extend(doc_results)
|
||||||
|
continue
|
||||||
|
|
||||||
|
has_refs = any(r.doc_item_refs for r in doc_results)
|
||||||
|
if not has_refs:
|
||||||
|
expanded_results.extend(doc_results)
|
||||||
|
continue
|
||||||
|
|
||||||
|
expanded = await expand_with_items(
|
||||||
|
client.document_item_repository,
|
||||||
|
doc_id,
|
||||||
|
doc_results,
|
||||||
|
max_chars,
|
||||||
|
)
|
||||||
|
expanded_results.extend(expanded)
|
||||||
|
|
||||||
|
expanded_results.sort(key=lambda r: r.score, reverse=True)
|
||||||
|
return expanded_results
|
||||||
|
|
||||||
|
|
||||||
|
async def visualize_chunk(client: "HaikuRAG", chunk: Chunk) -> list:
|
||||||
|
"""Render page images with bounding box highlights for a chunk.
|
||||||
|
|
||||||
|
Expands the chunk's context to find the full section, then resolves
|
||||||
|
bounding boxes from all items in the expanded range. This ensures
|
||||||
|
visualization covers all pages the expanded content spans.
|
||||||
|
|
||||||
|
Returns a list of PIL Image objects, one per page with bounding boxes.
|
||||||
|
Empty list if no bounding boxes or page images available.
|
||||||
|
"""
|
||||||
|
from copy import deepcopy
|
||||||
|
|
||||||
|
from PIL import ImageDraw
|
||||||
|
|
||||||
|
from haiku.rag.store.models.chunk import ChunkMetadata
|
||||||
|
|
||||||
|
if not chunk.document_id:
|
||||||
|
return []
|
||||||
|
|
||||||
|
doc = await client.document_repository.get_docling_data(chunk.document_id)
|
||||||
|
if not doc:
|
||||||
|
return []
|
||||||
|
|
||||||
|
docling_doc = doc.get_docling_document()
|
||||||
|
if not docling_doc:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Expand context to get all doc_item_refs in the section
|
||||||
|
chunk_meta = chunk.get_chunk_metadata()
|
||||||
|
if chunk_meta.doc_item_refs:
|
||||||
|
search_result = SearchResult(
|
||||||
|
content=chunk.content,
|
||||||
|
score=1.0,
|
||||||
|
chunk_id=chunk.id,
|
||||||
|
document_id=chunk.document_id,
|
||||||
|
doc_item_refs=chunk_meta.doc_item_refs,
|
||||||
|
page_numbers=chunk_meta.page_numbers,
|
||||||
|
)
|
||||||
|
expanded = await expand_context(client, [search_result])
|
||||||
|
refs = expanded[0].doc_item_refs if expanded else chunk_meta.doc_item_refs
|
||||||
|
meta = ChunkMetadata(doc_item_refs=refs)
|
||||||
|
else:
|
||||||
|
meta = chunk_meta
|
||||||
|
bounding_boxes = meta.resolve_bounding_boxes(docling_doc)
|
||||||
|
if not bounding_boxes:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Group bounding boxes by page
|
||||||
|
boxes_by_page: dict[int, list] = {}
|
||||||
|
for bbox in bounding_boxes:
|
||||||
|
if bbox.page_no not in boxes_by_page:
|
||||||
|
boxes_by_page[bbox.page_no] = []
|
||||||
|
boxes_by_page[bbox.page_no].append(bbox)
|
||||||
|
|
||||||
|
# Load only the needed page images
|
||||||
|
pages_doc = await client.document_repository.get_pages_data(chunk.document_id)
|
||||||
|
if not pages_doc:
|
||||||
|
return []
|
||||||
|
page_images = pages_doc.get_page_images(list(boxes_by_page.keys()))
|
||||||
|
|
||||||
|
# Render each page with its bounding boxes
|
||||||
|
images = []
|
||||||
|
for page_no in sorted(boxes_by_page.keys()):
|
||||||
|
if page_no not in page_images:
|
||||||
|
continue
|
||||||
|
|
||||||
|
page = page_images[page_no]
|
||||||
|
if page.image is None or page.image.pil_image is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
pil_image = page.image.pil_image
|
||||||
|
page_height = page.size.height
|
||||||
|
|
||||||
|
# Scale factor: image pixels vs document coordinates
|
||||||
|
scale_x = pil_image.width / page.size.width
|
||||||
|
scale_y = pil_image.height / page.size.height
|
||||||
|
|
||||||
|
image = deepcopy(pil_image)
|
||||||
|
draw = ImageDraw.Draw(image, "RGBA")
|
||||||
|
|
||||||
|
for bbox in boxes_by_page[page_no]:
|
||||||
|
# Document coords are bottom-left origin; PIL uses top-left
|
||||||
|
x0 = bbox.left * scale_x
|
||||||
|
y0 = (page_height - bbox.top) * scale_y
|
||||||
|
x1 = bbox.right * scale_x
|
||||||
|
y1 = (page_height - bbox.bottom) * scale_y
|
||||||
|
|
||||||
|
if y0 > y1:
|
||||||
|
y0, y1 = y1, y0
|
||||||
|
|
||||||
|
fill_color = (255, 255, 0, 40) # Yellow with transparency
|
||||||
|
outline_color = (255, 165, 0, 100) # Orange outline
|
||||||
|
|
||||||
|
draw.rectangle([(x0, y0), (x1, y1)], fill=fill_color, outline=None)
|
||||||
|
draw.rectangle([(x0, y0), (x1, y1)], outline=outline_color, width=1)
|
||||||
|
|
||||||
|
images.append(image)
|
||||||
|
|
||||||
|
return images
|
||||||
Loading…
Reference in a new issue