check_source_accessible narrowed its handler to ValueError, but Path.exists re-raises errno values outside its ignored set (EACCES, ENAMETOOLONG). Those were swallowed before and now escaped into the rebuild sweep the guard exists to protect. Catch OSError too. Restore the arity guard in _common_path_prefix: without it an empty list raises from min() and a single label yields a prefix covering the whole path. Two tests would have hung rather than failed on regression (the vacuum skip and the protected-wait cancellation); both are now bounded. The import vacuum test raced against the done-callback that discards the task, and now spies on the call instead, with a negative control. Replace assertions that could not fail: blank-query search against an empty corpus, a batch flush counted against an empty table, a picture description asserting its own input state, and an FS scheme check with nothing on disk to resolve. The get_model matrix asserted only the returned type across 26 cases and now pins the per-provider settings. The three batching tests now count flushes, which revealed embed-only writes through chunks_table.add rather than _flush_rebuild_batch.
433 lines
16 KiB
Python
433 lines
16 KiB
Python
import re
|
|
|
|
import pytest
|
|
from lancedb.table import AsyncTable, AsyncTags
|
|
|
|
from haiku.rag.store import ReadOnlyError, Store
|
|
from haiku.rag.store.engine import RESTORE_TABLE_ORDER
|
|
from haiku.rag.store.models import Document
|
|
from haiku.rag.store.repositories.document import DocumentRepository
|
|
|
|
SAFETY_TAG_PATTERN = r"before-restore-\d{8}T\d{6}Z"
|
|
|
|
|
|
async def _doc_contents(store: Store) -> set[str]:
|
|
docs = await DocumentRepository(store).list_all(include_content=True)
|
|
return {d.content for d in docs}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_tag_restores_all_tables(temp_db_path):
|
|
"""A complete tag restores every table; rows added after the tag are
|
|
absent from the restored latest state, which stays writable."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
pre_restore_docs_version = await store.documents_table.version()
|
|
|
|
safety_tag = await store.restore_tag("release-1")
|
|
|
|
assert re.fullmatch(SAFETY_TAG_PATTERN, safety_tag)
|
|
assert await _doc_contents(store) == {"First document"}
|
|
|
|
# restore writes a NEW latest version; the table is not a read-only
|
|
# checkout and stays writable.
|
|
assert await store.documents_table.version() > pre_restore_docs_version
|
|
await repo.create(Document(content="Third document"))
|
|
assert await _doc_contents(store) == {"First document", "Third document"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_safety_tag_matches_pre_restore_state(temp_db_path):
|
|
"""The safety tag records the exact pre-restore version map, and
|
|
restoring it returns the database to its prior logical state."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
snapshot = await store.current_table_versions()
|
|
|
|
safety_tag = await store.restore_tag("release-1")
|
|
|
|
tags = await store.list_tags()
|
|
assert tags[safety_tag].complete is True
|
|
assert tags[safety_tag].tables == snapshot
|
|
|
|
await store.restore_tag(safety_tag)
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_missing_tag_makes_no_changes(temp_db_path):
|
|
async with Store(temp_db_path, create=True) as store:
|
|
await DocumentRepository(store).create(Document(content="First document"))
|
|
versions = await store.current_table_versions()
|
|
|
|
with pytest.raises(ValueError, match="does not exist"):
|
|
await store.restore_tag("nope")
|
|
|
|
assert await store.current_table_versions() == versions
|
|
assert await store.list_tags() == {}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_partial_tag_makes_no_changes(temp_db_path):
|
|
"""A partial tag can never be restored; the error lists every missing
|
|
table and no safety tag is created."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
version = await store.chunks_table.version()
|
|
await store.chunks_table.tags.create("stale", version)
|
|
versions = await store.current_table_versions()
|
|
|
|
with pytest.raises(ValueError) as exc_info:
|
|
await store.restore_tag("stale")
|
|
|
|
msg = str(exc_info.value)
|
|
for table_name in ("documents", "document_meta", "document_items", "settings"):
|
|
assert table_name in msg
|
|
|
|
assert await store.current_table_versions() == versions
|
|
assert set(await store.list_tags()) == {"stale"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_safety_tag_name_collision(temp_db_path, monkeypatch):
|
|
"""A colliding safety-tag name gets a numeric suffix."""
|
|
import haiku.rag.store.engine as engine_mod
|
|
|
|
class FixedDatetime:
|
|
@staticmethod
|
|
def now(tz=None):
|
|
from datetime import UTC, datetime
|
|
|
|
return datetime(2026, 7, 15, 14, 30, 12, tzinfo=UTC)
|
|
|
|
monkeypatch.setattr(engine_mod, "datetime", FixedDatetime)
|
|
|
|
async with Store(temp_db_path, create=True) as store:
|
|
await DocumentRepository(store).create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await store.create_tag("before-restore-20260715T143012Z")
|
|
|
|
await store.create_tag("before-restore-20260715T143012Z-2")
|
|
|
|
safety_tag = await store.restore_tag("release-1")
|
|
assert safety_tag == "before-restore-20260715T143012Z-3"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_safety_tag_failure_leaves_state_untouched(
|
|
temp_db_path, monkeypatch
|
|
):
|
|
"""If the safety tag cannot be created, restore never begins."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
versions = await store.current_table_versions()
|
|
|
|
async def failing_create(self, name: str, version: int) -> None:
|
|
raise RuntimeError("tag boom")
|
|
|
|
monkeypatch.setattr(AsyncTags, "create", failing_create)
|
|
|
|
with pytest.raises(RuntimeError) as exc_info:
|
|
await store.restore_tag("release-1")
|
|
|
|
msg = str(exc_info.value)
|
|
assert "did not begin" in msg
|
|
assert "No table was changed" in msg
|
|
assert "tag boom" in msg
|
|
assert exc_info.value.__cause__ is not None
|
|
|
|
monkeypatch.undo()
|
|
assert await store.current_table_versions() == versions
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
assert set(await store.list_tags()) == {"release-1"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_midway_failure_rolls_back(temp_db_path, monkeypatch):
|
|
"""A restore failure after some tables were restored rolls every table
|
|
back to the pre-restore snapshot; the error names the failed table and
|
|
the safety tag."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
real_restore = AsyncTable.restore
|
|
calls = {"n": 0}
|
|
|
|
async def flaky_restore(self, version=None):
|
|
calls["n"] += 1
|
|
if calls["n"] == 3:
|
|
raise RuntimeError("restore boom")
|
|
return await real_restore(self, version)
|
|
|
|
monkeypatch.setattr(AsyncTable, "restore", flaky_restore)
|
|
|
|
with pytest.raises(RuntimeError) as exc_info:
|
|
await store.restore_tag("release-1")
|
|
|
|
msg = str(exc_info.value)
|
|
assert RESTORE_TABLE_ORDER[2] in msg
|
|
assert "rolled back" in msg
|
|
assert "before-restore-" in msg
|
|
|
|
monkeypatch.undo()
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
assert any(t.startswith("before-restore-") for t in await store.list_tags())
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_rollback_failure_reports_inconsistency(
|
|
temp_db_path, monkeypatch
|
|
):
|
|
"""When rollback also fails, the error lists the failed tables, names
|
|
the safety tag, and states manual recovery is required."""
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
real_restore = AsyncTable.restore
|
|
calls = {"n": 0}
|
|
|
|
async def flaky_restore(self, version=None):
|
|
calls["n"] += 1
|
|
if calls["n"] >= 3:
|
|
raise RuntimeError("restore boom")
|
|
return await real_restore(self, version)
|
|
|
|
monkeypatch.setattr(AsyncTable, "restore", flaky_restore)
|
|
|
|
with pytest.raises(RuntimeError) as exc_info:
|
|
await store.restore_tag("release-1")
|
|
|
|
msg = str(exc_info.value)
|
|
assert "inconsistent" in msg
|
|
assert "manual recovery" in msg
|
|
assert "before-restore-" in msg
|
|
for table_name in RESTORE_TABLE_ORDER:
|
|
assert table_name in msg
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_cancellation_rolls_back(temp_db_path, monkeypatch):
|
|
"""Cancellation mid-restore must not bypass rollback: the tables return
|
|
to the pre-restore snapshot and the cancellation re-raises."""
|
|
import asyncio
|
|
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
real_restore = AsyncTable.restore
|
|
calls = {"n": 0}
|
|
|
|
async def cancelled_restore(self, version=None):
|
|
calls["n"] += 1
|
|
if calls["n"] == 3:
|
|
raise asyncio.CancelledError()
|
|
return await real_restore(self, version)
|
|
|
|
monkeypatch.setattr(AsyncTable, "restore", cancelled_restore)
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await store.restore_tag("release-1")
|
|
|
|
monkeypatch.undo()
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
assert any(t.startswith("before-restore-") for t in await store.list_tags())
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_cancellation_with_failed_rollback_reports(
|
|
temp_db_path, monkeypatch
|
|
):
|
|
"""If rollback after a cancellation also fails, the manual-recovery
|
|
error is raised instead of the bare cancellation."""
|
|
import asyncio
|
|
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
real_restore = AsyncTable.restore
|
|
calls = {"n": 0}
|
|
|
|
async def broken_restore(self, version=None):
|
|
calls["n"] += 1
|
|
if calls["n"] < 3:
|
|
return await real_restore(self, version)
|
|
if calls["n"] == 3:
|
|
raise asyncio.CancelledError()
|
|
raise RuntimeError("restore boom")
|
|
|
|
monkeypatch.setattr(AsyncTable, "restore", broken_restore)
|
|
|
|
with pytest.raises(RuntimeError) as exc_info:
|
|
await store.restore_tag("release-1")
|
|
|
|
msg = str(exc_info.value)
|
|
assert "cancel" in msg.lower()
|
|
assert "manual recovery" in msg
|
|
assert "before-restore-" in msg
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_read_only_raises(temp_db_path):
|
|
async with Store(temp_db_path, create=True) as store:
|
|
await store.create_tag("release-1")
|
|
|
|
async with Store(temp_db_path, read_only=True) as store:
|
|
with pytest.raises(ReadOnlyError):
|
|
await store.restore_tag("release-1")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_rejected_during_rebuild(temp_db_path):
|
|
async with Store(temp_db_path, create=True) as store:
|
|
await store.create_tag("release-1")
|
|
|
|
async with store._rebuild_lock:
|
|
with pytest.raises(ValueError, match="[Rr]ebuild in progress"):
|
|
await store.restore_tag("release-1")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_old_version_marker_requires_explicit_migration(temp_db_path):
|
|
"""Restore never migrates: restoring a tag whose settings carry an old
|
|
version marker completes, the next normal open hits the migration gate,
|
|
explicit migration works, and the safety tag remains usable after it."""
|
|
from haiku.rag.store.exceptions import MigrationRequiredError
|
|
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
current_version = await store.get_haiku_version()
|
|
await store.set_haiku_version("0.63.0")
|
|
await store.create_tag("old-marker")
|
|
await store.set_haiku_version(current_version)
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
async with Store(temp_db_path) as store:
|
|
safety_tag = await store.restore_tag("old-marker")
|
|
assert await store.get_haiku_version() == "0.63.0"
|
|
assert await _doc_contents(store) == {"First document"}
|
|
|
|
with pytest.raises(MigrationRequiredError):
|
|
async with Store(temp_db_path):
|
|
pass
|
|
|
|
async with Store(temp_db_path, skip_migration_check=True) as store:
|
|
await store.migrate()
|
|
|
|
async with Store(temp_db_path) as store:
|
|
assert await _doc_contents(store) == {"First document"}
|
|
await store.restore_tag(safety_tag)
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_restore_failure_rollback_survives_cancellation(
|
|
temp_db_path, monkeypatch
|
|
):
|
|
"""Cancelling restore while it rolls back a failed restore must not
|
|
interrupt the rollback: all tables return to the snapshot before the
|
|
cancellation is delivered."""
|
|
import asyncio
|
|
|
|
async with Store(temp_db_path, create=True) as store:
|
|
repo = DocumentRepository(store)
|
|
await repo.create(Document(content="First document"))
|
|
await store.create_tag("release-1")
|
|
await repo.create(Document(content="Second document"))
|
|
|
|
real_restore = AsyncTable.restore
|
|
calls = {"n": 0}
|
|
rollback_started = asyncio.Event()
|
|
release = asyncio.Event()
|
|
|
|
async def flaky_restore(self, version=None):
|
|
calls["n"] += 1
|
|
if calls["n"] == 3:
|
|
raise RuntimeError("restore boom")
|
|
if calls["n"] == 4:
|
|
rollback_started.set()
|
|
await release.wait()
|
|
return await real_restore(self, version)
|
|
|
|
monkeypatch.setattr(AsyncTable, "restore", flaky_restore)
|
|
|
|
task = asyncio.create_task(store.restore_tag("release-1"))
|
|
await rollback_started.wait()
|
|
task.cancel()
|
|
release.set()
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await task
|
|
|
|
monkeypatch.undo()
|
|
# 3 forward calls (2 ok, 1 failed) + all 5 rollback calls ran.
|
|
assert calls["n"] == 8
|
|
assert await _doc_contents(store) == {"First document", "Second document"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_wait_protected_returns_result_on_same_tick_cancellation():
|
|
"""A cancellation landing after the recovery task completed but before
|
|
the waiter resumed must not discard the recovery result."""
|
|
import asyncio
|
|
|
|
from haiku.rag.store.engine import _wait_protected
|
|
|
|
async def recovery() -> str:
|
|
return "done"
|
|
|
|
outer = asyncio.create_task(_wait_protected(recovery()))
|
|
# First pass: outer starts, spawns the recovery task, suspends on shield.
|
|
await asyncio.sleep(0)
|
|
# Second pass: the recovery task completes; outer is scheduled to resume.
|
|
await asyncio.sleep(0)
|
|
# Cancellation beats the resumption: delivered at the shield await even
|
|
# though the recovery already finished.
|
|
outer.cancel()
|
|
|
|
result, cancelled = await outer
|
|
assert result == "done"
|
|
assert cancelled is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_wait_protected_reraises_when_recovery_itself_is_cancelled():
|
|
"""If the recovery coroutine ends cancelled there is nothing to wait for,
|
|
so the cancellation propagates instead of looping forever."""
|
|
import asyncio
|
|
|
|
from haiku.rag.store.engine import _wait_protected
|
|
|
|
async def self_cancelling_recovery() -> str:
|
|
current = asyncio.current_task()
|
|
assert current is not None
|
|
current.cancel()
|
|
await asyncio.sleep(0)
|
|
return "unreachable"
|
|
|
|
outer = asyncio.create_task(_wait_protected(self_cancelling_recovery()))
|
|
|
|
# Bounded: without the re-raise the retry loop spins forever on an
|
|
# already-cancelled task, and the timeout turns that into a clean failure.
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await asyncio.wait_for(outer, timeout=5)
|