Give one database its own session

SingleDatabaseSession owns the store, the repositories and the vacuum
machinery, so nothing above has to ask whether it has a store. The client
keeps every name callers already use.
This commit is contained in:
Yiorgis Gozadinos 2026-08-24 17:13:45 +03:00
parent fcfa4aefd8
commit 79a63a46d4
No known key found for this signature in database
6 changed files with 231 additions and 113 deletions

View file

@ -16,6 +16,7 @@ from urllib.parse import urlparse
import httpx import httpx
from haiku.rag.client.documents import DocumentImport from haiku.rag.client.documents import DocumentImport
from haiku.rag.client.session import SingleDatabaseSession
from haiku.rag.config import AppConfig, get_config from haiku.rag.config import AppConfig, get_config
from haiku.rag.converters import get_converter from haiku.rag.converters import get_converter
from haiku.rag.embeddings import get_embedder from haiku.rag.embeddings import get_embedder
@ -50,17 +51,6 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) 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 first_found( async def first_found(
clients: "list[HaikuRAG]", clients: "list[HaikuRAG]",
@ -154,14 +144,41 @@ class HaikuRAG:
self._skip_validation = skip_validation self._skip_validation = skip_validation
self._create = create self._create = create
self._read_only = read_only self._read_only = read_only
self._vacuum_tasks: set[asyncio.Task] = set()
self._last_vacuum_at: float | None = None
self._vacuum_dirty = False
self._requested_sources = sources self._requested_sources = sources
self._clients: dict[str, HaikuRAG] = {} self._clients: dict[str, HaikuRAG] = {}
self._federated: dict[str, str] = {} self._federated: dict[str, str] = {}
self._clients_lock = asyncio.Lock() self._clients_lock = asyncio.Lock()
self._source: str | None = None self._source: str | None = None
self._session: SingleDatabaseSession | None = None
@property
def store(self) -> Store:
"""The store of the database this client opened.
Absent while covering a set: a store has no unambiguous meaning across
several, and `clients_for` reaches the one holding a given database.
"""
if self._session is None:
raise AttributeError("store")
return self._session.store
@property
def document_repository(self) -> DocumentRepository:
if self._session is None:
raise AttributeError("document_repository")
return self._session.document_repository
@property
def chunk_repository(self) -> ChunkRepository:
if self._session is None:
raise AttributeError("chunk_repository")
return self._session.chunk_repository
@property
def document_item_repository(self) -> DocumentItemRepository:
if self._session is None:
raise AttributeError("document_item_repository")
return self._session.document_item_repository
@property @property
def is_read_only(self) -> bool: def is_read_only(self) -> bool:
@ -247,47 +264,14 @@ class HaikuRAG:
if db_path is not None: if db_path is not None:
self._db_path = db_path self._db_path = db_path
failure: str | None = None self._session = await SingleDatabaseSession(
try: self._db_path,
self.store = Store( self._config,
self._db_path, skip_validation=self._skip_validation,
config=self._config, create=self._create,
skip_validation=self._skip_validation, read_only=self._read_only,
create=self._create, source=self._source,
read_only=self._read_only, ).open()
)
# 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
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` client 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 return self
async def clients_for(self, names: list[str]) -> list["HaikuRAG"]: async def clients_for(self, names: list[str]) -> list["HaikuRAG"]:
@ -410,44 +394,12 @@ class HaikuRAG:
await _aclose_quietly(cached, name) await _aclose_quietly(cached, name)
async def _await_vacuum_tasks(self) -> None: async def _await_vacuum_tasks(self) -> None:
"""Drain background vacuum work and run a final collapse before teardown. if self._session is not None:
await self._session.drain_vacuum()
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: def _schedule_vacuum(self) -> None:
"""Schedule a background vacuum, throttled to at most one per if self._session is not None:
``_VACUUM_MIN_INTERVAL_S``. Sustained writes would otherwise trigger self._session.schedule_vacuum()
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 # Processing Primitives
@ -935,4 +887,5 @@ class HaikuRAG:
def close(self): def close(self):
"""Close the underlying store connection.""" """Close the underlying store connection."""
self._require_one_database("close") self._require_one_database("close")
self.store.close() assert self._session is not None
self._session.close()

View file

@ -0,0 +1,147 @@
import asyncio
import logging
from pathlib import Path
from time import monotonic
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
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)
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.
"""
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._skip_validation = skip_validation
self._create = create
self._read_only = read_only
self.source = source
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 close(self) -> None:
"""Close the underlying store connection."""
self.store.close()

View file

@ -1196,10 +1196,10 @@ async def test_delete_marks_vacuum_dirty(temp_db_path):
uri="mem://del", uri="mem://del",
) )
assert doc.id is not None assert doc.id is not None
client._vacuum_dirty = False # isolate the delete client._session._vacuum_dirty = False # isolate the delete
assert await client.delete_document(doc.id) is True assert await client.delete_document(doc.id) is True
assert client._vacuum_dirty is True assert client._session._vacuum_dirty is True
async def test_delete_rolls_back_on_partial_failure(temp_db_path, monkeypatch): async def test_delete_rolls_back_on_partial_failure(temp_db_path, monkeypatch):
@ -1284,9 +1284,9 @@ async def test_delete_missing_id_returns_false_without_vacuum(temp_db_path):
"""Deleting an id that doesn't exist returns False and owes no vacuum (the """Deleting an id that doesn't exist returns False and owes no vacuum (the
existence check is inside the lock, so a no-op delete stays a no-op).""" existence check is inside the lock, so a no-op delete stays a no-op)."""
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
client._vacuum_dirty = False client._session._vacuum_dirty = False
assert await client.delete_document("does-not-exist") is False assert await client.delete_document("does-not-exist") is False
assert client._vacuum_dirty is False assert client._session._vacuum_dirty is False
@pytest.mark.vcr() @pytest.mark.vcr()
@ -2843,7 +2843,7 @@ async def test_import_documents_schedules_vacuum_per_config(temp_db_path, auto_v
) )
] ]
) )
await asyncio.gather(*client._vacuum_tasks) await asyncio.gather(*client._session._vacuum_tasks)
assert vacuum.await_count == (1 if auto_vacuum else 0) assert vacuum.await_count == (1 if auto_vacuum else 0)

View file

@ -485,6 +485,23 @@ class TestOperationsThatNeedOneDatabase:
with pytest.raises(AmbiguousDatabaseError, match="clients_for"): with pytest.raises(AmbiguousDatabaseError, match="clients_for"):
await rag.vacuum() await rag.vacuum()
@pytest.mark.asyncio
async def test_a_set_has_no_store_of_its_own(self, tmp_path):
"""A store and its repositories belong to one database. `clients_for`
reaches the one holding a given database."""
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha one"])
async with HaikuRAG(config=config) as rag:
for name in (
"store",
"document_repository",
"chunk_repository",
"document_item_repository",
):
with pytest.raises(AttributeError, match=name):
getattr(rag, name)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_a_selected_database_is_still_writable(self, tmp_path): async def test_a_selected_database_is_still_writable(self, tmp_path):
"""Naming one of the set is how a write picks its database.""" """Naming one of the set is how a write picks its database."""

View file

@ -2,7 +2,7 @@ import asyncio
import pytest import pytest
import haiku.rag.client as client_mod import haiku.rag.client.session as session_mod
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.client.documents import _refresh_doc_metadata from haiku.rag.client.documents import _refresh_doc_metadata
from haiku.rag.config import get_config from haiku.rag.config import get_config
@ -23,7 +23,7 @@ async def test_schedule_vacuum_is_debounced(temp_db_path, monkeypatch):
"""Rapid writes within the throttle window schedule only one background """Rapid writes within the throttle window schedule only one background
vacuum; once the interval elapses, a new one is scheduled.""" vacuum; once the interval elapses, a new one is scheduled."""
t = {"now": 1000.0} t = {"now": 1000.0}
monkeypatch.setattr(client_mod, "monotonic", lambda: t["now"]) monkeypatch.setattr(session_mod, "monotonic", lambda: t["now"])
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
calls: list[int] = [] calls: list[int] = []
@ -34,13 +34,13 @@ async def test_schedule_vacuum_is_debounced(temp_db_path, monkeypatch):
monkeypatch.setattr(client.store, "vacuum", fake_vacuum) monkeypatch.setattr(client.store, "vacuum", fake_vacuum)
for _ in range(3): for _ in range(3):
client._schedule_vacuum() client._session.schedule_vacuum()
await asyncio.gather(*client._vacuum_tasks) await asyncio.gather(*client._session._vacuum_tasks)
assert len(calls) == 1 # debounced within the interval assert len(calls) == 1 # debounced within the interval
t["now"] += client_mod._VACUUM_MIN_INTERVAL_S + 1 t["now"] += session_mod._VACUUM_MIN_INTERVAL_S + 1
client._schedule_vacuum() client._session.schedule_vacuum()
await asyncio.gather(*client._vacuum_tasks) await asyncio.gather(*client._session._vacuum_tasks)
assert len(calls) == 2 # interval elapsed -> a new vacuum scheduled assert len(calls) == 2 # interval elapsed -> a new vacuum scheduled
@ -49,7 +49,7 @@ async def test_debounced_writes_still_collapse_on_close(temp_db_path, monkeypatc
"""Even when scheduled vacuums after the first are debounced, the writes are """Even when scheduled vacuums after the first are debounced, the writes are
marked dirty so the close-time drain runs a final collapse.""" marked dirty so the close-time drain runs a final collapse."""
t = {"now": 1000.0} t = {"now": 1000.0}
monkeypatch.setattr(client_mod, "monotonic", lambda: t["now"]) monkeypatch.setattr(session_mod, "monotonic", lambda: t["now"])
calls: list[int] = [] calls: list[int] = []
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
@ -59,13 +59,13 @@ async def test_debounced_writes_still_collapse_on_close(temp_db_path, monkeypatc
monkeypatch.setattr(client.store, "vacuum", fake_vacuum) monkeypatch.setattr(client.store, "vacuum", fake_vacuum)
client._schedule_vacuum() # schedules the first background pass client._session.schedule_vacuum() # schedules the first background pass
client._schedule_vacuum() # debounced (no task) client._session.schedule_vacuum() # debounced (no task)
await client._await_vacuum_tasks() await client._session.drain_vacuum()
# one scheduled background pass + one final collapse on drain # one scheduled background pass + one final collapse on drain
assert len(calls) == 2 assert len(calls) == 2
assert client._vacuum_dirty is False assert client._session._vacuum_dirty is False
@pytest.mark.asyncio @pytest.mark.asyncio
@ -82,7 +82,7 @@ async def test_metadata_refresh_sweep_schedules_vacuum(temp_db_path):
metadata={"source_revision": "r1"}, metadata={"source_revision": "r1"},
) )
# Isolate the refresh: the import already scheduled a vacuum. # Isolate the refresh: the import already scheduled a vacuum.
client._vacuum_dirty = False client._session._vacuum_dirty = False
await _refresh_doc_metadata( await _refresh_doc_metadata(
client, client,
@ -91,7 +91,7 @@ async def test_metadata_refresh_sweep_schedules_vacuum(temp_db_path):
user_metadata={}, user_metadata={},
source_metadata={"source_revision": "r2", "md5": "same"}, source_metadata={"source_revision": "r2", "md5": "same"},
) )
assert client._vacuum_dirty is True assert client._session._vacuum_dirty is True
@pytest.mark.asyncio @pytest.mark.asyncio

View file

@ -205,7 +205,7 @@ async def test_existing_database_checks_migrations(monkeypatch, temp_db_path):
async def _wait_for_background_vacuum(client): async def _wait_for_background_vacuum(client):
"""Wait for any in-flight background vacuum tasks to complete.""" """Wait for any in-flight background vacuum tasks to complete."""
await client._await_vacuum_tasks() await client._session.drain_vacuum()
@pytest.mark.vcr() @pytest.mark.vcr()
@ -438,7 +438,8 @@ async def test_close_suppresses_failing_drain_vacuum(temp_db_path, monkeypatch):
raise RuntimeError("vacuum boom") raise RuntimeError("vacuum boom")
# Writes happened, so close owes a final vacuum — force that drain branch. # Writes happened, so close owes a final vacuum — force that drain branch.
client._vacuum_dirty = True assert client._session is not None
client._session._vacuum_dirty = True
monkeypatch.setattr(client.store, "vacuum", boom) monkeypatch.setattr(client.store, "vacuum", boom)
# Must not raise despite the drain vacuum erroring. # Must not raise despite the drain vacuum erroring.