Ask the session which database an operation works on

One session field, whichever kind it is, with the covered databases derived
from it rather than kept beside it. `_single_session` returns the database a
write works on in place of ten guards that only asserted one existed, and the
single-database operations move to the session that owns them.
This commit is contained in:
Yiorgis Gozadinos 2026-08-25 09:28:20 +03:00
parent e8390ca747
commit 028217b7c0
No known key found for this signature in database
4 changed files with 147 additions and 109 deletions

View file

@ -139,13 +139,37 @@ class HaikuRAG:
self._read_only = read_only
self._requested_sources = sources
self._clients: dict[str, HaikuRAG] = {}
self._federated: dict[str, str] = {}
self._clients_lock = asyncio.Lock()
self._source: str | None = None
self._session: SingleDatabaseSession | None = None
self._federated_session: FederatedSession | None = None
self._session: SingleDatabaseSession | FederatedSession | None = None
self._owns_session = True
@property
def _federated(self) -> dict[str, str]:
"""The databases covered, name to location, empty while covering one.
Derived from the session rather than stored beside it: two answers to
"how many databases?" is what left every call site deciding for itself.
"""
if isinstance(self._session, FederatedSession):
return self._session.locations
return {}
def _single_session(self, operation: str) -> SingleDatabaseSession:
"""The one database this operation works on.
Writing, rebuilding and vacuuming all name a database, so they start
here: the session they get cannot be a set, and the refusal is the same
sentence whichever operation asked.
"""
if isinstance(self._session, SingleDatabaseSession):
return self._session
covered = ", ".join(sorted(self._federated))
raise AmbiguousDatabaseError(
f"{operation} works on one database, and this client covers "
f"{covered}; select one with clients_for([name])"
)
@property
def store(self) -> Store:
"""The store of the database this client opened.
@ -153,25 +177,25 @@ class HaikuRAG:
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:
if not isinstance(self._session, SingleDatabaseSession):
raise AttributeError("store")
return self._session.store
@property
def document_repository(self) -> DocumentRepository:
if self._session is None:
if not isinstance(self._session, SingleDatabaseSession):
raise AttributeError("document_repository")
return self._session.document_repository
@property
def chunk_repository(self) -> ChunkRepository:
if self._session is None:
if not isinstance(self._session, SingleDatabaseSession):
raise AttributeError("chunk_repository")
return self._session.chunk_repository
@property
def document_item_repository(self) -> DocumentItemRepository:
if self._session is None:
if not isinstance(self._session, SingleDatabaseSession):
raise AttributeError("document_item_repository")
return self._session.document_item_repository
@ -248,8 +272,7 @@ class HaikuRAG:
f"{', '.join(sorted(selected))}; name the one to create with "
"sources=[name]"
)
self._federated = selected
self._federated_session = FederatedSession(
self._session = FederatedSession(
DatabaseScope(
tuple(
DatabaseRef.configured(name, location)
@ -287,9 +310,9 @@ class HaikuRAG:
databases is typically queried a few at a time, and a database nobody
asked for must not be able to fail a query, or be opened for nothing.
"""
assert self._federated_session is not None
assert isinstance(self._session, FederatedSession)
names = _without_repeats(names)
sessions = await self._federated_session.sessions_for(names)
sessions = await self._session.sessions_for(names)
return [
self._facade_for(name, session)
for name, session in zip(names, sessions, strict=True)
@ -353,14 +376,13 @@ class HaikuRAG:
# Branch on what this client covers, not on what it happened to open:
# a federating client that answered no query has nothing open and no
# store either.
if self._federated:
assert self._federated_session is not None
if isinstance(self._session, FederatedSession):
# The wrappers built over covered databases hold rerankers of their
# own; the databases themselves are the federated session's to close.
for facade in self._clients.values():
await facade._release_own()
self._clients.clear()
await self._federated_session.aclose()
await self._session.aclose()
# The set shares this client's embedder and reranker, so this is the
# only place they are closed — and only if anything built them.
await self._aclose_cached("embedder")
@ -397,11 +419,11 @@ class HaikuRAG:
await aclose_quietly(cached, name)
async def _await_vacuum_tasks(self) -> None:
if self._session is not None:
if isinstance(self._session, SingleDatabaseSession):
await self._session.drain_vacuum()
def _schedule_vacuum(self) -> None:
if self._session is not None:
if isinstance(self._session, SingleDatabaseSession):
self._session.schedule_vacuum()
# =========================================================================
@ -465,7 +487,7 @@ class HaikuRAG:
) -> Document:
from haiku.rag.client.documents import create_document
self._require_one_database("create_document")
self._single_session("create_document")
return await create_document(self, content, uri, title, metadata, format)
@ -479,7 +501,7 @@ class HaikuRAG:
) -> Document:
from haiku.rag.client.documents import import_document
self._require_one_database("import_document")
self._single_session("import_document")
return await import_document(
self, docling_document, chunks, uri, title, metadata
@ -491,7 +513,7 @@ class HaikuRAG:
) -> list[Document]:
from haiku.rag.client.documents import import_documents
self._require_one_database("import_documents")
self._single_session("import_documents")
return await import_documents(self, imports)
@ -508,7 +530,7 @@ class HaikuRAG:
) -> Document | list[Document]:
from haiku.rag.client.documents import create_document_from_source
self._require_one_database("create_document_from_source")
self._single_session("create_document_from_source")
return await create_document_from_source(
self,
@ -534,7 +556,7 @@ class HaikuRAG:
) -> Document:
from haiku.rag.client.documents import update_document
self._require_one_database("update_document")
self._single_session("update_document")
return await update_document(
self,
@ -560,7 +582,9 @@ class HaikuRAG:
return await self._from_any_covered(
lambda owner: owner.get_document_by_id(document_id)
)
return self._name(await self.document_repository.get_by_id(document_id))
return await self._single_session("get_document_by_id").get_document_by_id(
document_id
)
async def get_chunk_by_id(self, chunk_id: str) -> Chunk | None:
"""Get a chunk by its ID.
@ -616,7 +640,9 @@ class HaikuRAG:
return await self._from_any_covered(
lambda owner: owner.get_document_by_uri(uri)
)
return self._name(await self.document_repository.get_by_uri(uri))
return await self._single_session("get_document_by_uri").get_document_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).
@ -644,46 +670,10 @@ class HaikuRAG:
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
self._require_one_database("delete_document")
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
``metadata.parent_uri``."""
return await self._single_session("delete_document").delete_document(
document_id
)
async def list_documents(
self,
@ -724,12 +714,9 @@ class HaikuRAG:
]
start = offset or 0
return merged[start:] if limit is None else merged[start : start + limit]
documents = await self.document_repository.list_all(
return await self._single_session("list_documents").list_documents(
limit=limit, offset=offset, filter=filter, include_content=include_content
)
for document in documents:
document.source = self._source
return documents
async def count_documents(self, filter: str | None = None) -> int:
"""Count documents with optional filtering.
@ -750,29 +737,6 @@ class HaikuRAG:
return sum(counts)
return await self.document_repository.count(filter=filter)
def _require_one_database(self, operation: str) -> None:
"""Refuse an operation that has no meaning across a set of databases.
Writing, rebuilding and vacuuming all have to name a database. Raised as
a domain error rather than surfacing the missing repository, so a caller
can tell an unsupported selection from a bug.
"""
if self._federated:
raise AmbiguousDatabaseError(
f"{operation} works on one database, and this client covers "
f"{', '.join(sorted(self._federated))}; select one with "
"clients_for([name])"
)
def _name(self, document: Document | None) -> Document | None:
"""`document`, told which configured database it came from.
None where no database is named, as with a single ``lancedb.uri``.
"""
if document is not None:
document.source = self._source
return document
async def _from_any_covered(
self, lookup: "Callable[[HaikuRAG], Coroutine[Any, Any, Any]]"
) -> Any:
@ -868,7 +832,7 @@ class HaikuRAG:
) -> list:
from haiku.rag.client.search import visualize_chunk
self._require_one_database("visualize_chunk")
self._single_session("visualize_chunk")
return await visualize_chunk(self, chunk, refs, expand)
@ -877,15 +841,14 @@ class HaikuRAG:
) -> AsyncGenerator[str, None]:
from haiku.rag.client.rebuild import rebuild_database
self._require_one_database("rebuild_database")
self._single_session("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."""
self._require_one_database("vacuum")
await self.store.vacuum()
await self._single_session("vacuum").store.vacuum()
def close(self):
"""Close the underlying store connection.
@ -893,8 +856,7 @@ class HaikuRAG:
A client covering one of a set borrows that database and never closes
it: the set opened it and the set closes it.
"""
self._require_one_database("close")
session = self._single_session("close")
if not self._owns_session:
return
assert self._session is not None
self._session.close()
session.close()

View file

@ -2,7 +2,7 @@ import asyncio
import logging
from pathlib import Path
from time import monotonic
from typing import Any
from typing import TYPE_CHECKING, Any
from haiku.rag.client.scope import DatabaseRef, DatabaseScope
from haiku.rag.config import AppConfig
@ -17,6 +17,9 @@ 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
@ -165,6 +168,73 @@ class SingleDatabaseSession:
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.
@ -209,6 +279,11 @@ class FederatedSession:
self._sessions: dict[str, SingleDatabaseSession] = {}
self._lock = asyncio.Lock()
@property
def locations(self) -> dict[str, str]:
"""The databases covered, name to location as configured."""
return {name: ref.uri or str(ref.db_path) for name, ref in self._refs.items()}
async def sessions_for(self, names: list[str]) -> list[SingleDatabaseSession]:
"""The sessions for these databases, opening any not yet open.

View file

@ -155,9 +155,9 @@ class TestOpeningDatabases:
await _seed(config, name, [f"{name} document about cats"])
async with HaikuRAG(config=config) as rag:
assert rag._federated_session is not None
assert isinstance(rag._session, FederatedSession)
barrier = asyncio.Barrier(len(names))
open_one = rag._federated_session._open
open_one = rag._session._open
async def gated(ref):
# Every open has to be in flight before any of them finishes, so
@ -165,7 +165,7 @@ class TestOpeningDatabases:
await barrier.wait()
return await open_one(ref)
rag._federated_session._open = gated
rag._session._open = gated
clients = await asyncio.wait_for(rag.clients_for(names), timeout=15)
assert {client._source for client in clients} == set(names)
@ -182,8 +182,8 @@ class TestOpeningDatabases:
with pytest.raises(SourceUnavailableError, match="beta"):
await rag.clients_for(["alpha", "beta"])
assert rag._federated_session is not None
assert set(rag._federated_session._sessions) == {"alpha"}
assert isinstance(rag._session, FederatedSession)
assert set(rag._session._sessions) == {"alpha"}
@pytest.mark.asyncio
async def test_a_database_named_twice_is_opened_once(self, tmp_path):
@ -358,9 +358,9 @@ class TestClosingASet:
drained: list[str | None] = []
async with HaikuRAG(config=config, read_only=True) as rag:
assert rag._federated_session is not None
assert isinstance(rag._session, FederatedSession)
await rag.clients_for(["alpha", "beta"])
for name, session in rag._federated_session._sessions.items():
for name, session in rag._session._sessions.items():
original = session.store.embedder.aclose
drain = session.drain_vacuum
@ -932,8 +932,8 @@ class TestRerankerFusion:
rag = HaikuRAG(config=config)
await rag.__aenter__()
await rag.clients_for(["alpha", "beta"])
assert rag._federated_session is not None
sessions = rag._federated_session._sessions
assert isinstance(rag._session, FederatedSession)
sessions = rag._session._sessions
async def boom():
raise RuntimeError("close failed")
@ -945,7 +945,7 @@ class TestRerankerFusion:
# The failure is swallowed, and the sibling is still closed after it.
assert rag._clients == {}
assert rag._federated_session._sessions == {}
assert rag._session._sessions == {}
assert not beta.db.is_open()
@pytest.mark.asyncio

View file

@ -3,6 +3,7 @@ import asyncio
import pytest
from haiku.rag.client import HaikuRAG
from haiku.rag.client.session import SingleDatabaseSession
from haiku.rag.store.engine import Store
@ -438,7 +439,7 @@ async def test_close_suppresses_failing_drain_vacuum(temp_db_path, monkeypatch):
raise RuntimeError("vacuum boom")
# Writes happened, so close owes a final vacuum — force that drain branch.
assert client._session is not None
assert isinstance(client._session, SingleDatabaseSession)
client._session._vacuum_dirty = True
monkeypatch.setattr(client.store, "vacuum", boom)