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 typing import TYPE_CHECKING, overload from urllib.parse import urlparse import httpx 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.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__) 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. 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() @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 before tearing down the connection. Each create_document / update_document can schedule its own vacuum task; all must be awaited, not just the most recently scheduled one. Vacuum skips when another is already running, so the cleanup for the final writes may have been a no-op. Run one more pass once the in-flight tasks are done to collapse versions created after the last vacuum took the lock. """ if not self._vacuum_tasks: return await asyncio.gather(*self._vacuum_tasks, return_exceptions=True) # __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 and track the task for later awaiting.""" 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 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, ) -> 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, ) 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``.""" from haiku.rag.client.documents import parent_uri_filter doc = await self.get_document_by_id(document_id) if doc is None: return False if doc.uri: children = await self.list_documents(filter=parent_uri_filter(doc.uri)) for child in children: if child.id and child.id != document_id: await self.delete_document(child.id) return await self.document_repository.delete(document_id) 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()