diff --git a/haiku_rag_slim/haiku/rag/store/engine.py b/haiku_rag_slim/haiku/rag/store/engine.py index 72c0dd40..bdc18b93 100644 --- a/haiku_rag_slim/haiku/rag/store/engine.py +++ b/haiku_rag_slim/haiku/rag/store/engine.py @@ -1,6 +1,8 @@ import asyncio +import contextlib import json import logging +from dataclasses import dataclass from datetime import datetime, timedelta from enum import Enum from importlib import metadata @@ -188,6 +190,23 @@ REQUIRED_TABLES: tuple[str, ...] = ( ) +@dataclass +class TagInfo: + """A database-level tag aggregated across all tables. + + A complete tag names the same tag on every table; a partial one (created + outside haiku.rag or left behind by a failure) lists the tables it is + missing from. + """ + + tables: dict[str, int] + missing_tables: list[str] + + @property + def complete(self) -> bool: + return not self.missing_tables + + async def get_database_stats(db: lancedb.AsyncConnection) -> dict: """Collect stats for every haiku.rag table on the connection. @@ -788,15 +807,19 @@ class Store: if hasattr(self, "db"): self.db.close() + def _tables(self) -> dict[str, lancedb.AsyncTable]: + """Map every haiku.rag table name to its open AsyncTable.""" + return { + "documents": self.documents_table, + "document_meta": self.document_meta_table, + "chunks": self.chunks_table, + "document_items": self.document_items_table, + "settings": self.settings_table, + } + async def current_table_versions(self) -> dict[str, int]: """Capture current versions of key tables for rollback using LanceDB's API.""" - return { - "documents": await self.documents_table.version(), - "document_meta": await self.document_meta_table.version(), - "chunks": await self.chunks_table.version(), - "document_items": await self.document_items_table.version(), - "settings": await self.settings_table.version(), - } + return {name: await table.version() for name, table in self._tables().items()} async def restore_table_versions(self, versions: dict[str, int]) -> bool: """Restore tables to the provided versions using LanceDB's API. @@ -805,13 +828,80 @@ class Store: ReadOnlyError: If the store is in read-only mode. """ self._assert_writable() - await self.documents_table.restore(int(versions["documents"])) - await self.document_meta_table.restore(int(versions["document_meta"])) - await self.chunks_table.restore(int(versions["chunks"])) - await self.document_items_table.restore(int(versions["document_items"])) - await self.settings_table.restore(int(versions["settings"])) + for name, table in self._tables().items(): + await table.restore(int(versions[name])) return True + async def create_tag(self, name: str) -> None: + """Tag the current version of every table with the given name. + + Raises: + ReadOnlyError: If the store is in read-only mode. + ValueError: If the tag already exists on any table. A partial tag + (present on some tables only) must be deleted before the name + can be reused. + """ + self._assert_writable() + tables = self._tables() + + existing = [ + table_name + for table_name, table in tables.items() + if name in await table.tags.list() + ] + if len(existing) == len(tables): + raise ValueError(f"Tag '{name}' already exists") + if existing: + raise ValueError( + f"Tag '{name}' already exists on some tables " + f"({', '.join(existing)}); delete it first with delete_tag" + ) + + versions = await self.current_table_versions() + created: list[str] = [] + try: + for table_name, table in tables.items(): + await table.tags.create(name, versions[table_name]) + created.append(table_name) + except Exception: + for table_name in created: + with contextlib.suppress(Exception): + await tables[table_name].tags.delete(name) + raise + + async def list_tags(self) -> dict[str, TagInfo]: + """Aggregate per-table tags into database-level tags. + + Returns: + Tag name mapped to a TagInfo with the tagged version per table + and the tables the tag is missing from (empty when complete). + """ + tables = self._tables() + tags: dict[str, TagInfo] = {} + for table_name, table in tables.items(): + for tag_name, tag in (await table.tags.list()).items(): + info = tags.setdefault(tag_name, TagInfo(tables={}, missing_tables=[])) + info.tables[table_name] = tag["version"] + for info in tags.values(): + info.missing_tables = [t for t in tables if t not in info.tables] + return tags + + async def delete_tag(self, name: str) -> None: + """Delete the tag from every table that has it. + + Raises: + ReadOnlyError: If the store is in read-only mode. + ValueError: If no table has the tag. + """ + self._assert_writable() + found = False + for table in self._tables().values(): + if name in await table.tags.list(): + await table.tags.delete(name) + found = True + if not found: + raise ValueError(f"Tag '{name}' does not exist") + async def _checkout_tables_before(self, before: datetime) -> None: """Checkout all tables to their state at or before the given datetime. @@ -830,15 +920,7 @@ class Store: # Already naive, assume local time before_local = before - tables = [ - ("documents", self.documents_table), - ("document_meta", self.document_meta_table), - ("chunks", self.chunks_table), - ("document_items", self.document_items_table), - ("settings", self.settings_table), - ] - - for table_name, table in tables: + for table in self._tables().values(): versions = await table.list_versions() # Find the latest version at or before the target datetime # Versions are sorted by version number, not timestamp, so we need to check all @@ -884,14 +966,7 @@ class Store: Returns: List of version info dicts with "version" and "timestamp" keys """ - table_map = { - "documents": self.documents_table, - "document_meta": self.document_meta_table, - "chunks": self.chunks_table, - "document_items": self.document_items_table, - "settings": self.settings_table, - } - table = table_map.get(table_name) + table = self._tables().get(table_name) if table is None: raise ValueError(f"Unknown table: {table_name}") diff --git a/tests/store/test_tags.py b/tests/store/test_tags.py new file mode 100644 index 00000000..04c30e91 --- /dev/null +++ b/tests/store/test_tags.py @@ -0,0 +1,125 @@ +import pytest +from lancedb.table import AsyncTags + +from haiku.rag.store import ReadOnlyError, Store +from haiku.rag.store.engine import REQUIRED_TABLES +from haiku.rag.store.models import Document +from haiku.rag.store.repositories.document import DocumentRepository + + +@pytest.mark.asyncio +async def test_create_and_list_tags(temp_db_path): + """create_tag tags every table at its current version; list_tags reports + the tag as complete with the exact versions.""" + async with Store(temp_db_path, create=True) as store: + repo = DocumentRepository(store) + await repo.create(Document(content="First document")) + + versions = await store.current_table_versions() + await store.create_tag("release-1") + + tags = await store.list_tags() + assert set(tags) == {"release-1"} + info = tags["release-1"] + assert info.complete is True + assert info.missing_tables == [] + assert info.tables == versions + + +@pytest.mark.asyncio +async def test_create_tag_rejects_existing(temp_db_path): + async with Store(temp_db_path, create=True) as store: + await store.create_tag("release-1") + + with pytest.raises(ValueError, match="already exists"): + await store.create_tag("release-1") + + tags = await store.list_tags() + assert tags["release-1"].complete is True + + +@pytest.mark.asyncio +async def test_create_tag_rejects_partial_existing(temp_db_path): + """A tag present on only some tables blocks creation before anything is + written; the error tells the user to delete it first.""" + async with Store(temp_db_path, create=True) as store: + version = await store.chunks_table.version() + await store.chunks_table.tags.create("stale", version) + + with pytest.raises(ValueError, match="delete"): + await store.create_tag("stale") + + tags = await store.list_tags() + assert tags["stale"].complete is False + assert set(tags["stale"].tables) == {"chunks"} + assert set(tags["stale"].missing_tables) == set(REQUIRED_TABLES) - {"chunks"} + + +@pytest.mark.asyncio +async def test_create_tag_rolls_back_own_tags_on_failure(temp_db_path, monkeypatch): + """A midway failure removes the tags this call created and leaves + pre-existing tags untouched.""" + async with Store(temp_db_path, create=True) as store: + await store.create_tag("keep") + + real_create = AsyncTags.create + calls = {"n": 0} + + async def flaky(self, name: str, version: int) -> None: + calls["n"] += 1 + if calls["n"] == 4: + raise RuntimeError("boom") + await real_create(self, name, version) + + monkeypatch.setattr(AsyncTags, "create", flaky) + + with pytest.raises(RuntimeError, match="boom"): + await store.create_tag("broken") + + monkeypatch.undo() + + tags = await store.list_tags() + assert "broken" not in tags + assert tags["keep"].complete is True + + +@pytest.mark.asyncio +async def test_delete_tag(temp_db_path): + async with Store(temp_db_path, create=True) as store: + await store.create_tag("release-1") + await store.delete_tag("release-1") + + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_delete_tag_heals_partial(temp_db_path): + async with Store(temp_db_path, create=True) as store: + version = await store.chunks_table.version() + await store.chunks_table.tags.create("stale", version) + + await store.delete_tag("stale") + + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_delete_tag_missing_raises(temp_db_path): + async with Store(temp_db_path, create=True) as store: + with pytest.raises(ValueError, match="does not exist"): + await store.delete_tag("nope") + + +@pytest.mark.asyncio +async def test_tag_writes_raise_when_read_only(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.create_tag("release-2") + with pytest.raises(ReadOnlyError): + await store.delete_tag("release-1") + + tags = await store.list_tags() + assert tags["release-1"].complete is True