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:
parent
cad472e999
commit
202cbd2d3f
3 changed files with 114 additions and 1 deletions
|
|
@ -184,6 +184,7 @@ async def _store_document_with_chunks(
|
|||
if session.config.storage.auto_vacuum:
|
||||
session.schedule_vacuum()
|
||||
|
||||
session.name(stored_doc)
|
||||
return stored_doc
|
||||
|
||||
|
||||
|
|
@ -237,6 +238,7 @@ async def _update_document_with_chunks(
|
|||
if session.config.storage.auto_vacuum:
|
||||
session.schedule_vacuum()
|
||||
|
||||
session.name(updated_doc)
|
||||
return updated_doc
|
||||
|
||||
|
||||
|
|
@ -343,7 +345,7 @@ async def _store_documents_with_chunks(
|
|||
if session.config.storage.auto_vacuum:
|
||||
session.schedule_vacuum()
|
||||
|
||||
return created
|
||||
return session.name_all(created)
|
||||
|
||||
|
||||
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.
|
||||
if session.config.storage.auto_vacuum:
|
||||
session.schedule_vacuum()
|
||||
session.name(result)
|
||||
return result
|
||||
session.name(doc)
|
||||
return doc
|
||||
|
||||
|
||||
|
|
@ -879,6 +883,7 @@ async def update_document(
|
|||
updated = await session.document_repository.update_meta(existing_doc)
|
||||
if session.config.storage.auto_vacuum:
|
||||
session.schedule_vacuum()
|
||||
session.name(updated)
|
||||
return updated
|
||||
|
||||
if chunks is not None:
|
||||
|
|
|
|||
|
|
@ -174,6 +174,12 @@ class SingleDatabaseSession:
|
|||
document.source = self.source
|
||||
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":
|
||||
return self.name(await self.document_repository.get_by_id(document_id))
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from docling_core.types.doc.document import DoclingDocument
|
|||
from docling_core.types.doc.labels import DocItemLabel
|
||||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.client.session import SingleDatabaseSession
|
||||
from haiku.rag.config import get_config
|
||||
from haiku.rag.store.models import Chunk
|
||||
from tests.multi_db.helpers import (
|
||||
|
|
@ -157,6 +158,107 @@ class TestLookupByIdentifier:
|
|||
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:
|
||||
"""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."""
|
||||
|
|
|
|||
Loading…
Reference in a new issue