Add database-level tag primitives to Store
This commit is contained in:
parent
aa620aeb60
commit
ae603d9b4d
2 changed files with 229 additions and 29 deletions
|
|
@ -1,6 +1,8 @@
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import contextlib
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from importlib import metadata
|
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:
|
async def get_database_stats(db: lancedb.AsyncConnection) -> dict:
|
||||||
"""Collect stats for every haiku.rag table on the connection.
|
"""Collect stats for every haiku.rag table on the connection.
|
||||||
|
|
||||||
|
|
@ -788,15 +807,19 @@ class Store:
|
||||||
if hasattr(self, "db"):
|
if hasattr(self, "db"):
|
||||||
self.db.close()
|
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]:
|
async def current_table_versions(self) -> dict[str, int]:
|
||||||
"""Capture current versions of key tables for rollback using LanceDB's API."""
|
"""Capture current versions of key tables for rollback using LanceDB's API."""
|
||||||
return {
|
return {name: await table.version() for name, table in self._tables().items()}
|
||||||
"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(),
|
|
||||||
}
|
|
||||||
|
|
||||||
async def restore_table_versions(self, versions: dict[str, int]) -> bool:
|
async def restore_table_versions(self, versions: dict[str, int]) -> bool:
|
||||||
"""Restore tables to the provided versions using LanceDB's API.
|
"""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.
|
ReadOnlyError: If the store is in read-only mode.
|
||||||
"""
|
"""
|
||||||
self._assert_writable()
|
self._assert_writable()
|
||||||
await self.documents_table.restore(int(versions["documents"]))
|
for name, table in self._tables().items():
|
||||||
await self.document_meta_table.restore(int(versions["document_meta"]))
|
await table.restore(int(versions[name]))
|
||||||
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"]))
|
|
||||||
return True
|
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:
|
async def _checkout_tables_before(self, before: datetime) -> None:
|
||||||
"""Checkout all tables to their state at or before the given datetime.
|
"""Checkout all tables to their state at or before the given datetime.
|
||||||
|
|
||||||
|
|
@ -830,15 +920,7 @@ class Store:
|
||||||
# Already naive, assume local time
|
# Already naive, assume local time
|
||||||
before_local = before
|
before_local = before
|
||||||
|
|
||||||
tables = [
|
for table in self._tables().values():
|
||||||
("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:
|
|
||||||
versions = await table.list_versions()
|
versions = await table.list_versions()
|
||||||
# Find the latest version at or before the target datetime
|
# Find the latest version at or before the target datetime
|
||||||
# Versions are sorted by version number, not timestamp, so we need to check all
|
# Versions are sorted by version number, not timestamp, so we need to check all
|
||||||
|
|
@ -884,14 +966,7 @@ class Store:
|
||||||
Returns:
|
Returns:
|
||||||
List of version info dicts with "version" and "timestamp" keys
|
List of version info dicts with "version" and "timestamp" keys
|
||||||
"""
|
"""
|
||||||
table_map = {
|
table = self._tables().get(table_name)
|
||||||
"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)
|
|
||||||
if table is None:
|
if table is None:
|
||||||
raise ValueError(f"Unknown table: {table_name}")
|
raise ValueError(f"Unknown table: {table_name}")
|
||||||
|
|
||||||
|
|
|
||||||
125
tests/store/test_tags.py
Normal file
125
tests/store/test_tags.py
Normal file
|
|
@ -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
|
||||||
Loading…
Reference in a new issue