covers_multiple, source_names, source and reader_for replace the private state seven modules were reading to work out how many databases they had. The configured selection is kept intact, so entering a client twice derives the same database rather than the last derivation.
333 lines
13 KiB
Python
333 lines
13 KiB
Python
import asyncio
|
|
import logging
|
|
from pathlib import Path
|
|
from time import monotonic
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
from haiku.rag.client.scope import DatabaseRef, DatabaseScope
|
|
from haiku.rag.config import AppConfig
|
|
from haiku.rag.store.engine import Store
|
|
from haiku.rag.store.exceptions import (
|
|
ConfigMismatchError,
|
|
MigrationRequiredError,
|
|
ReadOnlyError,
|
|
SourceUnavailableError,
|
|
)
|
|
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
|
|
|
|
if TYPE_CHECKING:
|
|
from haiku.rag.store.models.document import Document
|
|
|
|
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
|
|
|
|
|
|
# Failures whose message names the remedy and never the location, so the failing
|
|
# database is named alongside it instead of in place of it.
|
|
_NAMEABLE_FAILURES = (MigrationRequiredError, ConfigMismatchError, ReadOnlyError)
|
|
|
|
|
|
async def aclose_quietly(closeable: Any, what: str) -> None:
|
|
"""Close, reporting failure to the log rather than raising.
|
|
|
|
Teardown can run while an exception unwinds, so a raising close must
|
|
neither mask that exception nor stop a sibling from being closed.
|
|
"""
|
|
try:
|
|
await closeable.aclose()
|
|
except Exception:
|
|
logger.debug("Closing the %s failed on teardown", what, exc_info=True)
|
|
|
|
|
|
def default_db_path(config: AppConfig) -> Path:
|
|
"""Where a database lives when its location names no path."""
|
|
return config.storage.data_dir / "haiku.rag.lancedb"
|
|
|
|
|
|
class SingleDatabaseSession:
|
|
"""One database: its store, its repositories, and their lifecycle.
|
|
|
|
Everything that needs a store lives here, so nothing above has to ask whether
|
|
it has one. ``source`` is the configured name this database answers to, or
|
|
None where nothing names it.
|
|
|
|
``db_path``, ``config``, ``read_only`` and ``source`` are readable because a
|
|
client facade is built over a session it does not own, and needs to report
|
|
the same things it would have reported had it opened the database itself.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
db_path: Path | str,
|
|
config: AppConfig,
|
|
*,
|
|
skip_validation: bool = False,
|
|
create: bool = False,
|
|
read_only: bool = False,
|
|
source: str | None = None,
|
|
) -> None:
|
|
self.db_path = db_path
|
|
self.config = config
|
|
self.read_only = read_only
|
|
self.source = source
|
|
self._skip_validation = skip_validation
|
|
self._create = create
|
|
self._vacuum_tasks: set[asyncio.Task] = set()
|
|
self._last_vacuum_at: float | None = None
|
|
self._vacuum_dirty = False
|
|
|
|
async def open(self) -> "SingleDatabaseSession":
|
|
"""Connect, validate, and build the repositories."""
|
|
failure: str | None = None
|
|
try:
|
|
self.store = Store(
|
|
self.db_path,
|
|
config=self.config,
|
|
skip_validation=self._skip_validation,
|
|
create=self._create,
|
|
read_only=self.read_only,
|
|
)
|
|
# If _initialize fails mid-way (e.g. migration check raises after
|
|
# connect), close the store so we don't leak the LanceDB connection —
|
|
# the caller's `async with` never entered, so its exit won't run.
|
|
try:
|
|
await self.store._initialize()
|
|
except BaseException:
|
|
self.store.close()
|
|
raise
|
|
except _NAMEABLE_FAILURES as error:
|
|
# These say what to run and never where the database is, so the name
|
|
# is added to the message rather than replacing it: the operator needs
|
|
# both which database failed and what to do about it.
|
|
if self.source is None:
|
|
raise
|
|
raise type(error)(f"database {self.source!r}: {error}") from error
|
|
except Exception as error:
|
|
# A legacy `uri` or `db_path` session has no name to report instead,
|
|
# so its error passes through as it always has.
|
|
if self.source is None:
|
|
raise
|
|
failure = type(error).__name__
|
|
if failure is not None:
|
|
# Raised outside the except block on purpose. A database named in
|
|
# config is reported by name, and the original spells out the path or
|
|
# the bucket: `from None` would only stop it being *printed*, leaving
|
|
# it on `__context__` for anything that walks the chain.
|
|
raise SourceUnavailableError(
|
|
f"database {self.source!r} could not be opened: {failure}"
|
|
)
|
|
self.document_repository = DocumentRepository(self.store)
|
|
self.chunk_repository = ChunkRepository(self.store)
|
|
self.document_item_repository = DocumentItemRepository(self.store)
|
|
return self
|
|
|
|
async def drain_vacuum(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
|
|
# Teardown 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)
|
|
|
|
def name(self, document: "Document | None") -> "Document | None":
|
|
"""`document`, told which database it came from."""
|
|
if document is not None:
|
|
document.source = self.source
|
|
return document
|
|
|
|
async def get_document_by_id(self, document_id: str) -> "Document | None":
|
|
return self.name(await self.document_repository.get_by_id(document_id))
|
|
|
|
async def get_document_by_uri(self, uri: str) -> "Document | None":
|
|
return self.name(await self.document_repository.get_by_uri(uri))
|
|
|
|
async def list_documents(
|
|
self,
|
|
limit: int | None = None,
|
|
offset: int | None = None,
|
|
filter: str | None = None,
|
|
include_content: bool = False,
|
|
) -> "list[Document]":
|
|
documents = await self.document_repository.list_all(
|
|
limit=limit, offset=offset, filter=filter, include_content=include_content
|
|
)
|
|
for document in documents:
|
|
document.source = self.source
|
|
return documents
|
|
|
|
async def delete_document(self, document_id: str) -> bool:
|
|
"""Delete a document, cascading 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_transaction():
|
|
# 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
|
|
|
|
for doc_id in ids_to_delete:
|
|
await self.document_repository.delete(doc_id)
|
|
|
|
if self.config.storage.auto_vacuum:
|
|
self.schedule_vacuum()
|
|
return True
|
|
|
|
async def aclose(self) -> None:
|
|
"""Drain, release the embedder, and close the connection.
|
|
|
|
The store owns the embedder, so releasing it belongs here rather than
|
|
with whoever happened to hold the session.
|
|
"""
|
|
await self.drain_vacuum()
|
|
await aclose_quietly(self.store.embedder, "embedder")
|
|
self.close()
|
|
|
|
def close(self) -> None:
|
|
"""Close the underlying store connection."""
|
|
self.store.close()
|
|
|
|
|
|
class FederatedSession:
|
|
"""Several databases, read as one.
|
|
|
|
Composes single-database sessions and owns their teardown. They open on first
|
|
use rather than at entry: which databases a query covers is a per-query
|
|
choice, so a database nobody asked for must neither be opened for nothing nor
|
|
be able to fail a query.
|
|
|
|
Reads only. Writing names a database, and naming one is what
|
|
``SingleDatabaseSession`` is.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
scope: DatabaseScope,
|
|
config: AppConfig,
|
|
*,
|
|
skip_validation: bool = False,
|
|
read_only: bool = False,
|
|
) -> None:
|
|
self._refs: dict[str, DatabaseRef] = {
|
|
ref.name: ref for ref in scope.databases if ref.name is not None
|
|
}
|
|
self._config = config
|
|
self._skip_validation = skip_validation
|
|
self._read_only = read_only
|
|
self._sessions: dict[str, SingleDatabaseSession] = {}
|
|
self._lock = asyncio.Lock()
|
|
|
|
@property
|
|
def names(self) -> tuple[str, ...]:
|
|
"""The databases covered, in configured order."""
|
|
return tuple(self._refs)
|
|
|
|
async def sessions_for(self, names: list[str]) -> list[SingleDatabaseSession]:
|
|
"""The sessions for these databases, opening any not yet open.
|
|
|
|
Missing ones open together: on object storage a serial loop makes the
|
|
first query cost the sum of the opens.
|
|
"""
|
|
unknown = [name for name in names if name not in self._refs]
|
|
if unknown:
|
|
raise KeyError(
|
|
f"unknown database(s) {', '.join(sorted(unknown))}; configured: "
|
|
f"{', '.join(sorted(self._refs))}"
|
|
)
|
|
async with self._lock:
|
|
missing = [name for name in names if name not in self._sessions]
|
|
if missing:
|
|
opened = await asyncio.gather(
|
|
*(self._open(self._refs[name]) for name in missing),
|
|
return_exceptions=True,
|
|
)
|
|
# Whatever opened is tracked before the failure is reported, so
|
|
# teardown closes it: `gather` does not cancel the siblings of the
|
|
# one that raised, and an untracked connection leaks.
|
|
failure: BaseException | None = None
|
|
for name, result in zip(missing, opened, strict=True):
|
|
if isinstance(result, BaseException):
|
|
failure = failure or result
|
|
else:
|
|
self._sessions[name] = result
|
|
if failure is not None:
|
|
raise failure
|
|
return [self._sessions[name] for name in names]
|
|
|
|
async def _open(self, ref: DatabaseRef) -> SingleDatabaseSession:
|
|
one, db_path = ref.connection(self._config)
|
|
return await SingleDatabaseSession(
|
|
db_path if db_path is not None else default_db_path(one),
|
|
one,
|
|
skip_validation=self._skip_validation,
|
|
read_only=self._read_only,
|
|
source=ref.name,
|
|
).open()
|
|
|
|
async def aclose(self) -> None:
|
|
"""Close every database this session opened."""
|
|
for session in self._sessions.values():
|
|
await aclose_quietly(session, "database")
|
|
self._sessions.clear()
|