haiku.rag/haiku_rag_slim/haiku/rag/client/__init__.py
2026-06-16 16:37:26 +03:00

514 lines
18 KiB
Python

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()