import asyncio import hashlib import json import logging import mimetypes import tempfile from collections.abc import AsyncGenerator from datetime import datetime from enum import Enum from functools import cached_property from pathlib import Path from time import monotonic from typing import TYPE_CHECKING, overload from urllib.parse import urlparse import httpx from haiku.rag.client.documents import DocumentImport from haiku.rag.config import AppConfig, Config from haiku.rag.converters import get_converter from haiku.rag.reranking import get_reranker from haiku.rag.store.engine import Store from haiku.rag.store.models.chunk import Chunk, SearchResult, SearchType from haiku.rag.store.models.document import Document from haiku.rag.store.models.document_item import extract_items from haiku.rag.store.repositories.chunk import ChunkRepository from haiku.rag.store.repositories.document import DocumentRepository from haiku.rag.store.repositories.document_item import DocumentItemRepository from haiku.rag.store.repositories.settings import SettingsRepository from haiku.rag.utils import escape_sql_string if TYPE_CHECKING: from docling_core.types.doc.document import DoclingDocument from PIL import Image as PILImage from haiku.rag.embeddings import EmbedderWrapper from haiku.rag.ingester.metadata import MetadataProvider from haiku.rag.ingester.sources.base import Source from haiku.rag.reranking.base import RerankerBase from haiku.rag.sandbox import AnalysisResult from haiku.rag.store.models.citation import Citation logger = logging.getLogger(__name__) # Throttle for the background auto-vacuum: under sustained ingestion, scheduling # a compaction on every write degenerates into back-to-back optimize() passes # that churn the blob-bearing documents table. Fire at most one per interval; a # final vacuum on close collapses anything throttled here. _VACUUM_MIN_INTERVAL_S = 300.0 class RebuildMode(Enum): """Mode for rebuilding the database.""" FULL = "full" # Re-convert from source, re-chunk, re-embed RECHUNK = "rechunk" # Re-chunk from existing content, re-embed EMBED_ONLY = "embed_only" # Keep chunks, only regenerate embeddings TITLE_ONLY = "title_only" # Only generate titles for untitled documents DESCRIPTIONS = "descriptions" # Run the VLM over already-stored picture # bytes, patch descriptions into the docling blob, then re-chunk + re-embed. SET_EMBEDDER = "set_embedder" # Adopt the current embedder identity without # re-embedding, when the vector dimension is unchanged. class HaikuRAG: """High-level haiku-rag client.""" def __init__( self, db_path: Path | None = None, config: AppConfig = Config, skip_validation: bool = False, create: bool = False, read_only: bool = False, before: datetime | None = None, ): """Initialize the RAG client with a database path. Args: db_path: Path to the database file. If None, uses config.storage.data_dir. config: Configuration to use. Defaults to global Config. skip_validation: Whether to skip configuration validation on database load. create: Whether to create the database if it doesn't exist. read_only: Whether to open the database in read-only mode. before: Query the database as it existed at this datetime. Implies read_only=True. """ self._config = config if db_path is None: db_path = self._config.storage.data_dir / "haiku.rag.lancedb" self._db_path = db_path self._skip_validation = skip_validation self._create = create self._read_only = read_only self._before = before self._vacuum_tasks: set[asyncio.Task] = set() self._last_vacuum_at: float | None = None self._vacuum_dirty = False @property def is_read_only(self) -> bool: """Whether the client is in read-only mode.""" return self.store.is_read_only @property def embedder(self) -> "EmbedderWrapper": """The embedder owned by the Store, reused across all operations.""" return self.store.embedder @cached_property def reranker(self) -> "RerankerBase | None": """The configured reranker, built once and reused across searches. None when reranking is disabled. Local rerankers load model weights on construction, so building per search would reload them on every query. """ return get_reranker(config=self._config) async def __aenter__(self): """Async context manager entry — initializes store and repositories.""" self.store = Store( self._db_path, config=self._config, skip_validation=self._skip_validation, create=self._create, read_only=self._read_only, before=self._before, ) # If _initialize fails mid-way (e.g. migration check raises after # connect), close the store so we don't leak the LanceDB connection — # __aexit__ won't run because the `async with` never entered. try: await self.store._initialize() except BaseException: self.store.close() raise self.document_repository = DocumentRepository(self.store) self.chunk_repository = ChunkRepository(self.store) self.document_item_repository = DocumentItemRepository(self.store) return self async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002 """Async context manager exit.""" await self._await_vacuum_tasks() self.close() return False async def _await_vacuum_tasks(self) -> None: """Drain background vacuum work and run a final collapse before teardown. Writes schedule a throttled background vacuum; many are debounced or skip because another vacuum holds the lock. The final pass collapses the versions those left behind. It runs whenever writes happened (``_vacuum_dirty``) — not gated on in-flight tasks remaining, since a debounced run may have scheduled none — but never when nothing was written (so opening + closing a store still never writes). """ if self._vacuum_tasks: await asyncio.gather(*self._vacuum_tasks, return_exceptions=True) if not self._vacuum_dirty: return self._vacuum_dirty = False # __aexit__ runs during exception unwinding; a raising vacuum here would # mask the original exception, so the drain stays best-effort. try: await self.store.vacuum() except Exception: logger.debug("Final vacuum on close failed", exc_info=True) def _schedule_vacuum(self) -> None: """Schedule a background vacuum, throttled to at most one per ``_VACUUM_MIN_INTERVAL_S``. Sustained writes would otherwise trigger back-to-back compaction of the blob-bearing documents table. The throttle only skips the background task — ``_vacuum_dirty`` still marks that a final vacuum on close is owed.""" self._vacuum_dirty = True now = monotonic() if ( self._last_vacuum_at is not None and now - self._last_vacuum_at < _VACUUM_MIN_INTERVAL_S ): return self._last_vacuum_at = now task = asyncio.create_task(self.store.vacuum()) self._vacuum_tasks.add(task) task.add_done_callback(self._vacuum_tasks.discard) # ========================================================================= # Processing Primitives # ========================================================================= @overload async def convert( self, source: Path, *, source_uri: str | None = None ) -> "DoclingDocument": ... @overload async def convert( self, source: str, *, format: str = "md", source_uri: str | None = None ) -> "DoclingDocument": ... async def convert( self, source: Path | str, *, format: str = "md", source_uri: str | None = None, ) -> "DoclingDocument": from haiku.rag.client.processing import convert return await convert(self._config, source, format=format, source_uri=source_uri) async def chunk( self, docling_document: "DoclingDocument", *, existing_picture_data: dict[str, bytes] | None = None, document_id: str | None = None, ) -> list[Chunk]: from haiku.rag.client.processing import chunk return await chunk( self._config, docling_document, embedder=self.embedder, existing_picture_data=existing_picture_data, document_id=document_id, ) # ========================================================================= # Title Generation # ========================================================================= async def generate_title(self, document: Document) -> str | None: from haiku.rag.client.titles import generate_title return await generate_title(self._config, document) async def create_document( self, content: str, uri: str | None = None, title: str | None = None, metadata: dict | None = None, format: str = "md", ) -> Document: from haiku.rag.client.documents import create_document return await create_document(self, content, uri, title, metadata, format) async def import_document( self, docling_document: "DoclingDocument", chunks: list[Chunk], uri: str | None = None, title: str | None = None, metadata: dict | None = None, ) -> Document: from haiku.rag.client.documents import import_document return await import_document( self, docling_document, chunks, uri, title, metadata ) async def import_documents( self, imports: "list[DocumentImport]", ) -> list[Document]: from haiku.rag.client.documents import import_documents return await import_documents(self, imports) async def create_document_from_source( self, source: str | Path, title: str | None = None, metadata: dict | None = None, uri: str | None = None, storage_options: dict[str, str] | None = None, sources: "list[Source] | None" = None, source_id: str | None = None, metadata_provider: "MetadataProvider | None" = None, ) -> Document | list[Document]: from haiku.rag.client.documents import create_document_from_source return await create_document_from_source( self, source, title, metadata, uri=uri, storage_options=storage_options, sources=sources, source_id=source_id, metadata_provider=metadata_provider, ) async def update_document( self, document_id: str, content: str | None = None, metadata: dict | None = None, chunks: list[Chunk] | None = None, title: str | None = None, docling_document: "DoclingDocument | None" = None, ) -> Document: from haiku.rag.client.documents import update_document return await update_document( self, document_id, content, metadata, chunks, title, docling_document, ) async def get_document_by_id(self, document_id: str) -> Document | None: """Get a document by its ID. Args: document_id: The unique identifier of the document. Returns: The Document instance if found, None otherwise. """ return await self.document_repository.get_by_id(document_id) async def get_chunk_by_id(self, chunk_id: str) -> Chunk | None: """Get a chunk by its ID. Args: chunk_id: The unique identifier of the chunk. Returns: The Chunk instance if found, None otherwise. """ return await self.chunk_repository.get_by_id(chunk_id) async def get_document_by_uri(self, uri: str) -> Document | None: """Get a document by its URI. Args: uri: The URI identifier of the document. Returns: The Document instance if found, None otherwise. """ return await self.document_repository.get_by_uri(uri) async def resolve_document(self, id_or_title: str) -> Document | None: """Resolve a document by ID, title, or URI (in that order). Args: id_or_title: Document ID, title, or URI to look up. Returns: The Document instance if found, None otherwise. """ doc = await self.get_document_by_id(id_or_title) if doc: return doc safe_input = escape_sql_string(id_or_title) docs = await self.list_documents(filter=f"title = '{safe_input}'") if docs and docs[0].id: return await self.get_document_by_id(docs[0].id) docs = await self.list_documents(filter=f"uri = '{safe_input}'") if docs and docs[0].id: return await self.get_document_by_id(docs[0].id) return None async def delete_document(self, document_id: str) -> bool: """Delete a document by its ID. Cascades to children linked via ``metadata.parent_uri``. The whole subtree (root + transitive children) is deleted under a single write lock and a single version snapshot, so the cascade is atomic: any failure restores every table to the pre-delete state, and no other write can interleave between deleting a child and its parent. """ from haiku.rag.client.documents import parent_uri_filter async with self.store._write_lock: # Resolve existence and collect the subtree under the lock so two # concurrent deletes of the same id can't both proceed, and children # can't appear or move between collection and deletion. parent_uri # links a child to its parent's uri; walk transitively, guarding # against cycles. ids_to_delete: list[str] = [] seen: set[str] = set() queue = [await self.get_document_by_id(document_id)] while queue: doc = queue.pop() if doc is None or doc.id is None or doc.id in seen: continue seen.add(doc.id) ids_to_delete.append(doc.id) if doc.uri: queue.extend( await self.list_documents(filter=parent_uri_filter(doc.uri)) ) if not ids_to_delete: return False versions = await self.store.current_table_versions() try: for doc_id in ids_to_delete: await self.document_repository.delete(doc_id) except Exception: await self.store.restore_table_versions(versions) raise if self._config.storage.auto_vacuum: self._schedule_vacuum() return True async def list_documents( self, limit: int | None = None, offset: int | None = None, filter: str | None = None, include_content: bool = False, ) -> list[Document]: """List all documents with optional pagination and filtering. Args: limit: Maximum number of documents to return. offset: Number of documents to skip. filter: Optional SQL WHERE clause to filter documents. include_content: Whether to load content and docling_document. Defaults to False to avoid loading large blobs. Returns: List of Document instances matching the criteria. """ return await self.document_repository.list_all( limit=limit, offset=offset, filter=filter, include_content=include_content ) async def count_documents(self, filter: str | None = None) -> int: """Count documents with optional filtering. Args: filter: Optional SQL WHERE clause to filter documents. Returns: Number of documents matching the criteria. """ return await self.document_repository.count(filter=filter) async def search( self, query: "str | bytes | PILImage.Image", limit: int | None = None, search_type: SearchType | None = None, filter: str | None = None, include_images: bool = True, ) -> list[SearchResult]: from haiku.rag.client.search import search return await search(self, query, limit, search_type, filter, include_images) async def expand_context( self, search_results: list[SearchResult], ) -> list[SearchResult]: from haiku.rag.client.search import expand_context return await expand_context(self, search_results) async def ask( self, question: str, filter: str | None = None, ) -> "tuple[str, list[Citation]]": from haiku.rag.client.agents import ask return await ask(self, question, filter) async def analyze( self, question: str, filter: str | None = None, ) -> "AnalysisResult": from haiku.rag.client.agents import analyze return await analyze(self, question, filter) async def visualize_chunk(self, chunk: Chunk) -> list: from haiku.rag.client.search import visualize_chunk return await visualize_chunk(self, chunk) async def rebuild_database( self, mode: RebuildMode = RebuildMode.FULL ) -> AsyncGenerator[str, None]: from haiku.rag.client.rebuild import rebuild_database async for doc_id in rebuild_database(self, mode): yield doc_id async def vacuum(self) -> None: """Optimize and clean up old versions across all tables.""" await self.store.vacuum() def close(self): """Close the underlying store connection.""" self.store.close()