haiku.rag/tests/test_versioning.py
Yiorgis Gozadinos 2da294b850
Put multi-table writes behind one transaction boundary
Four call sites repeated lock, snapshot, try/except, restore. Each used
restore_table_versions, which restores in _tables() order — documents
first, contradicting RESTORE_TABLE_ORDER — and each caught Exception, so
a cancellation mid-write skipped rollback and left the earlier table
writes committed.

Store.write_transaction() holds the lock, snapshots under it, and rolls
back through _rollback_to_snapshot: RESTORE_TABLE_ORDER, shielded from
cancellation, absorbed cancellation re-delivered, rollback failure raised
with the original as cause. The two single-table update_meta sites keep
the bare lock.

The batch documents write moves inside the guarded body; it was outside
the try, so a failure there was never rolled back. Auto-vacuum is
scheduled after the transaction rather than inside it.

restore_table_versions is removed; those four sites were its only callers.
2026-08-19 14:05:40 +03:00

447 lines
17 KiB
Python

import asyncio
import pytest
from haiku.rag.client import HaikuRAG
from haiku.rag.store.engine import Store
@pytest.mark.vcr()
async def test_version_rollback_on_create_failure(temp_db_path):
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Patch chunk_repository.create to succeed then fail, triggering rollback
orig_create = client.chunk_repository.create
async def succeed_then_fail(chunks):
await orig_create(chunks)
raise RuntimeError("boom")
client.chunk_repository.create = succeed_then_fail
# Attempt to create document; expect failure and rollback
content = "Hello, rollback!"
with pytest.raises(RuntimeError):
await client.create_document(content=content)
# State should be restored (no documents/chunks)
docs = await client.list_documents()
assert len(docs) == 0
all_chunks = await client.chunk_repository.list_all()
assert len(all_chunks) == 0
@pytest.mark.vcr()
async def test_version_rollback_on_update_failure(temp_db_path):
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Create a valid document first
base_content = "Base content"
created = await client.create_document(content=base_content)
# Patch chunk replacement to succeed then fail during update
orig_replace = client.chunk_repository.replace_for_document
async def succeed_then_fail(document_id, chunks):
await orig_replace(document_id, chunks)
raise RuntimeError("update fail")
client.chunk_repository.replace_for_document = succeed_then_fail
# Attempt update
with pytest.raises(RuntimeError):
await client.update_document(
document_id=created.id,
content="Updated content",
)
# Content and chunks should remain the original
persisted = await client.get_document_by_id(created.id)
assert persisted is not None
assert persisted.content == base_content
original_chunks = await client.chunk_repository.get_by_document_id(created.id)
assert len(original_chunks) > 0
@pytest.mark.vcr()
async def test_cancellation_mid_write_rolls_back(temp_db_path):
"""A cancellation between two table writes must roll back, not leave the
chunks write committed without its document."""
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
orig_create = client.chunk_repository.create
async def succeed_then_cancel(chunks):
await orig_create(chunks)
raise asyncio.CancelledError()
client.chunk_repository.create = succeed_then_cancel
with pytest.raises(asyncio.CancelledError):
await client.create_document(content="cancelled mid-write")
assert await client.list_documents() == []
assert await client.chunk_repository.list_all() == []
@pytest.mark.vcr()
async def test_rollback_failure_keeps_the_original_cause(temp_db_path):
"""A failed rollback must not hide what failed first."""
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
async def failing_restore(versions, *, best_effort=False):
return [("chunks", RuntimeError("restore refused"))]
client.store._restore_tables = failing_restore
async def boom(chunks):
raise RuntimeError("original failure")
client.chunk_repository.create = boom
with pytest.raises(RuntimeError, match="rollback failed on: chunks") as excinfo:
await client.create_document(content="doomed")
assert isinstance(excinfo.value.__cause__, RuntimeError)
assert str(excinfo.value.__cause__) == "original failure"
@pytest.mark.vcr()
async def test_cancellation_during_rollback_is_redelivered(temp_db_path):
"""A cancellation arriving while rollback runs cannot cut it short, and is
delivered to the caller once the restore has completed."""
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
rollback_started = asyncio.Event()
rollback_finished = asyncio.Event()
async def slow_restore(versions, *, best_effort=False):
rollback_started.set()
await asyncio.sleep(0.05)
rollback_finished.set()
return []
client.store._restore_tables = slow_restore
async def boom(chunks):
raise RuntimeError("first failure")
client.chunk_repository.create = boom
task = asyncio.create_task(client.create_document(content="cancel in rollback"))
await rollback_started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert rollback_finished.is_set()
@pytest.mark.vcr()
async def test_batch_import_rolls_back_the_documents_write(temp_db_path):
"""The batch document write is inside the guarded body, so a failure after
it lands restores the documents table too."""
from haiku.rag.client.documents import DocumentImport
from tests.store.test_document_items import _docling_doc_with_picture
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
orig_create = client.document_repository.create
async def succeed_then_fail(documents):
await orig_create(documents)
raise RuntimeError("after the documents write")
client.document_repository.create = succeed_then_fail
with pytest.raises(RuntimeError, match="after the documents write"):
await client.import_documents(
[
DocumentImport(
docling_document=_docling_doc_with_picture(),
chunks=[],
uri="test://batch-rollback",
)
]
)
assert await client.store.documents_table.count_rows() == 0
async def test_new_database_does_not_run_upgrades(monkeypatch, temp_db_path):
def fail_if_called(*_args, **_kwargs):
raise AssertionError("run_pending_upgrades should not be called for new DB")
monkeypatch.setattr(
"haiku.rag.store.upgrades.run_pending_upgrades",
fail_if_called,
)
async with Store(temp_db_path, create=True):
pass
async def test_existing_database_checks_migrations(monkeypatch, temp_db_path):
async with Store(temp_db_path, create=True):
pass
from haiku.rag.store import upgrades
called = {"value": False}
original_get_pending = upgrades.get_pending_upgrades
def mark_called(*args, **kwargs):
called["value"] = True
return original_get_pending(*args, **kwargs)
monkeypatch.setattr(
"haiku.rag.store.upgrades.get_pending_upgrades",
mark_called,
)
# Opening an existing database should check for pending migrations
async with Store(temp_db_path):
pass
assert called["value"]
async def _wait_for_background_vacuum(client):
"""Wait for any in-flight background vacuum tasks to complete."""
await client._await_vacuum_tasks()
@pytest.mark.vcr()
async def test_vacuum_with_retention_threshold(temp_db_path):
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Create first document
await client.create_document(content="First document")
await _wait_for_background_vacuum(client)
# Create second document
await client.create_document(content="Second document")
await _wait_for_background_vacuum(client)
store = client.store
# Get initial version counts (should have multiple versions from creates)
initial_doc_versions = len(await store.documents_table.list_versions())
initial_chunk_versions = len(await store.chunks_table.list_versions())
assert initial_doc_versions > 1, "Should have multiple document table versions"
assert initial_chunk_versions > 1, "Should have multiple chunk table versions"
# Vacuum with default threshold (60 seconds) - should keep recent versions
# Note: vacuum may create new versions even when not cleaning up old ones
await store.vacuum()
after_default_doc_versions = len(await store.documents_table.list_versions())
after_default_chunk_versions = len(await store.chunks_table.list_versions())
# After vacuum with retention, version count should stay the same or increase
# (optimize may create new versions) but not decrease
assert after_default_doc_versions >= initial_doc_versions, (
"Default vacuum should not remove recent versions"
)
assert after_default_chunk_versions >= initial_chunk_versions, (
"Default vacuum should not remove recent versions"
)
# Vacuum with 0 threshold - should significantly reduce versions
await store.vacuum(retention_seconds=0)
after_zero_doc_versions = len(await store.documents_table.list_versions())
after_zero_chunk_versions = len(await store.chunks_table.list_versions())
# After aggressive vacuum, should have minimal versions (1-2)
# Note: optimize operation may create a version after cleanup
assert after_zero_doc_versions <= 2, (
f"Should have minimal document versions after vacuum(0), got {after_zero_doc_versions}"
)
assert after_zero_chunk_versions <= 2, (
f"Should have minimal chunk versions after vacuum(0), got {after_zero_chunk_versions}"
)
# And it should be significantly fewer than before
assert after_zero_doc_versions < initial_doc_versions, (
"Should have fewer versions after vacuum(0)"
)
assert after_zero_chunk_versions < initial_chunk_versions, (
"Should have fewer versions after vacuum(0)"
)
@pytest.mark.vcr()
async def test_vacuum_completes_before_context_exit(temp_db_path, monkeypatch):
"""Test that background vacuum completes when context manager exits."""
from haiku.rag.config import Config
# Set aggressive vacuum retention for this test
monkeypatch.setattr(Config.storage, "vacuum_retention_seconds", 0)
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Create multiple documents - each creation triggers automatic vacuum with retention=0
# This aggressively cleans up old versions between operations
for i in range(3):
await client.create_document(content=f"Test document {i}")
# After context exit, automatic vacuum should have kept versions minimal
async with Store(temp_db_path, create=True) as store:
final_versions = len(await store.documents_table.list_versions())
# With retention_seconds=0, vacuum aggressively cleans up between operations
# Should have very few versions remaining (1-2)
assert final_versions <= 2, (
f"Aggressive vacuum should keep minimal versions, got {final_versions}"
)
assert final_versions >= 1, "Should have at least one version remaining"
@pytest.mark.vcr()
async def test_aexit_awaits_background_vacuum(temp_db_path, monkeypatch):
"""__aexit__ must await any in-flight background vacuum, not just release the lock.
Background vacuum runs as an asyncio task; the event loop may not have scheduled
it yet when __aexit__ runs. Simply acquiring the vacuum lock (which is free until
the task actually starts) would let close() proceed before vacuum runs.
"""
from haiku.rag.config import Config
monkeypatch.setattr(Config.storage, "auto_vacuum", True)
vacuum_started = asyncio.Event()
vacuum_completed = asyncio.Event()
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
original_vacuum = client.store.vacuum
async def instrumented_vacuum(*args, **kwargs):
vacuum_started.set()
# Delay so __aexit__ would see an unstarted/incomplete task if it
# relied on the lock rather than awaiting the task directly.
await asyncio.sleep(0.05)
await original_vacuum(*args, **kwargs)
vacuum_completed.set()
client.store.vacuum = instrumented_vacuum
await client.create_document(content="triggers background vacuum")
assert vacuum_started.is_set(), "Background vacuum task never ran"
assert vacuum_completed.is_set(), "__aexit__ exited before vacuum finished"
@pytest.mark.vcr()
async def test_aexit_awaits_all_background_vacuums(temp_db_path, monkeypatch):
"""Multiple create_document calls schedule multiple vacuum tasks; __aexit__
must await all of them, not just the last-scheduled one.
Scenario: Task A acquires the vacuum lock and is slow. Task B is scheduled
while Task A still holds the lock — Task B sees the lock held and returns
immediately. If the client only tracks the most recently scheduled task,
__aexit__ awaits the fast no-op B and closes the connection while Task A
is still running.
"""
from haiku.rag.config import Config
monkeypatch.setattr(Config.storage, "auto_vacuum", True)
first_vacuum_completed = asyncio.Event()
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
call_count = 0
async def slow_vacuum(*_args, **_kwargs):
nonlocal call_count
call_count += 1
my_num = call_count
# Mimic the real vacuum's skip-if-running behavior.
if client.store._vacuum_lock.locked():
return
async with client.store._vacuum_lock:
if my_num == 1:
# Hold the lock longer than any other operation in the
# test so Task A cannot finish incidentally. __aexit__
# must explicitly wait for this task.
await asyncio.sleep(2.0)
first_vacuum_completed.set()
client.store.vacuum = slow_vacuum
await client.create_document(content="triggers first vacuum")
# Let Task A start and acquire the vacuum lock before scheduling B.
await asyncio.sleep(0.02)
await client.create_document(content="triggers second vacuum")
assert first_vacuum_completed.is_set(), (
"__aexit__ returned before the first vacuum task finished"
)
@pytest.mark.vcr()
async def test_auto_vacuum_disabled_skips_vacuum(temp_db_path, monkeypatch):
"""Test that auto_vacuum=False prevents automatic vacuum after operations."""
from haiku.rag.config import Config
# Disable auto-vacuum
monkeypatch.setattr(Config.storage, "auto_vacuum", False)
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Create multiple documents
for i in range(3):
await client.create_document(content=f"Test document {i}")
# Count versions - should accumulate without vacuum
doc_versions = len(await client.store.documents_table.list_versions())
chunk_versions = len(await client.store.chunks_table.list_versions())
# Without auto-vacuum, versions should accumulate (more than 3 from creates)
assert doc_versions >= 3, (
f"Without auto-vacuum, should have accumulated versions, got {doc_versions}"
)
assert chunk_versions >= 3, (
f"Without auto-vacuum, should have accumulated versions, got {chunk_versions}"
)
@pytest.mark.vcr()
async def test_auto_vacuum_enabled_triggers_vacuum(temp_db_path, monkeypatch):
"""Test that auto_vacuum=True (default) triggers vacuum after operations."""
from haiku.rag.config import Config
# Enable auto-vacuum with aggressive retention
monkeypatch.setattr(Config.storage, "auto_vacuum", True)
monkeypatch.setattr(Config.storage, "vacuum_retention_seconds", 0)
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
# Create multiple documents
for i in range(3):
await client.create_document(content=f"Test document {i}")
# After context exit, vacuum should have cleaned up
async with Store(temp_db_path, create=True) as store:
final_versions = len(await store.documents_table.list_versions())
# With auto_vacuum=True and retention=0, should have minimal versions
assert final_versions <= 2, (
f"With auto-vacuum enabled, should have minimal versions, got {final_versions}"
)
async def test_close_suppresses_failing_drain_vacuum(temp_db_path, monkeypatch):
"""A failing final vacuum on close must not raise out of __aexit__,
where it would mask an in-flight exception from the context body."""
client = HaikuRAG(db_path=temp_db_path, create=True)
await client.__aenter__()
calls: list[int] = []
async def boom(*args, **kwargs):
calls.append(1)
raise RuntimeError("vacuum boom")
# Writes happened, so close owes a final vacuum — force that drain branch.
client._vacuum_dirty = True
monkeypatch.setattr(client.store, "vacuum", boom)
# Must not raise despite the drain vacuum erroring.
await client.__aexit__(None, None, None)
assert calls, "drain vacuum should have been attempted"