Name the database a write wrote to

Reads tagged the document they returned and writes did not, so creating in
`alpha` came back with `source=None` while reading the same document came back
with `"alpha"`. Every outward write result goes through the owning session.
This commit is contained in:
Yiorgis Gozadinos 2026-08-28 08:41:45 +03:00
parent cad472e999
commit 202cbd2d3f
No known key found for this signature in database
3 changed files with 114 additions and 1 deletions

View file

@ -184,6 +184,7 @@ async def _store_document_with_chunks(
if session.config.storage.auto_vacuum: if session.config.storage.auto_vacuum:
session.schedule_vacuum() session.schedule_vacuum()
session.name(stored_doc)
return stored_doc return stored_doc
@ -237,6 +238,7 @@ async def _update_document_with_chunks(
if session.config.storage.auto_vacuum: if session.config.storage.auto_vacuum:
session.schedule_vacuum() session.schedule_vacuum()
session.name(updated_doc)
return updated_doc return updated_doc
@ -343,7 +345,7 @@ async def _store_documents_with_chunks(
if session.config.storage.auto_vacuum: if session.config.storage.auto_vacuum:
session.schedule_vacuum() session.schedule_vacuum()
return created return session.name_all(created)
async def import_documents( async def import_documents(
@ -402,7 +404,9 @@ async def _refresh_doc_metadata(
# The vacuum is debounced, and document_meta is tiny, so this is cheap. # The vacuum is debounced, and document_meta is tiny, so this is cheap.
if session.config.storage.auto_vacuum: if session.config.storage.auto_vacuum:
session.schedule_vacuum() session.schedule_vacuum()
session.name(result)
return result return result
session.name(doc)
return doc return doc
@ -879,6 +883,7 @@ async def update_document(
updated = await session.document_repository.update_meta(existing_doc) updated = await session.document_repository.update_meta(existing_doc)
if session.config.storage.auto_vacuum: if session.config.storage.auto_vacuum:
session.schedule_vacuum() session.schedule_vacuum()
session.name(updated)
return updated return updated
if chunks is not None: if chunks is not None:

View file

@ -174,6 +174,12 @@ class SingleDatabaseSession:
document.source = self.source document.source = self.source
return document return document
def name_all(self, documents: "list[Document]") -> "list[Document]":
"""`documents`, each told which database it came from."""
for document in documents:
document.source = self.source
return documents
async def get_document_by_id(self, document_id: str) -> "Document | None": async def get_document_by_id(self, document_id: str) -> "Document | None":
return self.name(await self.document_repository.get_by_id(document_id)) return self.name(await self.document_repository.get_by_id(document_id))

View file

@ -5,6 +5,7 @@ from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel from docling_core.types.doc.labels import DocItemLabel
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.client.session import SingleDatabaseSession
from haiku.rag.config import get_config from haiku.rag.config import get_config
from haiku.rag.store.models import Chunk from haiku.rag.store.models import Chunk
from tests.multi_db.helpers import ( from tests.multi_db.helpers import (
@ -157,6 +158,107 @@ class TestLookupByIdentifier:
assert await rag.get_document_by_uri("test://nowhere") is None assert await rag.get_document_by_uri("test://nowhere") is None
class TestWritesNameTheirDatabase:
"""A write returns the document it wrote, and it came from a database. A
read of the same document names it, so the write has to as well."""
@staticmethod
def _doc(text: str):
from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel
doc = DoclingDocument(name=text)
doc.add_text(label=DocItemLabel.TEXT, text=text)
return doc
@pytest.mark.asyncio
async def test_import_names_the_database(self, tmp_path):
config = _config(tmp_path, ["alpha"])
dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag:
written = await rag.import_document(
self._doc("cats"),
[Chunk(content="cats", embedding=[0.1] * dim, order=0)],
uri="test://alpha/cats",
)
assert written.id is not None
read = await rag.get_document_by_id(written.id)
assert written.source == "alpha"
assert read is not None and read.source == written.source
@pytest.mark.asyncio
async def test_a_batch_import_names_every_document(self, tmp_path):
from haiku.rag.client.documents import DocumentImport
config = _config(tmp_path, ["alpha"])
dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag:
written = await rag.import_documents(
[
DocumentImport(
docling_document=self._doc(text),
chunks=[Chunk(content=text, embedding=[0.1] * dim, order=0)],
uri=f"test://alpha/{text}",
)
for text in ("cats", "dogs")
]
)
assert [d.source for d in written] == ["alpha", "alpha"]
@pytest.mark.asyncio
async def test_a_metadata_only_update_names_the_database(self, tmp_path):
"""Changing only metadata rewrites the row without re-chunking, so it
never reaches the paths that name a document on the way through."""
config = _config(tmp_path, ["alpha"])
dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag:
stored = await rag.import_document(
self._doc("cats"),
[Chunk(content="cats", embedding=[0.1] * dim, order=0)],
uri="test://alpha/cats",
)
assert stored.id is not None
updated = await rag.update_document(
stored.id, title="Cats", metadata={"k": "v"}
)
assert updated is not None
assert updated.title == "Cats"
assert updated.source == "alpha"
@pytest.mark.asyncio
async def test_the_revision_short_circuit_names_the_database(self, tmp_path):
"""`create_document_from_source` refreshes metadata in place when the
revision is unchanged, returning the document it rewrote."""
from haiku.rag.client.documents import _refresh_doc_metadata
config = _config(tmp_path, ["alpha"])
dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag:
stored = await rag.import_document(
self._doc("cats"),
[Chunk(content="cats", embedding=[0.1] * dim, order=0)],
uri="test://alpha/cats",
)
assert isinstance(rag._session, SingleDatabaseSession)
refreshed = await _refresh_doc_metadata(
rag._session,
stored,
title="Cats",
user_metadata={"k": "v"},
source_metadata=None,
)
assert refreshed.source == "alpha"
class TestDocumentsNameTheirDatabase: class TestDocumentsNameTheirDatabase:
"""A listing that spans databases is unreadable when the documents do not """A listing that spans databases is unreadable when the documents do not
say which one they came from, the same reason a search result carries one.""" say which one they came from, the same reason a search result carries one."""