From af7731f4e10b0cfdb172e54e8f16e04dbeca2765 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 20 Apr 2026 14:25:16 +0300 Subject: [PATCH] bulk-fetch items.jsonl via lazy cache to avoid per-document query timeout --- .../haiku/rag/agents/analysis/sandbox.py | 80 +++++++++++-------- .../rag/store/repositories/document_item.py | 26 ++++++ 2 files changed, 73 insertions(+), 33 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/agents/analysis/sandbox.py b/haiku_rag_slim/haiku/rag/agents/analysis/sandbox.py index 90976a49..88a61ea9 100644 --- a/haiku_rag_slim/haiku/rag/agents/analysis/sandbox.py +++ b/haiku_rag_slim/haiku/rag/agents/analysis/sandbox.py @@ -55,6 +55,7 @@ class Sandbox: _config: AppConfig _context: AnalysisContext _search_results: "list[SearchResult]" + _items_cache: dict[str, str] | None def __init__( self, @@ -66,6 +67,7 @@ class Sandbox: self._config = config self._context = context self._search_results = [] + self._items_cache = None def _build_external_functions(self) -> dict[str, Any]: """Build async external functions for the Monty interpreter.""" @@ -147,6 +149,51 @@ class Sandbox: async with HaikuRAG(db_path, config=config, read_only=True) as rag: docs = await rag.list_documents(filter=self._context.filter) + doc_ids = [doc.id for doc in docs if doc.id] + + def _load_items_cache() -> dict[str, str]: + """Bulk-fetch all document items in one query, serialize to JSONL.""" + + async def _fetch() -> dict[str, str]: + from haiku.rag.client import HaikuRAG + + async with HaikuRAG(db_path, config=config, read_only=True) as rag: + grouped = await rag.document_item_repository.get_all_items_grouped( + doc_ids + ) + result: dict[str, str] = {} + for did, items in grouped.items(): + lines = [] + for item in items: + lines.append( + json.dumps( + { + "position": item.position, + "self_ref": item.self_ref, + "label": item.label, + "text": item.text, + "page_numbers": item.page_numbers, + }, + ensure_ascii=False, + ) + ) + result[did] = "\n".join(lines) + return result + + return _run_async(_fetch()) + + sandbox = self + + def _make_items_reader( + did: str, + ) -> Callable[["PurePosixPath"], str]: + def read_items(_path: "PurePosixPath") -> str: + if sandbox._items_cache is None: + sandbox._items_cache = _load_items_cache() + return sandbox._items_cache.get(did, "") + + return read_items + for doc in docs: if not doc.id: continue @@ -181,39 +228,6 @@ class Sandbox: return read_content - def _make_items_reader( - did: str, - ) -> Callable[["PurePosixPath"], str]: - def read_items(_path: "PurePosixPath") -> str: - async def _fetch() -> str: - from haiku.rag.client import HaikuRAG - - async with HaikuRAG( - db_path, config=config, read_only=True - ) as rag: - items = await rag.document_item_repository.get_all_items( - did - ) - lines = [] - for item in items: - lines.append( - json.dumps( - { - "position": item.position, - "self_ref": item.self_ref, - "label": item.label, - "text": item.text, - "page_numbers": item.page_numbers, - }, - ensure_ascii=False, - ) - ) - return "\n".join(lines) - - return _run_async(_fetch()) - - return read_items - files.append( CallbackFile( f"{doc_dir}/content.txt", diff --git a/haiku_rag_slim/haiku/rag/store/repositories/document_item.py b/haiku_rag_slim/haiku/rag/store/repositories/document_item.py index eeb5523d..1e68c736 100644 --- a/haiku_rag_slim/haiku/rag/store/repositories/document_item.py +++ b/haiku_rag_slim/haiku/rag/store/repositories/document_item.py @@ -52,6 +52,32 @@ class DocumentItemRepository: items.sort(key=lambda x: x.position) return items + async def get_all_items_grouped( + self, document_ids: list[str] | None = None + ) -> dict[str, list[DocumentItem]]: + """Get all items grouped by document_id in a single query. + + Args: + document_ids: If provided, only fetch items for these documents. + If None, fetches all items. + + Returns: + Dict mapping document_id to sorted list of DocumentItem. + """ + query = self.store.document_items_table.search() + if document_ids is not None: + safe_ids = ", ".join(f"'{escape_sql_string(did)}'" for did in document_ids) + query = query.where(f"document_id IN ({safe_ids})") + rows = query.to_list() + + grouped: dict[str, list[DocumentItem]] = {} + for row in rows: + item = self._record_to_item(row) + grouped.setdefault(item.document_id, []).append(item) + for items in grouped.values(): + items.sort(key=lambda x: x.position) + return grouped + async def get_items_in_range( self, document_id: str, start: int, end: int ) -> list[DocumentItem]: