Serialize multi-table document writes and bound update version churn
This commit is contained in:
parent
ab9fea4833
commit
4bc52f0710
7 changed files with 218 additions and 113 deletions
|
|
@ -74,30 +74,31 @@ async def _store_document_with_chunks(
|
||||||
"""
|
"""
|
||||||
chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
|
chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
|
||||||
|
|
||||||
versions = await client.store.current_table_versions()
|
async with client.store._write_lock:
|
||||||
|
versions = await client.store.current_table_versions()
|
||||||
|
|
||||||
created_doc = await client.document_repository.create(document)
|
created_doc = await client.document_repository.create(document)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
assert created_doc.id is not None, (
|
assert created_doc.id is not None, (
|
||||||
"Document ID should not be None after creation"
|
"Document ID should not be None after creation"
|
||||||
)
|
)
|
||||||
for order, chunk in enumerate(chunks):
|
for order, chunk in enumerate(chunks):
|
||||||
chunk.document_id = created_doc.id
|
chunk.document_id = created_doc.id
|
||||||
chunk.order = order
|
chunk.order = order
|
||||||
|
|
||||||
await client.chunk_repository.create(chunks)
|
await client.chunk_repository.create(chunks)
|
||||||
|
|
||||||
items = extract_items(created_doc.id, docling_document)
|
items = extract_items(created_doc.id, docling_document)
|
||||||
await client.document_item_repository.create_items(created_doc.id, items)
|
await client.document_item_repository.create_items(created_doc.id, items)
|
||||||
|
|
||||||
if client._config.storage.auto_vacuum:
|
if client._config.storage.auto_vacuum:
|
||||||
client._schedule_vacuum()
|
client._schedule_vacuum()
|
||||||
|
|
||||||
return created_doc
|
return created_doc
|
||||||
except Exception:
|
except Exception:
|
||||||
await client.store.restore_table_versions(versions)
|
await client.store.restore_table_versions(versions)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
async def _update_document_with_chunks(
|
async def _update_document_with_chunks(
|
||||||
|
|
@ -124,36 +125,36 @@ async def _update_document_with_chunks(
|
||||||
|
|
||||||
chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
|
chunks = await ensure_chunks_embedded(client._config, chunks, client.embedder)
|
||||||
|
|
||||||
versions = await client.store.current_table_versions()
|
async with client.store._write_lock:
|
||||||
|
versions = await client.store.current_table_versions()
|
||||||
|
|
||||||
await client.chunk_repository.delete_by_document_id(document.id)
|
try:
|
||||||
|
updated_doc = await client.document_repository.update(document)
|
||||||
|
|
||||||
try:
|
assert updated_doc.id is not None
|
||||||
updated_doc = await client.document_repository.update(document)
|
for order, chunk in enumerate(chunks):
|
||||||
|
chunk.document_id = updated_doc.id
|
||||||
|
chunk.order = order
|
||||||
|
|
||||||
assert updated_doc.id is not None
|
await client.chunk_repository.replace_for_document(updated_doc.id, chunks)
|
||||||
for order, chunk in enumerate(chunks):
|
|
||||||
chunk.document_id = updated_doc.id
|
|
||||||
chunk.order = order
|
|
||||||
|
|
||||||
await client.chunk_repository.create(chunks)
|
if docling_document is not None:
|
||||||
|
items = extract_items(
|
||||||
|
updated_doc.id,
|
||||||
|
docling_document,
|
||||||
|
existing_picture_data=existing_picture_data,
|
||||||
|
)
|
||||||
|
await client.document_item_repository.replace_for_document(
|
||||||
|
updated_doc.id, items
|
||||||
|
)
|
||||||
|
|
||||||
if docling_document is not None:
|
if client._config.storage.auto_vacuum:
|
||||||
await client.document_item_repository.delete_by_document_id(updated_doc.id)
|
client._schedule_vacuum()
|
||||||
items = extract_items(
|
|
||||||
updated_doc.id,
|
|
||||||
docling_document,
|
|
||||||
existing_picture_data=existing_picture_data,
|
|
||||||
)
|
|
||||||
await client.document_item_repository.create_items(updated_doc.id, items)
|
|
||||||
|
|
||||||
if client._config.storage.auto_vacuum:
|
return updated_doc
|
||||||
client._schedule_vacuum()
|
except Exception:
|
||||||
|
await client.store.restore_table_versions(versions)
|
||||||
return updated_doc
|
raise
|
||||||
except Exception:
|
|
||||||
await client.store.restore_table_versions(versions)
|
|
||||||
raise
|
|
||||||
|
|
||||||
|
|
||||||
async def create_document(
|
async def create_document(
|
||||||
|
|
@ -235,33 +236,36 @@ async def _store_documents_with_chunks(
|
||||||
for _, chunks, _ in prepared
|
for _, chunks, _ in prepared
|
||||||
]
|
]
|
||||||
|
|
||||||
versions = await client.store.current_table_versions()
|
async with client.store._write_lock:
|
||||||
|
versions = await client.store.current_table_versions()
|
||||||
|
|
||||||
created = await client.document_repository.create([doc for doc, _, _ in prepared])
|
created = await client.document_repository.create(
|
||||||
|
[doc for doc, _, _ in prepared]
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
all_chunks: list[Chunk] = []
|
all_chunks: list[Chunk] = []
|
||||||
all_items = []
|
all_items = []
|
||||||
for doc, doc_chunks, docling_document in zip(
|
for doc, doc_chunks, docling_document in zip(
|
||||||
created, embedded, (d for _, _, d in prepared)
|
created, embedded, (d for _, _, d in prepared)
|
||||||
):
|
):
|
||||||
assert doc.id is not None
|
assert doc.id is not None
|
||||||
for order, chunk in enumerate(doc_chunks):
|
for order, chunk in enumerate(doc_chunks):
|
||||||
chunk.document_id = doc.id
|
chunk.document_id = doc.id
|
||||||
chunk.order = order
|
chunk.order = order
|
||||||
all_chunks.extend(doc_chunks)
|
all_chunks.extend(doc_chunks)
|
||||||
all_items.extend(extract_items(doc.id, docling_document))
|
all_items.extend(extract_items(doc.id, docling_document))
|
||||||
|
|
||||||
await client.chunk_repository.create(all_chunks)
|
await client.chunk_repository.create(all_chunks)
|
||||||
await client.document_item_repository.create_all(all_items)
|
await client.document_item_repository.create_all(all_items)
|
||||||
|
|
||||||
if client._config.storage.auto_vacuum:
|
if client._config.storage.auto_vacuum:
|
||||||
client._schedule_vacuum()
|
client._schedule_vacuum()
|
||||||
|
|
||||||
return created
|
return created
|
||||||
except Exception:
|
except Exception:
|
||||||
await client.store.restore_table_versions(versions)
|
await client.store.restore_table_versions(versions)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
async def import_documents(
|
async def import_documents(
|
||||||
|
|
|
||||||
|
|
@ -345,6 +345,7 @@ class Store:
|
||||||
self._skip_validation = skip_validation
|
self._skip_validation = skip_validation
|
||||||
self._skip_migration_check = skip_migration_check
|
self._skip_migration_check = skip_migration_check
|
||||||
self._vacuum_lock = asyncio.Lock()
|
self._vacuum_lock = asyncio.Lock()
|
||||||
|
self._write_lock = asyncio.Lock()
|
||||||
self._is_new_db = False
|
self._is_new_db = False
|
||||||
|
|
||||||
# Check if database exists (for local filesystem only)
|
# Check if database exists (for local filesystem only)
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ from lancedb.rerankers import RRFReranker
|
||||||
|
|
||||||
from haiku.rag.store.engine import Store, query_to_pydantic
|
from haiku.rag.store.engine import Store, query_to_pydantic
|
||||||
from haiku.rag.store.models.chunk import Chunk, SearchType
|
from haiku.rag.store.models.chunk import Chunk, SearchType
|
||||||
|
from haiku.rag.utils import escape_sql_string
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -42,6 +43,21 @@ class ChunkRepository:
|
||||||
return "\n".join(meta.headings) + "\n" + chunk.content
|
return "\n".join(meta.headings) + "\n" + chunk.content
|
||||||
return chunk.content
|
return chunk.content
|
||||||
|
|
||||||
|
def _to_record(self, chunk: Chunk, chunk_id: str):
|
||||||
|
assert chunk.document_id is not None
|
||||||
|
assert chunk.embedding is not None
|
||||||
|
return self.store.ChunkRecord(
|
||||||
|
id=chunk_id,
|
||||||
|
document_id=chunk.document_id,
|
||||||
|
content=chunk.content,
|
||||||
|
content_fts=self._contextualize_content(chunk),
|
||||||
|
metadata=json.dumps(
|
||||||
|
{k: v for k, v in chunk.metadata.items() if k != "order"}
|
||||||
|
),
|
||||||
|
order=int(chunk.order),
|
||||||
|
vector=chunk.embedding,
|
||||||
|
)
|
||||||
|
|
||||||
async def create(self, entity: Chunk | list[Chunk]) -> Chunk | list[Chunk]:
|
async def create(self, entity: Chunk | list[Chunk]) -> Chunk | list[Chunk]:
|
||||||
"""Create one or more chunks in the database.
|
"""Create one or more chunks in the database.
|
||||||
|
|
||||||
|
|
@ -55,18 +71,7 @@ class ChunkRepository:
|
||||||
assert entity.embedding is not None, "Chunk must have an embedding"
|
assert entity.embedding is not None, "Chunk must have an embedding"
|
||||||
|
|
||||||
chunk_id = str(uuid4())
|
chunk_id = str(uuid4())
|
||||||
|
chunk_record = self._to_record(entity, chunk_id)
|
||||||
chunk_record = self.store.ChunkRecord(
|
|
||||||
id=chunk_id,
|
|
||||||
document_id=entity.document_id,
|
|
||||||
content=entity.content,
|
|
||||||
content_fts=self._contextualize_content(entity),
|
|
||||||
metadata=json.dumps(
|
|
||||||
{k: v for k, v in entity.metadata.items() if k != "order"}
|
|
||||||
),
|
|
||||||
order=int(entity.order),
|
|
||||||
vector=entity.embedding,
|
|
||||||
)
|
|
||||||
|
|
||||||
await self.store.chunks_table.add([chunk_record])
|
await self.store.chunks_table.add([chunk_record])
|
||||||
|
|
||||||
|
|
@ -88,19 +93,7 @@ class ChunkRepository:
|
||||||
for chunk in chunks:
|
for chunk in chunks:
|
||||||
chunk_id = str(uuid4())
|
chunk_id = str(uuid4())
|
||||||
|
|
||||||
assert chunk.document_id is not None
|
chunk_record = self._to_record(chunk, chunk_id)
|
||||||
assert chunk.embedding is not None
|
|
||||||
chunk_record = self.store.ChunkRecord(
|
|
||||||
id=chunk_id,
|
|
||||||
document_id=chunk.document_id,
|
|
||||||
content=chunk.content,
|
|
||||||
content_fts=self._contextualize_content(chunk),
|
|
||||||
metadata=json.dumps(
|
|
||||||
{k: v for k, v in chunk.metadata.items() if k != "order"}
|
|
||||||
),
|
|
||||||
order=int(chunk.order),
|
|
||||||
vector=chunk.embedding,
|
|
||||||
)
|
|
||||||
chunk_records.append(chunk_record)
|
chunk_records.append(chunk_record)
|
||||||
chunk.id = chunk_id
|
chunk.id = chunk_id
|
||||||
|
|
||||||
|
|
@ -109,6 +102,38 @@ class ChunkRepository:
|
||||||
|
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
|
async def replace_for_document(
|
||||||
|
self, document_id: str, chunks: list[Chunk]
|
||||||
|
) -> list[Chunk]:
|
||||||
|
"""Replace all chunks for a document with one scoped merge operation."""
|
||||||
|
self.store._assert_writable()
|
||||||
|
|
||||||
|
if not chunks:
|
||||||
|
await self.delete_by_document_id(document_id)
|
||||||
|
return []
|
||||||
|
|
||||||
|
for chunk in chunks:
|
||||||
|
assert chunk.document_id == document_id, (
|
||||||
|
"All chunks must belong to the replaced document"
|
||||||
|
)
|
||||||
|
assert chunk.embedding is not None, "All chunks must have embeddings"
|
||||||
|
|
||||||
|
records = []
|
||||||
|
for chunk in chunks:
|
||||||
|
chunk_id = str(uuid4())
|
||||||
|
records.append(self._to_record(chunk, chunk_id))
|
||||||
|
chunk.id = chunk_id
|
||||||
|
|
||||||
|
safe_id = escape_sql_string(document_id)
|
||||||
|
await (
|
||||||
|
self.store.chunks_table.merge_insert(["document_id", "order"])
|
||||||
|
.when_matched_update_all()
|
||||||
|
.when_not_matched_insert_all()
|
||||||
|
.when_not_matched_by_source_delete(f"document_id = '{safe_id}'")
|
||||||
|
.execute(records)
|
||||||
|
)
|
||||||
|
return chunks
|
||||||
|
|
||||||
async def get_by_id(self, entity_id: str) -> Chunk | None:
|
async def get_by_id(self, entity_id: str) -> Chunk | None:
|
||||||
"""Get a chunk by its ID."""
|
"""Get a chunk by its ID."""
|
||||||
results = await query_to_pydantic(
|
results = await query_to_pydantic(
|
||||||
|
|
|
||||||
|
|
@ -62,7 +62,13 @@ class DocumentRepository:
|
||||||
else datetime.now(),
|
else datetime.now(),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _to_record(self, entity: Document, doc_id: str, now: str) -> DocumentRecord:
|
def _to_record(
|
||||||
|
self,
|
||||||
|
entity: Document,
|
||||||
|
doc_id: str,
|
||||||
|
created_at: str,
|
||||||
|
updated_at: str,
|
||||||
|
) -> DocumentRecord:
|
||||||
return DocumentRecord(
|
return DocumentRecord(
|
||||||
id=doc_id,
|
id=doc_id,
|
||||||
content=entity.content,
|
content=entity.content,
|
||||||
|
|
@ -72,8 +78,8 @@ class DocumentRepository:
|
||||||
docling_document=entity.docling_document,
|
docling_document=entity.docling_document,
|
||||||
docling_pages=entity.docling_pages,
|
docling_pages=entity.docling_pages,
|
||||||
docling_version=entity.docling_version,
|
docling_version=entity.docling_version,
|
||||||
created_at=now,
|
created_at=created_at,
|
||||||
updated_at=now,
|
updated_at=updated_at,
|
||||||
)
|
)
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
|
|
@ -94,7 +100,9 @@ class DocumentRepository:
|
||||||
if isinstance(entity, Document):
|
if isinstance(entity, Document):
|
||||||
doc_id = str(uuid4())
|
doc_id = str(uuid4())
|
||||||
now = datetime.now().isoformat()
|
now = datetime.now().isoformat()
|
||||||
await self.store.documents_table.add([self._to_record(entity, doc_id, now)])
|
await self.store.documents_table.add(
|
||||||
|
[self._to_record(entity, doc_id, now, now)]
|
||||||
|
)
|
||||||
entity.id = doc_id
|
entity.id = doc_id
|
||||||
entity.created_at = datetime.fromisoformat(now)
|
entity.created_at = datetime.fromisoformat(now)
|
||||||
entity.updated_at = datetime.fromisoformat(now)
|
entity.updated_at = datetime.fromisoformat(now)
|
||||||
|
|
@ -109,7 +117,7 @@ class DocumentRepository:
|
||||||
records = []
|
records = []
|
||||||
for document in documents:
|
for document in documents:
|
||||||
doc_id = str(uuid4())
|
doc_id = str(uuid4())
|
||||||
records.append(self._to_record(document, doc_id, now))
|
records.append(self._to_record(document, doc_id, now, now))
|
||||||
document.id = doc_id
|
document.id = doc_id
|
||||||
document.created_at = created_at
|
document.created_at = created_at
|
||||||
document.updated_at = created_at
|
document.updated_at = created_at
|
||||||
|
|
@ -199,20 +207,16 @@ class DocumentRepository:
|
||||||
now = datetime.now().isoformat()
|
now = datetime.now().isoformat()
|
||||||
entity.updated_at = datetime.fromisoformat(now)
|
entity.updated_at = datetime.fromisoformat(now)
|
||||||
|
|
||||||
# Update the record
|
record = self._to_record(
|
||||||
safe_id = escape_sql_string(entity.id)
|
entity,
|
||||||
await self.store.documents_table.update(
|
entity.id,
|
||||||
{
|
entity.created_at.isoformat() if entity.created_at else now,
|
||||||
"content": entity.content,
|
now,
|
||||||
"uri": entity.uri,
|
)
|
||||||
"title": entity.title,
|
await (
|
||||||
"metadata": json.dumps(entity.metadata),
|
self.store.documents_table.merge_insert("id")
|
||||||
"docling_document": entity.docling_document,
|
.when_matched_update_all()
|
||||||
"docling_pages": entity.docling_pages,
|
.execute([record])
|
||||||
"docling_version": entity.docling_version,
|
|
||||||
"updated_at": now,
|
|
||||||
},
|
|
||||||
where=f"id = '{safe_id}'",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return entity
|
return entity
|
||||||
|
|
|
||||||
|
|
@ -69,6 +69,31 @@ class DocumentItemRepository:
|
||||||
records = [self._to_record(item.document_id, item) for item in items]
|
records = [self._to_record(item.document_id, item) for item in items]
|
||||||
await self.store.document_items_table.add(records)
|
await self.store.document_items_table.add(records)
|
||||||
|
|
||||||
|
async def replace_for_document(
|
||||||
|
self, document_id: str, items: list[DocumentItem]
|
||||||
|
) -> None:
|
||||||
|
"""Replace all items for a document with one scoped merge operation."""
|
||||||
|
self.store._assert_writable()
|
||||||
|
|
||||||
|
if not items:
|
||||||
|
await self.delete_by_document_id(document_id)
|
||||||
|
return
|
||||||
|
|
||||||
|
for item in items:
|
||||||
|
assert item.document_id == document_id, (
|
||||||
|
"All items must belong to the replaced document"
|
||||||
|
)
|
||||||
|
|
||||||
|
safe_id = escape_sql_string(document_id)
|
||||||
|
records = [self._to_record(document_id, item) for item in items]
|
||||||
|
await (
|
||||||
|
self.store.document_items_table.merge_insert(["document_id", "self_ref"])
|
||||||
|
.when_matched_update_all()
|
||||||
|
.when_not_matched_insert_all()
|
||||||
|
.when_not_matched_by_source_delete(f"document_id = '{safe_id}'")
|
||||||
|
.execute(records)
|
||||||
|
)
|
||||||
|
|
||||||
async def get_all_items(self, document_id: str) -> list[DocumentItem]:
|
async def get_all_items(self, document_id: str) -> list[DocumentItem]:
|
||||||
"""Get all items for a document, sorted by position."""
|
"""Get all items for a document, sorted by position."""
|
||||||
safe_id = escape_sql_string(document_id)
|
safe_id = escape_sql_string(document_id)
|
||||||
|
|
|
||||||
|
|
@ -861,6 +861,52 @@ async def test_client_import_documents_empty(temp_db_path):
|
||||||
assert after == before
|
assert after == before
|
||||||
|
|
||||||
|
|
||||||
|
async def test_client_update_document_replaces_rows_with_bounded_versions(
|
||||||
|
temp_db_path,
|
||||||
|
):
|
||||||
|
"""Updating one document should replace stale rows with bounded versions."""
|
||||||
|
dim = Config.embeddings.model.vector_dim
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
created = await client.import_document(
|
||||||
|
_docling_doc("original", "Original body"),
|
||||||
|
[Chunk(content="Original body", embedding=[0.1] * dim, order=0)],
|
||||||
|
uri="mem://replace",
|
||||||
|
title="Replace",
|
||||||
|
)
|
||||||
|
assert created.id is not None
|
||||||
|
|
||||||
|
updated_docling = _docling_doc("updated", "Updated body")
|
||||||
|
updated_chunks = [
|
||||||
|
Chunk(content="Updated body A", embedding=[0.2] * dim, order=0),
|
||||||
|
Chunk(content="Updated body B", embedding=[0.3] * dim, order=1),
|
||||||
|
]
|
||||||
|
|
||||||
|
before = await client.store.current_table_versions()
|
||||||
|
updated = await client.update_document(
|
||||||
|
created.id,
|
||||||
|
docling_document=updated_docling,
|
||||||
|
chunks=updated_chunks,
|
||||||
|
)
|
||||||
|
after = await client.store.current_table_versions()
|
||||||
|
|
||||||
|
assert updated.id == created.id
|
||||||
|
assert after["documents"] - before["documents"] == 1
|
||||||
|
# Indexed LanceDB tables record one additional physical version for
|
||||||
|
# merge replacement in 0.30.x.
|
||||||
|
assert after["chunks"] - before["chunks"] <= 2
|
||||||
|
assert after["document_items"] - before["document_items"] <= 2
|
||||||
|
|
||||||
|
stored_chunks = await client.chunk_repository.get_by_document_id(created.id)
|
||||||
|
assert [chunk.content for chunk in stored_chunks] == [
|
||||||
|
"Updated body A",
|
||||||
|
"Updated body B",
|
||||||
|
]
|
||||||
|
stored_items = await client.document_item_repository.get_all_items(created.id)
|
||||||
|
assert len(stored_items) == 1
|
||||||
|
assert stored_items[0].text == "Updated body"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
async def test_client_ask(allow_model_requests, temp_db_path):
|
async def test_client_ask(allow_model_requests, temp_db_path):
|
||||||
"""Test asking questions returns answer and citations (VCR recorded)."""
|
"""Test asking questions returns answer and citations (VCR recorded)."""
|
||||||
|
|
|
||||||
|
|
@ -38,14 +38,14 @@ async def test_version_rollback_on_update_failure(temp_db_path):
|
||||||
base_content = "Base content"
|
base_content = "Base content"
|
||||||
created = await client.create_document(content=base_content)
|
created = await client.create_document(content=base_content)
|
||||||
|
|
||||||
# Patch chunk_repository.create to succeed then fail during update
|
# Patch chunk replacement to succeed then fail during update
|
||||||
orig_create = client.chunk_repository.create
|
orig_replace = client.chunk_repository.replace_for_document
|
||||||
|
|
||||||
async def succeed_then_fail(chunks):
|
async def succeed_then_fail(document_id, chunks):
|
||||||
await orig_create(chunks)
|
await orig_replace(document_id, chunks)
|
||||||
raise RuntimeError("update fail")
|
raise RuntimeError("update fail")
|
||||||
|
|
||||||
client.chunk_repository.create = succeed_then_fail
|
client.chunk_repository.replace_for_document = succeed_then_fail
|
||||||
|
|
||||||
# Attempt update
|
# Attempt update
|
||||||
with pytest.raises(RuntimeError):
|
with pytest.raises(RuntimeError):
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue