diff --git a/CHANGELOG.md b/CHANGELOG.md index 534e16c9..689fa53b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,18 @@ # Changelog ## [Unreleased] +### Added + +- Database tags: `haiku-rag tag create/list/delete/restore`, tags shown in `history`. `tag restore` creates a `before-restore-*` safety tag before changing live state. Vacuum retains versions back to the oldest tag. + +### Changed + +- `lancedb` bumped to 0.34.0. + +### Removed + +- `--before` global flag and the `before` constructor arguments on `HaikuRAG`, `Store`, `HaikuRAGApp`, `ChatApp`/`run_chat`, and `InspectorApp`/`run_inspector`. There is no read-only replacement; create tags prospectively before important changes and use `tag restore` during a maintenance window. + ### Fixed - `docling-local` text conversion no longer misroutes markdown/HTML content whose first bytes collide with a binary magic signature (e.g. `BM`, `ID3`) to an image or audio backend. diff --git a/README.md b/README.md index 6e57596d..d04300a5 100644 --- a/README.md +++ b/README.md @@ -23,7 +23,7 @@ Agentic RAG built on [LanceDB](https://lancedb.com/), [Pydantic AI](https://ai.p - **MCP server** — Expose as tools for AI assistants (Claude Desktop, etc.) - **Visual grounding** — View chunks highlighted on original page images - **Production ingester** — Long-lived `haiku-ingester` service with persistent SQLite queue, async worker pool with retries and a dead-letter queue, FS / HTTP / S3 / WebDAV source adapters, FastAPI control plane, and a browser dashboard for operators. See [docs/ingester.md](docs/ingester.md). -- **Time travel** — Query the database at any historical point with `--before` +- **Tags** — Name database states with `haiku-rag tag` and roll back to them - **Inspector** — TUI for browsing documents, chunks, and search results ## Installation diff --git a/docs/chat.md b/docs/chat.md index 966b6135..c0cdbc99 100644 --- a/docs/chat.md +++ b/docs/chat.md @@ -76,4 +76,4 @@ For everyday Q&A, the rag skill alone is faster and cheaper. Attaching both lets Run "Filter documents" from the command palette to restrict searches to a subset. The filter applies to every search the agent runs for the rest of the session. -Chat also honors the global `--read-only` and `--before` flags. See the [CLI reference](cli.md) for details. +Chat also honors the global `--read-only` flag. See the [CLI reference](cli.md) for details. diff --git a/docs/cli.md b/docs/cli.md index 1fe6d14d..ad7252a4 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -7,7 +7,6 @@ The `haiku-rag` CLI provides complete document management functionality. - `--config` - Specify custom configuration file - `--read-only` - Open database in read-only mode (blocks writes, skips upgrades) - - `--before` - Query database as it existed before a datetime (implies `--read-only`) - `--version` / `-v` - Show version and exit Per-command options: @@ -20,7 +19,6 @@ The `haiku-rag` CLI provides complete document management functionality. haiku-rag --config /path/to/config.yaml list haiku-rag --config /path/to/config.yaml list --db /path/to/custom.db haiku-rag --read-only search "query" - haiku-rag --before "2025-01-15" search "query" haiku-rag add -h ``` @@ -542,35 +540,53 @@ haiku-skills chat --use-entrypoints --skill medic └── haiku.rag.yaml # Optional config ``` -## Time Travel +## Tags -LanceDB maintains version history for tables, enabling you to query the database as it existed at a previous point in time. This is useful for: - -- **Debugging**: Investigate data before a problematic change -- **Auditing**: Verify what knowledge was available when a support ticket was filed - -### Query Historical State - -Use `--before` to query the database as it existed before a specific datetime: +A tag names the current database state. It is a logical snapshot composed of one LanceDB tag on each of the five tables, created from a single version snapshot. ```bash -# Query documents as of January 15, 2025 -haiku-rag --before "2025-01-15" list +# Tag the current state, e.g. at deploy time or after an ingestion run +haiku-rag tag create release-1 -# Search historical state -haiku-rag --before "2025-01-15T14:30:00" search "machine learning" +# List tags with the versions they point to +haiku-rag tag list -# Ask questions against historical data -haiku-rag --before "2025-01-15" ask "What documents existed?" +# Delete a tag, releasing its versions for cleanup +haiku-rag tag delete release-1 ``` -Supported datetime formats: +A tag present on every table is complete. A tag missing from some tables (created outside haiku.rag, or left behind by a failure) is partial. `tag list` marks partial tags. Partial tags can be listed and deleted but never restored. -- ISO 8601: `2025-01-15T14:30:00`, `2025-01-15T14:30:00Z`, `2025-01-15T14:30:00+00:00` -- Date only: `2025-01-15` (interpreted as start of day) +Create tags with other writers stopped. Tag creation coordinates writers within one process only; a writer in another process can commit between the per-table snapshot reads, and the tag then captures a mixed state. -!!! note - Time travel mode automatically enables read-only mode. You cannot modify the database while viewing historical state. +Tagged versions survive `vacuum`. Vacuum retains the oldest tagged version and every newer version; versions older than the oldest tag remain eligible for cleanup. Delete tags you no longer need so cleanup can advance. + +### Restore + +`tag restore` brings the database back to a tagged state: + +```bash +haiku-rag tag restore release-1 +``` + +Restore changes the live state. It is not a read-only view: each table gets a new latest version equal to the tagged one, and reads and writes continue from there. Versions written after the tag remain in history until vacuum removes them. + +Before changing anything, restore creates a complete safety tag (`before-restore-`) for the current state and reports it, so you always have a named path back: + +```bash +haiku-rag tag create release-1 --db /path/to/db.lancedb +# Stop all writers before either restore. +haiku-rag tag restore release-1 --db /path/to/db.lancedb --yes +haiku-rag tag list --db /path/to/db.lancedb +haiku-rag tag restore before-restore-YYYYMMDDTHHMMSSZ --db /path/to/db.lancedb --yes +``` + +Restore is a maintenance operation: + +- Stop all ingestion and other writers before restoring and keep them stopped until it finishes. +- The operation is coordinated but not transactionally atomic across tables. On failure it attempts to roll back to the pre-restore state and reports whether the rollback succeeded. +- `--yes` only skips the confirmation prompt. It provides no locking and no concurrent-writer protection. +- Restore never migrates. Restoring a tag from an older haiku.rag version completes normally, and the next open reports the required migration. Run `haiku-rag migrate` explicitly. ### Version History @@ -587,20 +603,18 @@ haiku-rag history --table documents haiku-rag history --limit 10 ``` -Output shows version numbers and timestamps, sorted newest first: +Output shows version numbers and timestamps, sorted newest first, with tags marked: ``` Version History documents - v5: 2025-01-15 14:30:00 + v5: 2025-01-15 14:30:00 <- release-1 v4: 2025-01-14 10:00:00 v3: 2025-01-13 09:15:00 chunks - v8: 2025-01-15 14:30:00 + v8: 2025-01-15 14:30:00 <- release-1 v7: 2025-01-14 10:00:00 ... ``` - -Use the timestamps from `history` to construct `--before` queries. diff --git a/docs/python.md b/docs/python.md index c4aa1507..0e377024 100644 --- a/docs/python.md +++ b/docs/python.md @@ -434,6 +434,32 @@ await client.vacuum() This compacts tables and removes historical versions to keep disk usage in check. It’s safe to run anytime, for example after bulk imports or periodically in long‑running apps. +### Tags + +Tag the current database state and restore it later, for example after an ingestion run. A tag covers all five tables and is created from a single version snapshot. Create tags with other writers stopped: the snapshot is coordinated within one process only, and a writer in another process can commit between the per-table reads. + +```python +await client.store.create_tag("release-1") + +tags = await client.store.list_tags() +for name, info in tags.items(): + print(name, info.tables, info.complete) +``` + +`restore_tag` brings the live database back to a tagged state. It creates a complete safety tag for the current state before changing any table and returns its name: + +```python +safety_tag = await client.store.restore_tag("release-1") +``` + +Restore is a maintenance operation: stop all other writers first. A tag present on only some tables is partial; `list_tags` reports it via `missing_tables`, and partial tags can be deleted but never restored. + +Delete tags you no longer need. Vacuum retains the oldest tagged version and everything newer: + +```python +await client.store.delete_tag("release-1") +``` + ### Rebuilding the Database ```python diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index b3eba176..c7f95427 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -1,10 +1,10 @@ import logging -from datetime import datetime from pathlib import Path from typing import TYPE_CHECKING from rich.console import Console from rich.markdown import Markdown +from rich.markup import escape from rich.progress import ( BarColumn, DownloadColumn, @@ -22,6 +22,7 @@ from haiku.rag.store.models.chunk import SearchType from haiku.rag.store.models.document import Document if TYPE_CHECKING: + from haiku.rag.store.engine import Store from haiku.rag.store.models import SearchResult from haiku.rag.utils import format_bytes, format_citations_rich @@ -34,12 +35,10 @@ class HaikuRAGApp: # pragma: no cover db_path: Path, config: AppConfig = Config, read_only: bool = False, - before: datetime | None = None, ): self.db_path = db_path self.config = config self.read_only = read_only - self.before = before self.console = Console() from haiku.rag.store.engine import ConnectionMode @@ -67,11 +66,6 @@ class HaikuRAGApp: # pragma: no cover from haiku.rag.store.engine import gather_database_info - if self.before is not None: - self.console.print( - "[yellow]Note: --before is not supported by info; showing current state.[/yellow]" - ) - # Basic: show path/URI self.console.print("[bold]haiku.rag database info[/bold]") self.console.print( @@ -288,7 +282,6 @@ class HaikuRAGApp: # pragma: no cover skip_validation=True, read_only=True, skip_migration_check=True, - before=self.before, ) as store: tables = [ "documents", @@ -307,6 +300,14 @@ class HaikuRAGApp: # pragma: no cover self.console.print("[bold]Version History[/bold]") + try: + tags = await store.list_tags() + except Exception as exc: + tags = {} + self.console.print( + f"[yellow]Tag annotations unavailable: {escape(str(exc))}[/yellow]" + ) + for table_name in tables: versions = await store.list_table_versions(table_name) @@ -316,6 +317,12 @@ class HaikuRAGApp: # pragma: no cover if limit: versions = versions[:limit] + version_tags: dict[int, list[str]] = {} + for tag_name, info in tags.items(): + tagged_version = info.tables.get(table_name) + if tagged_version is not None: + version_tags.setdefault(tagged_version, []).append(tag_name) + self.console.print(f"\n[bold cyan]{table_name}[/bold cyan]") if not versions: @@ -325,16 +332,105 @@ class HaikuRAGApp: # pragma: no cover for v in versions: version_num = v["version"] timestamp = v["timestamp"] + suffix = "" + if version_num in version_tags: + names = ", ".join( + escape(n) for n in sorted(version_tags[version_num]) + ) + suffix = f" [magenta]<- {names}[/magenta]" self.console.print( - f" [repr.attrib_name]v{version_num}[/repr.attrib_name]: {timestamp}" + f" [repr.attrib_name]v{version_num}[/repr.attrib_name]: {timestamp}{suffix}" ) + def _tag_write_store(self) -> "Store": + """Writable store for tag create/delete with normal validation and + migration checks. + + A coordinated tag is only reliable when the database schema is + current, and a writable open of a legacy database would create + missing tables as a side effect. + """ + from haiku.rag.store.engine import Store + + return Store(self.db_path, config=self.config, read_only=self.read_only) + + def _tag_read_store(self) -> "Store": + """Read-only store for tag inspection; works on old or drifted DBs.""" + from haiku.rag.store.engine import Store + + return Store( + self.db_path, + config=self.config, + skip_validation=True, + skip_migration_check=True, + read_only=True, + ) + + async def create_tag(self, name: str): + """Tag the current version of every table.""" + if self._is_local and not self.db_path.exists(): + raise ValueError(f"Database path does not exist: {self.db_path}") + async with self._tag_write_store() as store: + await store.create_tag(name) + self.console.print(f"[green]Created tag '{escape(name)}'[/green]") + + async def list_tags(self): + """List database tags, flagging partial ones.""" + if self._is_local and not self.db_path.exists(): + raise ValueError(f"Database path does not exist: {self.db_path}") + async with self._tag_read_store() as store: + tags = await store.list_tags() + + if not tags: + self.console.print("No tags") + return + + self.console.print("[bold]Tags[/bold]") + for name in sorted(tags): + info = tags[name] + versions = " ".join(f"{t}=v{v}" for t, v in info.tables.items()) + line = f" [repr.attrib_name]{escape(name)}[/repr.attrib_name]: {versions}" + if not info.complete: + missing = ", ".join(info.missing_tables) + line += f" [yellow](partial - missing: {missing})[/yellow]" + self.console.print(line) + + async def delete_tag(self, name: str): + """Delete a tag from every table that has it.""" + if self._is_local and not self.db_path.exists(): + raise ValueError(f"Database path does not exist: {self.db_path}") + async with self._tag_write_store() as store: + await store.delete_tag(name) + self.console.print(f"[green]Deleted tag '{escape(name)}'[/green]") + + async def restore_tag(self, name: str): + """Restore the database to a tagged state and report the outcome. + + The Store context exits before anything is printed; no high-level + database access happens after the restore. + + Raises: + ValueError: If the database path does not exist. + """ + if self._is_local and not self.db_path.exists(): + raise ValueError(f"Database path does not exist: {self.db_path}") + async with self._tag_write_store() as store: + safety_tag = await store.restore_tag(name) + self.console.print(f"[green]Restored database to tag '{escape(name)}'.[/green]") + self.console.print( + f"The previous state is preserved as '{escape(safety_tag)}'." + ) + self.console.print( + "The restored state is now live. Later historical versions remain " + "until eligible for vacuum. Run [cyan]haiku-rag migrate[/cyan] if " + "migration is required." + ) + async def list_documents(self, filter: str | None = None): async with HaikuRAG( db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: documents = await self.client.list_documents(filter=filter) for doc in documents: @@ -347,7 +443,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=self.read_only, - before=self.before, ) as self.client: doc = await self.client.create_document( text, title=title, metadata=metadata @@ -364,7 +459,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=self.read_only, - before=self.before, ) as self.client: result = await self.client.create_document_from_source( source, title=title, metadata=metadata @@ -386,7 +480,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: doc = await self.client.get_document_by_id(doc_id) if doc is None: @@ -399,7 +492,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=self.read_only, - before=self.before, ) as self.client: deleted = await self.client.delete_document(doc_id) if deleted: @@ -443,7 +535,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: results = await self.client.search( search_input, @@ -465,7 +556,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: chunk = await self.client.get_chunk_by_id(chunk_id) if not chunk: @@ -509,7 +599,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: answer, citations = await self.client.ask(question, filter=filter) @@ -537,7 +626,6 @@ class HaikuRAGApp: # pragma: no cover db_path=self.db_path, config=self.config, read_only=True, - before=self.before, ) as self.client: self.console.print(f"[bold blue]Question:[/bold blue] {question}") self.console.print() @@ -561,7 +649,6 @@ class HaikuRAGApp: # pragma: no cover config=self.config, skip_validation=True, read_only=self.read_only, - before=self.before, ) as client: if mode == RebuildMode.SET_EMBEDDER: async for _ in client.rebuild_database(mode=mode): @@ -605,7 +692,6 @@ class HaikuRAGApp: # pragma: no cover config=self.config, skip_validation=True, read_only=self.read_only, - before=self.before, ) as client: await client.vacuum() self.console.print("[bold green]Vacuum completed successfully.[/bold green]") @@ -634,7 +720,6 @@ class HaikuRAGApp: # pragma: no cover config=self.config, skip_validation=True, read_only=self.read_only, - before=self.before, ) as client: row_count = await client.store.chunks_table.count_rows() self.console.print(f"Chunks in database: {row_count}") @@ -802,7 +887,6 @@ class HaikuRAGApp: # pragma: no cover self.db_path, config=self.config, read_only=self.read_only, - before=self.before, ): server = create_mcp_server( self.db_path, config=self.config, read_only=self.read_only diff --git a/haiku_rag_slim/haiku/rag/chat/__init__.py b/haiku_rag_slim/haiku/rag/chat/__init__.py index 00674990..6c94784b 100644 --- a/haiku_rag_slim/haiku/rag/chat/__init__.py +++ b/haiku_rag_slim/haiku/rag/chat/__init__.py @@ -1,11 +1,9 @@ -from datetime import datetime from pathlib import Path def run_chat( db_path: Path | None = None, read_only: bool = False, - before: datetime | None = None, model: str | None = None, skills: list[str] | None = None, ) -> None: @@ -14,7 +12,6 @@ def run_chat( Args: db_path: Path to the LanceDB database. If None, uses default from config. read_only: Whether to open the database in read-only mode. - before: Query database as it existed before this datetime. model: Model to use for the chat. skills: Skills to enable ("rag", "analysis"). Defaults to ["rag"]. """ @@ -55,7 +52,6 @@ def run_chat( db_path, skills=skill_list, read_only=read_only, - before=before, model=model or get_model(config.qa.model, config), ) app.run() diff --git a/haiku_rag_slim/haiku/rag/chat/app.py b/haiku_rag_slim/haiku/rag/chat/app.py index 70439286..c3464ebb 100644 --- a/haiku_rag_slim/haiku/rag/chat/app.py +++ b/haiku_rag_slim/haiku/rag/chat/app.py @@ -2,7 +2,6 @@ import asyncio import json import uuid from collections.abc import Iterable -from datetime import datetime from pathlib import Path from typing import TYPE_CHECKING, Any @@ -85,14 +84,12 @@ class ChatApp(App): db_path: Path, skills: list[Skill], read_only: bool = False, - before: datetime | None = None, model: str | None = None, ) -> None: super().__init__() self.db_path = db_path self._skills = skills self.read_only = read_only - self.before = before self._model = model self.client: HaikuRAG | None = None self.config = get_config() @@ -150,7 +147,6 @@ class ChatApp(App): db_path=self.db_path, config=self.config, read_only=self.read_only, - before=self.before, ) await self.client.__aenter__() diff --git a/haiku_rag_slim/haiku/rag/cli.py b/haiku_rag_slim/haiku/rag/cli.py index a14e5bfe..20c7318a 100644 --- a/haiku_rag_slim/haiku/rag/cli.py +++ b/haiku_rag_slim/haiku/rag/cli.py @@ -2,7 +2,6 @@ import asyncio import json import sys import warnings -from datetime import datetime from importlib.metadata import version from pathlib import Path from typing import Any @@ -47,7 +46,6 @@ def cli(): # Module-level flags set by callback _read_only: bool = False -_before: datetime | None = None def create_app(db: Path | None = None) -> HaikuRAGApp: # pragma: no cover @@ -61,9 +59,7 @@ def create_app(db: Path | None = None) -> HaikuRAGApp: # pragma: no cover """ config = get_config() db_path = db if db else config.storage.data_dir / "haiku.rag.lancedb" - return HaikuRAGApp( - db_path=db_path, config=config, read_only=_read_only, before=_before - ) + return HaikuRAGApp(db_path=db_path, config=config, read_only=_read_only) async def check_version(): # pragma: no cover @@ -102,28 +98,10 @@ def main( "--read-only", help="Open database in read-only mode", ), - before: str | None = typer.Option( - None, - "--before", - help="Query database as it existed before this datetime (implies --read-only). " - "Accepts ISO 8601 format (e.g., 2025-01-15T14:30:00) or date (e.g., 2025-01-15)", - ), ): """haiku.rag CLI - Vector database RAG system""" - global _read_only, _before + global _read_only _read_only = read_only - - # Parse and store before datetime - if before is not None: # pragma: no cover - from haiku.rag.utils import parse_datetime, to_utc - - try: - _before = to_utc(parse_datetime(before)) - except ValueError as e: - typer.echo(f"Error: {e}") - raise typer.Exit(1) - else: - _before = None # Load config from --config, local folder, or default directory config_path = find_config_file(cli_path=config) if config_path: @@ -635,6 +613,98 @@ def history( # pragma: no cover asyncio.run(app.history(table=table, limit=limit)) +tag_cli = typer.Typer( + help="Manage database tags (named versions across all tables)", + no_args_is_help=True, +) +_cli.add_typer(tag_cli, name="tag") + + +@tag_cli.command("create", help="Tag the current database state") +def tag_create( # pragma: no cover + name: str = typer.Argument(help="Name of the tag to create"), + db: Path | None = typer.Option( + None, + "--db", + help="Path to the LanceDB database file", + ), +): + app = create_app(db) + try: + asyncio.run(app.create_tag(name)) + except (ValueError, RuntimeError) as e: + typer.echo(f"Error: {e}", err=True) + raise typer.Exit(1) + + +@tag_cli.command("list", help="List database tags") +def tag_list( # pragma: no cover + db: Path | None = typer.Option( + None, + "--db", + help="Path to the LanceDB database file", + ), +): + app = create_app(db) + try: + asyncio.run(app.list_tags()) + except (ValueError, RuntimeError) as e: + typer.echo(f"Error: {e}", err=True) + raise typer.Exit(1) + + +@tag_cli.command("delete", help="Delete a tag") +def tag_delete( # pragma: no cover + name: str = typer.Argument(help="Name of the tag to delete"), + db: Path | None = typer.Option( + None, + "--db", + help="Path to the LanceDB database file", + ), +): + app = create_app(db) + try: + asyncio.run(app.delete_tag(name)) + except (ValueError, RuntimeError) as e: + typer.echo(f"Error: {e}", err=True) + raise typer.Exit(1) + + +@tag_cli.command("restore", help="Restore the database to a tagged state") +def tag_restore( # pragma: no cover + name: str = typer.Argument(help="Name of the tag to restore"), + yes: bool = typer.Option( + False, + "--yes", + help="Skip the confirmation prompt. Provides no locking or " + "concurrent-writer protection.", + ), + db: Path | None = typer.Option( + None, + "--db", + help="Path to the LanceDB database file", + ), +): + app = create_app(db) + if app._is_local and not app.db_path.exists(): + typer.echo(f"Error: Database path does not exist: {app.db_path}", err=True) + raise typer.Exit(1) + if not yes: + typer.echo(f"Database: {app.db_path}") + typer.echo(f"Tag: {name}") + typer.echo("This changes the live database state across all tables.") + typer.echo("Stop all ingestion and other writers before continuing.") + typer.echo("The operation is coordinated but not transactionally atomic.") + typer.echo("A safety tag will preserve the current state.") + if not typer.confirm("Continue?", default=False): + raise typer.Exit(1) + try: + asyncio.run(app.restore_tag(name)) + except (ValueError, RuntimeError) as e: + typer.echo(f"Error: {e}", err=True) + raise typer.Exit(1) + + @_cli.command("download-models", help="Download Docling and Ollama models per config") def download_models_cmd(): # pragma: no cover app = HaikuRAGApp(db_path=Path(), config=get_config(), read_only=True) @@ -661,7 +731,7 @@ def inspect( # pragma: no cover raise typer.Exit(1) from e db_path = db if db else get_config().storage.data_dir / "haiku.rag.lancedb" - run_inspector(db_path, read_only=True, before=_before) + run_inspector(db_path, read_only=True) @_cli.command("chat", help="Launch interactive chat TUI for conversational RAG") @@ -692,7 +762,6 @@ def chat( # pragma: no cover run_chat( db_path, read_only=True, - before=_before, model=model, skills=skills, ) diff --git a/haiku_rag_slim/haiku/rag/client/__init__.py b/haiku_rag_slim/haiku/rag/client/__init__.py index c817180f..bff6f32b 100644 --- a/haiku_rag_slim/haiku/rag/client/__init__.py +++ b/haiku_rag_slim/haiku/rag/client/__init__.py @@ -5,7 +5,6 @@ import logging import mimetypes import tempfile from collections.abc import AsyncGenerator, Sequence -from datetime import datetime from enum import Enum from functools import cached_property from pathlib import Path @@ -72,7 +71,6 @@ class HaikuRAG: skip_validation: bool = False, create: bool = False, read_only: bool = False, - before: datetime | None = None, ): """Initialize the RAG client with a database path. @@ -82,8 +80,6 @@ class HaikuRAG: skip_validation: Whether to skip configuration validation on database load. create: Whether to create the database if it doesn't exist. read_only: Whether to open the database in read-only mode. - before: Query the database as it existed at this datetime. - Implies read_only=True. """ self._config = config if db_path is None: @@ -93,7 +89,6 @@ class HaikuRAG: self._skip_validation = skip_validation self._create = create self._read_only = read_only - self._before = before self._vacuum_tasks: set[asyncio.Task] = set() self._last_vacuum_at: float | None = None self._vacuum_dirty = False @@ -125,7 +120,6 @@ class HaikuRAG: skip_validation=self._skip_validation, create=self._create, read_only=self._read_only, - before=self._before, ) # If _initialize fails mid-way (e.g. migration check raises after # connect), close the store so we don't leak the LanceDB connection — diff --git a/haiku_rag_slim/haiku/rag/client/documents.py b/haiku_rag_slim/haiku/rag/client/documents.py index f49effeb..41e25599 100644 --- a/haiku_rag_slim/haiku/rag/client/documents.py +++ b/haiku_rag_slim/haiku/rag/client/documents.py @@ -395,7 +395,8 @@ async def _refresh_doc_metadata( updated = True if updated: - result = await client.document_repository.update_meta(doc) + async with client.store._write_lock: + result = await client.document_repository.update_meta(doc) # Reclaim the document_meta churn from rolling source_revision sweeps. # The vacuum is debounced, and document_meta is tiny, so this is cheap. if client._config.storage.auto_vacuum: @@ -864,7 +865,8 @@ async def update_document( existing_doc.uri = uri if content is None and chunks is None and docling_document is None: - updated = await client.document_repository.update_meta(existing_doc) + async with client.store._write_lock: + updated = await client.document_repository.update_meta(existing_doc) if client._config.storage.auto_vacuum: client._schedule_vacuum() return updated diff --git a/haiku_rag_slim/haiku/rag/client/rebuild.py b/haiku_rag_slim/haiku/rag/client/rebuild.py index 38be804e..fb7f5b5a 100644 --- a/haiku_rag_slim/haiku/rag/client/rebuild.py +++ b/haiku_rag_slim/haiku/rag/client/rebuild.py @@ -71,15 +71,30 @@ async def rebuild_database( """Rebuild the database with the specified mode. Yields the ID of each document as it is processed. + + Holds the store's rebuild lock for the whole run so tag operations fail + fast instead of snapshotting a half-rebuilt database. The lock is held + across yields; an abandoned generator releases it when closed or + garbage-collected. """ from haiku.rag.client import RebuildMode if mode is None: mode = RebuildMode.FULL - if mode == RebuildMode.SET_EMBEDDER: - await _set_embedder(client) - return + async with client.store._rebuild_lock: + if mode == RebuildMode.SET_EMBEDDER: + await _set_embedder(client) + return + + async for doc_id in _rebuild_locked(client, mode): + yield doc_id + + +async def _rebuild_locked( + client: "HaikuRAG", mode: "RebuildMode" +) -> AsyncGenerator[str, None]: + from haiku.rag.client import RebuildMode # Resolve any leftover staging/marker tables from a previously # interrupted rebuild. Returns True only when phase 1 was already diff --git a/haiku_rag_slim/haiku/rag/inspector/app.py b/haiku_rag_slim/haiku/rag/inspector/app.py index b744ffe9..369e4f45 100644 --- a/haiku_rag_slim/haiku/rag/inspector/app.py +++ b/haiku_rag_slim/haiku/rag/inspector/app.py @@ -1,4 +1,3 @@ -from datetime import datetime from pathlib import Path from typing import TYPE_CHECKING @@ -67,13 +66,10 @@ class InspectorApp(App): Binding("c", "show_context", "Context", show=True), ] - def __init__( - self, db_path: Path, read_only: bool = False, before: datetime | None = None - ): + def __init__(self, db_path: Path, read_only: bool = False): super().__init__() self.db_path = db_path self.read_only = read_only - self.before = before self.client: HaikuRAG | None = None def compose(self) -> "ComposeResult": @@ -91,7 +87,6 @@ class InspectorApp(App): db_path=self.db_path, config=config, read_only=self.read_only, - before=self.before, ) await self.client.__aenter__() @@ -234,18 +229,16 @@ class InspectorApp(App): def run_inspector( db_path: Path | None = None, read_only: bool = False, - before: datetime | None = None, ) -> None: """Run the inspector TUI. Args: db_path: Path to the LanceDB database. If None, uses default from config. read_only: Whether to open the database in read-only mode. - before: Query database as it existed before this datetime. """ config = get_config() if db_path is None: db_path = config.storage.data_dir / "haiku.rag.lancedb" - app = InspectorApp(db_path, read_only=read_only, before=before) + app = InspectorApp(db_path, read_only=read_only) app.run() diff --git a/haiku_rag_slim/haiku/rag/store/engine.py b/haiku_rag_slim/haiku/rag/store/engine.py index 72c0dd40..1c5cd926 100644 --- a/haiku_rag_slim/haiku/rag/store/engine.py +++ b/haiku_rag_slim/haiku/rag/store/engine.py @@ -1,7 +1,9 @@ import asyncio import json import logging -from datetime import datetime, timedelta +from collections.abc import Coroutine +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta from enum import Enum from importlib import metadata from pathlib import Path @@ -187,6 +189,68 @@ REQUIRED_TABLES: tuple[str, ...] = ( "settings", ) +# Keeps the vacuum cleanup cutoff safely older than the oldest tagged +# version; guards against timestamp precision at the boundary. +TAG_RETENTION_MARGIN = timedelta(seconds=1) + +# Restore order for multi-table restore and its rollback. documents restores +# last: writes land in it last on the ingest path, making it the closest +# available database commit point. +RESTORE_TABLE_ORDER: tuple[str, ...] = tuple( + name for name in REQUIRED_TABLES if name != "documents" +) + ("documents",) + + +async def _wait_protected[T](coro: Coroutine[Any, Any, T]) -> tuple[T, bool]: + """Await a recovery coroutine that a cancellation cannot interrupt. + + Runs the coroutine as a task and keeps waiting for it even if this + coroutine is cancelled, so a Ctrl-C cannot leave recovery half applied. + Returns the result and whether a cancellation was absorbed; the caller + must re-deliver an absorbed cancellation. + """ + task = asyncio.ensure_future(coro) + cancelled = False + while True: + try: + return await asyncio.shield(task), cancelled + except asyncio.CancelledError: + if task.cancelled(): + # The recovery coroutine itself ended cancelled; there is + # nothing left to wait for. A task that completed (even in + # the same tick as the cancellation) still returns its + # result on the next pass. + raise + cancelled = True + + +def _safety_tag_name(existing: set[str]) -> str: + """Collision-resistant name for the pre-restore safety tag.""" + base = f"before-restore-{datetime.now(UTC).strftime('%Y%m%dT%H%M%SZ')}" + if base not in existing: + return base + n = 2 + while f"{base}-{n}" in existing: + n += 1 + return f"{base}-{n}" + + +@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. @@ -347,19 +411,19 @@ class Store: skip_validation: bool = False, create: bool = False, read_only: bool = False, - before: datetime | None = None, skip_migration_check: bool = False, ): self.db_path: Path = db_path self._config = config - self._before = before - # Time-travel mode is always read-only - self._read_only = read_only or (before is not None) + self._read_only = read_only self._create = create self._skip_validation = skip_validation self._skip_migration_check = skip_migration_check self._vacuum_lock = asyncio.Lock() self._write_lock = asyncio.Lock() + # Held by rebuild_database for its whole run; tag operations check it + # and fail fast instead of snapshotting a half-rebuilt database. + self._rebuild_lock = asyncio.Lock() self._is_new_db = False # Check if database exists (for local filesystem only) @@ -409,10 +473,6 @@ class Store: # pending, before creating any newly-introduced table. await self._init_tables(is_new_db) - # Checkout tables to historical state if before is specified - if self._before is not None: - await self._checkout_tables_before(self._before) - # Set version for new databases. if is_new_db and not self._read_only: await self._set_initial_version() @@ -474,6 +534,13 @@ class Store: if self._read_only: raise ReadOnlyError("Cannot modify database in read-only mode") + def _assert_not_rebuilding(self) -> None: + """Raise if a rebuild is in progress in this process.""" + if self._rebuild_lock.locked(): + raise ValueError( + "Rebuild in progress; tag operations are unavailable until it completes" + ) + async def vacuum(self, retention_seconds: int | None = None) -> None: """Optimize and clean up old versions across all tables to reduce disk usage. @@ -487,6 +554,8 @@ class Store: Raises: ReadOnlyError: If the store is in read-only mode. + RuntimeError: On lance errors during optimize; only OSError + (resource pressure) skips the pass. """ self._assert_writable() @@ -497,25 +566,54 @@ class Store: if self._vacuum_lock.locked(): return - async with self._vacuum_lock: + async with self._vacuum_lock, self._write_lock: try: # Evaluate config at runtime to allow dynamic changes if retention_seconds is None: retention_seconds = self._config.storage.vacuum_retention_seconds # Perform maintenance per table using optimize() with configurable retention retention = timedelta(seconds=retention_seconds) - for table in [ - self.documents_table, - self.document_meta_table, - self.chunks_table, - self.document_items_table, - self.settings_table, - ]: - await table.optimize(cleanup_older_than=retention) - except (RuntimeError, OSError) as e: - # Handle resource errors gracefully + for table in self._tables().values(): + await table.optimize( + cleanup_older_than=await self._tag_safe_retention( + table, retention + ) + ) + except OSError as e: + # Resource errors (e.g. disk pressure) skip the pass; lance + # errors surface as RuntimeError and must not be swallowed — + # a silently skipped cleanup hides tag-interaction bugs. logger.debug(f"Vacuum skipped due to resource constraints: {e}") + async def _tag_safe_retention( + self, table: lancedb.AsyncTable, retention: timedelta + ) -> timedelta: + """Grow the retention so the cleanup cutoff stays older than the + table's oldest tagged version. + + Lance hard-errors when a tagged version falls inside the cleanup + window and the Python API exposes no way to skip tagged versions, so + the oldest tagged version and everything newer are retained; versions + older than the oldest tag remain eligible for cleanup. + """ + tags = await table.tags.list() + if not tags: + return retention + + timestamps = {v["version"]: v["timestamp"] for v in await table.list_versions()} + tagged = [ + timestamps[tag["version"]] + for tag in tags.values() + if tag["version"] in timestamps + ] + if not tagged: + return retention + + # LanceDB version timestamps are naive datetimes in local time. + oldest = min(ts.replace(tzinfo=None) for ts in tagged) + needed = datetime.now() - oldest + TAG_RETENTION_MARGIN + return max(retention, needed) + @property def _connection_mode(self) -> ConnectionMode: return ConnectionMode.from_config(self._config) @@ -788,15 +886,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,74 +907,243 @@ 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 _checkout_tables_before(self, before: datetime) -> None: - """Checkout all tables to their state at or before the given datetime. + async def create_tag(self, name: str) -> None: + """Tag the current version of every table with the given name. - Args: - before: The datetime to checkout to + Serializes with client writes via the write lock so a write cannot + land between the version snapshot and the per-table tag creation. + This is in-process coordination only: a writer in another process + can commit between the per-table version reads, so create tags with + all other writers stopped when a consistent snapshot matters. Raises: - ValueError: If no version exists before the given datetime + ReadOnlyError: If the store is in read-only mode. + ValueError: If a rebuild is in progress, or 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. """ - # LanceDB stores timestamps as naive datetimes in local time. - # Convert 'before' to naive local time for comparison. - if before.tzinfo is not None: - # Convert to local time and make naive - before_local = before.astimezone().replace(tzinfo=None) - else: - # Already naive, assume local time - before_local = before + self._assert_writable() + self._assert_not_rebuilding() - 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), + async with self._rebuild_lock, self._write_lock: + await self._create_tag_locked(name) + + async def _create_tag_locked(self, name: str) -> None: + """Create a tag on every table; the caller must hold the write lock.""" + 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" + ) - for table_name, table in tables: - 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 - best_version = None - best_timestamp = None + versions = await self.current_table_versions() + try: + for table_name, table in tables.items(): + await table.tags.create(name, versions[table_name]) + except BaseException as exc: + # BaseException: cancellation must also trigger cleanup, and the + # cleanup itself is protected from further cancellation. The + # sweep covers all tables, not only the recorded ones: a + # cancellation can land after lance committed a table's tag but + # before this attempt recorded it, and preflight guarantees the + # name was unused, so any occurrence belongs to this attempt. + (_, failed_cleanup), cancelled = await _wait_protected( + self._delete_tag_locked(name) + ) + if failed_cleanup: + raise RuntimeError( + f"Tag '{name}' creation failed ({exc!r}) and cleanup " + f"failed on: {', '.join(failed_cleanup)}. A partial " + "tag may remain; delete it with delete_tag." + ) from exc + if cancelled and not isinstance(exc, asyncio.CancelledError): + raise asyncio.CancelledError() + raise - for v in versions: - # LanceDB version timestamps are naive datetime objects in local time - v_timestamp = v["timestamp"] - # Make sure it's naive for comparison - if v_timestamp.tzinfo is not None: - v_timestamp = v_timestamp.replace(tzinfo=None) + async def _delete_tag_locked(self, name: str) -> tuple[bool, list[str]]: + """Delete the tag from every table that has it; the caller must + hold the write lock. - if v_timestamp <= before_local: - if best_timestamp is None or v_timestamp > best_timestamp: - best_version = v["version"] - best_timestamp = v_timestamp + Returns whether the tag was found anywhere and the tables where + listing or deletion failed. + """ + found = False + failed: list[str] = [] + for table_name, table in self._tables().items(): + try: + if name in await table.tags.list(): + found = True + await table.tags.delete(name) + except Exception: + failed.append(table_name) + return found, failed - if best_version is None: - # Find the earliest version to report in error message - if versions: - earliest = min(versions, key=lambda v: v["timestamp"]) - earliest_ts = earliest["timestamp"] - raise ValueError( - f"No data exists before {before}. " - f"Database was created on {earliest_ts}" - ) - else: - raise ValueError( - f"No data exists before {before}. Table has no versions." - ) + async def list_tags(self) -> dict[str, TagInfo]: + """Aggregate per-table tags into database-level tags. - # Checkout to the found version - await table.checkout(best_version) + 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. + + Serializes with create_tag and client writes via the write lock. + + Raises: + ReadOnlyError: If the store is in read-only mode. + ValueError: If a rebuild is in progress or no table has the tag. + RuntimeError: If deletion failed on some tables; remnants remain + until a retry succeeds. + """ + self._assert_writable() + self._assert_not_rebuilding() + async with self._rebuild_lock, self._write_lock: + found, failed = await self._delete_tag_locked(name) + if failed: + # A listing failure obscures whether the tag exists on that + # table, so failures take precedence over not-found. + raise RuntimeError( + f"Tag '{name}' deletion failed on: {', '.join(failed)}. " + "Remnants may remain; retry delete_tag." + ) + if not found: + raise ValueError(f"Tag '{name}' does not exist") + + async def _restore_tables( + self, versions: dict[str, int], *, best_effort: bool = False + ) -> list[tuple[str, Exception]]: + """Restore every table to the given versions, documents last. + + Stops at the first failure by default; with best_effort, continues + through all tables. Returns the failures either way. + """ + tables = self._tables() + failures: list[tuple[str, Exception]] = [] + for table_name in RESTORE_TABLE_ORDER: + try: + await tables[table_name].restore(int(versions[table_name])) + except Exception as exc: + failures.append((table_name, exc)) + if not best_effort: + break + return failures + + async def _rollback_to_snapshot( + self, snapshot: dict[str, int] + ) -> tuple[list[tuple[str, Exception]], bool]: + """Best-effort rollback that a cancellation cannot interrupt. + + Returns the rollback failures and whether a cancellation was + absorbed; the caller must re-deliver an absorbed cancellation. + """ + return await _wait_protected(self._restore_tables(snapshot, best_effort=True)) + + async def restore_tag(self, name: str) -> str: + """Restore every table to the versions of a complete tag. + + Creates a complete safety tag for the pre-restore state before + changing any table and returns its name. Each table restore writes a + new latest version; nothing is left checked out read-only. + + In-process coordination only: all other writers must be stopped for + the duration of the operation. + + Raises: + ReadOnlyError: If the store is in read-only mode. + ValueError: If a rebuild is in progress, the tag does not exist, + or the tag is partial. + RuntimeError: If the safety tag could not be created (no table + changed), or a table restore failed (the error states whether + rollback succeeded). + """ + self._assert_writable() + self._assert_not_rebuilding() + + async with self._rebuild_lock, self._write_lock: + tags = await self.list_tags() + info = tags.get(name) + if info is None: + raise ValueError(f"Tag '{name}' does not exist") + if not info.complete: + raise ValueError( + f"Tag '{name}' is partial (missing tables: " + f"{', '.join(info.missing_tables)}) and cannot be " + "restored; delete it with delete_tag" + ) + + snapshot = await self.current_table_versions() + safety_tag = _safety_tag_name(set(tags)) + try: + await self._create_tag_locked(safety_tag) + except Exception as exc: + raise RuntimeError( + f"Restore of tag '{name}' did not begin: safety tag " + f"creation failed ({exc}). No table was changed." + ) from exc + + try: + failures = await self._restore_tables(info.tables) + except asyncio.CancelledError: + # CancelledError is a BaseException and escapes the + # per-table handler; roll back before re-raising. + rollback_failures, _ = await self._rollback_to_snapshot(snapshot) + if rollback_failures: + failed_names = ", ".join(t for t, _ in rollback_failures) + raise RuntimeError( + f"Restore of tag '{name}' was cancelled and rollback " + f"failed on: {failed_names}. The database may be " + f"cross-table inconsistent; manual recovery is " + f"required using safety tag '{safety_tag}'." + ) from None + raise + if failures: + failed_table, cause = failures[0] + rollback_failures, cancelled = await self._rollback_to_snapshot( + snapshot + ) + if rollback_failures: + failed_names = ", ".join(t for t, _ in rollback_failures) + raise RuntimeError( + f"Restore of tag '{name}' failed on table " + f"'{failed_table}' and rollback failed on: " + f"{failed_names}. The database may be cross-table " + f"inconsistent; manual recovery is required using " + f"safety tag '{safety_tag}'." + ) from cause + if cancelled: + raise asyncio.CancelledError() + raise RuntimeError( + f"Restore of tag '{name}' failed on table " + f"'{failed_table}'; all tables were rolled back to the " + f"pre-restore state. Safety tag '{safety_tag}' is " + "preserved." + ) from cause + + return safety_tag async def list_table_versions(self, table_name: str) -> list[dict[str, Any]]: """List version history for a table. @@ -884,14 +1155,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/haiku_rag_slim/haiku/rag/store/upgrades/v0_45_0.py b/haiku_rag_slim/haiku/rag/store/upgrades/v0_45_0.py index 4cd6e781..a89ed8aa 100644 --- a/haiku_rag_slim/haiku/rag/store/upgrades/v0_45_0.py +++ b/haiku_rag_slim/haiku/rag/store/upgrades/v0_45_0.py @@ -157,10 +157,12 @@ async def _apply_extract_picture_bytes(store: Store) -> None: schema=_V0_45_0_ITEMS_SCHEMA, ) + # Update-only merge: v0.40.0 guarantees a matching row per + # self_ref, and an insert branch would require the source to + # carry every non-nullable column of the live schema. await ( store.document_items_table.merge_insert(["document_id", "self_ref"]) .when_matched_update_all() - .when_not_matched_insert_all() .execute(new_records) ) wrote_items = True diff --git a/haiku_rag_slim/haiku/rag/utils.py b/haiku_rag_slim/haiku/rag/utils.py index 94be9222..fa8f48e9 100644 --- a/haiku_rag_slim/haiku/rag/utils.py +++ b/haiku_rag_slim/haiku/rag/utils.py @@ -1,11 +1,9 @@ import math import sys -from datetime import UTC, datetime from importlib import metadata from pathlib import Path from typing import TYPE_CHECKING, Any, cast -from dateutil import parser as dateutil_parser from packaging.version import Version, parse if TYPE_CHECKING: @@ -38,55 +36,6 @@ def cosine_similarity(vec1: list[float], vec2: list[float]) -> float: return dot_product / (norm1 * norm2) -def parse_datetime(s: str) -> datetime: - """Parse a datetime string into a datetime object. - - Supports: - - ISO 8601 format: "2025-01-15T14:30:00", "2025-01-15T14:30:00Z", "2025-01-15T14:30:00+00:00" - - Date only: "2025-01-15" (interpreted as 00:00:00) - - Various other formats via dateutil - - Args: - s: String to parse - - Returns: - Parsed datetime object - - Raises: - ValueError: If the string cannot be parsed - """ - try: - return dateutil_parser.parse(s) - except (ValueError, TypeError) as e: - raise ValueError( - f"Could not parse datetime: {s}. " - "Use ISO 8601 format (e.g., 2025-01-15T14:30:00) or date (e.g., 2025-01-15)" - ) from e - - -def to_utc(dt: datetime) -> datetime: - """Convert a datetime to UTC. - - - Naive datetimes are assumed to be local time and converted to UTC - - Datetimes with timezone info are converted to UTC - - UTC datetimes are returned as-is - - Args: - dt: Datetime to convert - - Returns: - Datetime in UTC timezone - """ - if dt.tzinfo is None: - # Naive datetime - assume local time - local_dt = dt.astimezone() # Adds local timezone - return local_dt.astimezone(UTC) - elif dt.tzinfo == UTC: - return dt - else: - return dt.astimezone(UTC) - - def apply_common_settings( settings: Any | None, settings_class: type[Any], diff --git a/haiku_rag_slim/pyproject.toml b/haiku_rag_slim/pyproject.toml index 7ba291ee..a04e5681 100644 --- a/haiku_rag_slim/pyproject.toml +++ b/haiku_rag_slim/pyproject.toml @@ -28,7 +28,7 @@ dependencies = [ "jinja2>=3.1.0", "jsonpatch>=1.33", "fastmcp>=3.3.0", - "lancedb==0.30.2", + "lancedb==0.34.0", "pathspec>=1.0.4", "pydantic>=2.12.5", "pydantic-ai-slim[openai,logfire,ag-ui]>=1.100.0", diff --git a/tests/ingester/test_revision_round_trip.py b/tests/ingester/test_revision_round_trip.py index 7bffbde0..e71e5750 100644 --- a/tests/ingester/test_revision_round_trip.py +++ b/tests/ingester/test_revision_round_trip.py @@ -38,7 +38,9 @@ async def test_fs_second_sweep_emits_unchanged_after_ingest(temp_db_path, tmp_pa """The full round-trip: ingest a file, build a sync_state-shaped snapshot from document.metadata, hand it to FSSource.discover() — must see UNCHANGED, not UPSERT. This is exactly what the periodic poller does.""" - file_path = tmp_path / "doc.md" + docs_dir = tmp_path / "docs" + docs_dir.mkdir() + file_path = docs_dir / "doc.md" file_path.write_text("hello") async with HaikuRAG(temp_db_path, create=True) as client: @@ -47,7 +49,7 @@ async def test_fs_second_sweep_emits_unchanged_after_ingest(temp_db_path, tmp_pa assert doc.uri is not None snapshot = {doc.uri: doc.metadata["source_revision"]} - src = FSSource(root=tmp_path) + src = FSSource(root=docs_dir) kinds: list[SourceEventKind] = [] async for event in src.discover(since=snapshot): kinds.append(event.kind) @@ -60,7 +62,9 @@ async def test_fs_second_sweep_emits_unchanged_after_ingest(temp_db_path, tmp_pa async def test_fs_second_sweep_emits_upsert_when_file_changes(temp_db_path, tmp_path): """Counterpart to the unchanged test: a file modified after ingest still triggers UPSERT. Ensures the round-trip doesn't accidentally over-skip.""" - file_path = tmp_path / "doc.md" + docs_dir = tmp_path / "docs" + docs_dir.mkdir() + file_path = docs_dir / "doc.md" file_path.write_text("hello") async with HaikuRAG(temp_db_path, create=True) as client: @@ -74,7 +78,7 @@ async def test_fs_second_sweep_emits_upsert_when_file_changes(temp_db_path, tmp_ # on any sane filesystem, but assert anyway to make the intent explicit. assert str(file_path.stat().st_mtime_ns) != doc.metadata["source_revision"] - src = FSSource(root=tmp_path) + src = FSSource(root=docs_dir) kinds: list[SourceEventKind] = [] async for event in src.discover(since=snapshot): kinds.append(event.kind) @@ -242,7 +246,9 @@ async def test_directory_ingest_threads_configured_source_to_provider( """Directory ingestion with a configured source passes that source's id and fetch context to each child, so the provider sees the configured source id rather than an ad-hoc fs: identity.""" - (tmp_path / "doc.md").write_text("hello") + docs_dir = tmp_path / "docs" + docs_dir.mkdir() + (docs_dir / "doc.md").write_text("hello") seen_source_ids: list[str] = [] @@ -251,11 +257,11 @@ async def test_directory_ingest_threads_configured_source_to_provider( seen_source_ids.append(source_id) return {"collection": source_id} - source = FSSource(root=tmp_path, source_id="docs") + source = FSSource(root=docs_dir, source_id="docs") async with HaikuRAG(temp_db_path, create=True) as client: docs = await client.create_document_from_source( - tmp_path, + docs_dir, sources=[source], source_id="docs", metadata_provider=Provider(), diff --git a/tests/store/test_restore.py b/tests/store/test_restore.py new file mode 100644 index 00000000..5202bb97 --- /dev/null +++ b/tests/store/test_restore.py @@ -0,0 +1,408 @@ +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") + + safety_tag = await store.restore_tag("release-1") + assert safety_tag == "before-restore-20260715T143012Z-2" + + +@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 diff --git a/tests/store/test_tags.py b/tests/store/test_tags.py new file mode 100644 index 00000000..9b20f10a --- /dev/null +++ b/tests/store/test_tags.py @@ -0,0 +1,534 @@ +import asyncio + +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_create_tag_reports_failed_cleanup(temp_db_path, monkeypatch): + """When midway-failure cleanup also fails, the error reports both the + original failure and the remaining partial-tag risk.""" + async with Store(temp_db_path, create=True) as store: + real_create = AsyncTags.create + calls = {"n": 0} + + async def flaky_create(self, name: str, version: int) -> None: + calls["n"] += 1 + if calls["n"] == 4: + raise RuntimeError("create boom") + await real_create(self, name, version) + + async def failing_delete(self, name: str) -> None: + raise RuntimeError("delete boom") + + monkeypatch.setattr(AsyncTags, "create", flaky_create) + monkeypatch.setattr(AsyncTags, "delete", failing_delete) + + with pytest.raises(RuntimeError) as exc_info: + await store.create_tag("broken") + + msg = str(exc_info.value) + assert "create boom" in msg + assert "partial" in msg + assert exc_info.value.__cause__ is not None + + monkeypatch.undo() + tags = await store.list_tags() + assert tags["broken"].complete is False + + +@pytest.mark.asyncio +async def test_delete_tag_reports_failed_tables(temp_db_path, monkeypatch): + """delete_tag never claims success when remnants remain: it names the + tables where deletion failed.""" + async with Store(temp_db_path, create=True) as store: + await store.create_tag("release-1") + + real_delete = AsyncTags.delete + calls = {"n": 0} + + async def flaky_delete(self, name: str) -> None: + calls["n"] += 1 + if calls["n"] == 2: + raise RuntimeError("delete boom") + await real_delete(self, name) + + monkeypatch.setattr(AsyncTags, "delete", flaky_delete) + + with pytest.raises(RuntimeError) as exc_info: + await store.delete_tag("release-1") + + assert "document_meta" in str(exc_info.value) + + monkeypatch.undo() + tags = await store.list_tags() + assert set(tags["release-1"].tables) == {"document_meta"} + + await store.delete_tag("release-1") + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_create_tag_waits_for_write_lock(temp_db_path): + """create_tag serializes with client writes so a write cannot land + between the version snapshot and the per-table tag creation.""" + async with Store(temp_db_path, create=True) as store: + async with store._write_lock: + task = asyncio.create_task(store.create_tag("release-1")) + await asyncio.sleep(0.1) + assert not task.done() + await task + + tags = await store.list_tags() + assert tags["release-1"].complete is True + + +@pytest.mark.asyncio +async def test_delete_tag_waits_for_write_lock(temp_db_path): + """delete_tag serializes with create_tag and client writes so it cannot + remove tags out from under a concurrent create_tag.""" + async with Store(temp_db_path, create=True) as store: + await store.create_tag("release-1") + + async with store._write_lock: + task = asyncio.create_task(store.delete_tag("release-1")) + await asyncio.sleep(0.1) + assert not task.done() + await task + + assert await store.list_tags() == {} + + +@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_vacuum_cleans_untagged_versions_and_keeps_tagged(temp_db_path): + """Vacuum must both preserve tagged versions (lance hard-errors when a + tagged version falls inside the cleanup window, which vacuum would + swallow) and still clean untagged versions older than the oldest tag's + safety margin.""" + async with Store(temp_db_path, create=True) as store: + repo = DocumentRepository(store) + await repo.create(Document(content="First document")) + versions_before = [ + v["version"] for v in await store.list_table_versions("documents") + ] + + # Age the pre-tag versions past the retention safety margin. + await asyncio.sleep(1.5) + + await repo.create(Document(content="Second document")) + await store.create_tag("release-1") + tagged_version = (await store.list_tags())["release-1"].tables["documents"] + + await store.vacuum(retention_seconds=0) + + remaining = [v["version"] for v in await store.list_table_versions("documents")] + assert tagged_version in remaining + assert min(versions_before) not in remaining + + await store.documents_table.checkout("release-1") + rows = await store.documents_table.count_rows() + await store.documents_table.checkout_latest() + assert rows == 2 + + +@pytest.mark.asyncio +async def test_vacuum_reraises_runtime_error(temp_db_path, monkeypatch): + """Vacuum suppresses OSError only; lance errors (RuntimeError) surface + instead of silently skipping cleanup.""" + from lancedb.table import AsyncTable + + async with Store(temp_db_path, create=True) as store: + + async def failing_optimize(self, **kwargs): + raise RuntimeError("lance error: boom") + + monkeypatch.setattr(AsyncTable, "optimize", failing_optimize) + with pytest.raises(RuntimeError, match="boom"): + await store.vacuum(retention_seconds=0) + + async def failing_optimize_os(self, **kwargs): + raise OSError("disk full") + + monkeypatch.setattr(AsyncTable, "optimize", failing_optimize_os) + await store.vacuum(retention_seconds=0) + + +@pytest.mark.asyncio +async def test_vacuum_multiple_tags_uses_oldest_cutoff(temp_db_path): + """With several tags the retention clamp must key off the oldest one; + clamping to a newer tag would put the older tagged version inside the + cleanup window and lance would hard-error.""" + async with Store(temp_db_path, create=True) as store: + repo = DocumentRepository(store) + await repo.create(Document(content="First document")) + await store.create_tag("old") + + await asyncio.sleep(1.5) + + await repo.create(Document(content="Second document")) + await store.create_tag("new") + + await store.vacuum(retention_seconds=0) + + tags = await store.list_tags() + remaining = [v["version"] for v in await store.list_table_versions("documents")] + assert tags["old"].tables["documents"] in remaining + assert tags["new"].tables["documents"] in remaining + + +@pytest.mark.asyncio +async def test_vacuum_partial_tag_protects_its_tables(temp_db_path): + """A partial tag still protects the versions of the tables it exists on, + while untagged tables clean up normally.""" + async with Store(temp_db_path, create=True) as store: + repo = DocumentRepository(store) + await repo.create(Document(content="First document")) + + chunks_version = await store.chunks_table.version() + await store.chunks_table.tags.create("stale", chunks_version) + docs_versions_before = [ + v["version"] for v in await store.list_table_versions("documents") + ] + + await asyncio.sleep(1.5) + + await repo.create(Document(content="Second document")) + await store.vacuum(retention_seconds=0) + + chunk_versions = [ + v["version"] for v in await store.list_table_versions("chunks") + ] + assert chunks_version in chunk_versions + + docs_versions_after = [ + v["version"] for v in await store.list_table_versions("documents") + ] + assert min(docs_versions_before) not in docs_versions_after + + +@pytest.mark.asyncio +async def test_deleting_oldest_tag_advances_cleanup(temp_db_path): + """Versions pinned by a tag become cleanable once the tag is deleted; + the cleanup cutoff advances to the next retained tag without removing + its version.""" + async with Store(temp_db_path, create=True) as store: + repo = DocumentRepository(store) + await repo.create(Document(content="First document")) + await store.create_tag("old") + old_version = (await store.list_tags())["old"].tables["documents"] + + await asyncio.sleep(1.5) + + await repo.create(Document(content="Second document")) + await store.create_tag("new") + new_version = (await store.list_tags())["new"].tables["documents"] + + await store.vacuum(retention_seconds=0) + remaining = [v["version"] for v in await store.list_table_versions("documents")] + assert old_version in remaining + assert new_version in remaining + + await store.delete_tag("old") + await asyncio.sleep(1.5) + await store.vacuum(retention_seconds=0) + + remaining = [v["version"] for v in await store.list_table_versions("documents")] + assert old_version not in remaining + assert new_version in remaining + + +@pytest.mark.asyncio +async def test_tag_operations_rejected_during_rebuild(temp_db_path): + """While a rebuild holds the rebuild lock, tag operations fail fast + instead of snapshotting a half-rebuilt database.""" + async with Store(temp_db_path, create=True) as store: + await store.create_tag("keep") + + async with store._rebuild_lock: + with pytest.raises(ValueError, match="[Rr]ebuild in progress"): + await store.create_tag("release-1") + with pytest.raises(ValueError, match="[Rr]ebuild in progress"): + await store.delete_tag("keep") + + await store.create_tag("release-1") + await store.delete_tag("keep") + assert set(await store.list_tags()) == {"release-1"} + + +@pytest.mark.asyncio +async def test_vacuum_waits_for_write_lock(temp_db_path): + """Vacuum serializes with writers and tag operations so a tag cannot be + created between _tag_safe_retention's read and the optimize call.""" + async with Store(temp_db_path, create=True) as store: + async with store._write_lock: + task = asyncio.create_task(store.vacuum(retention_seconds=0)) + await asyncio.sleep(0.1) + assert not task.done() + await task + + +@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 + + +@pytest.mark.asyncio +async def test_current_table_versions_returns_versions(temp_db_path): + """current_table_versions returns dict of table versions.""" + async with Store(temp_db_path, create=True) as store: + versions = await store.current_table_versions() + + assert "documents" in versions + assert "chunks" in versions + assert "settings" in versions + assert all(isinstance(v, int) for v in versions.values()) + + +@pytest.mark.asyncio +async def test_list_table_versions_returns_history(temp_db_path): + """list_table_versions returns version history for a table.""" + async with Store(temp_db_path, create=True) as store: + versions = await store.list_table_versions("documents") + + assert len(versions) >= 1 + for v in versions: + assert "version" in v + assert "timestamp" in v + + +@pytest.mark.asyncio +async def test_delete_tag_reports_listing_failures(temp_db_path, monkeypatch): + """A tags.list() failure mid-delete is reported with the table named and + a recovery hint, instead of escaping raw after earlier deletions.""" + async with Store(temp_db_path, create=True) as store: + await store.create_tag("release-1") + + real_list = AsyncTags.list + calls = {"n": 0} + + async def flaky_list(self): + calls["n"] += 1 + if calls["n"] == 2: + raise RuntimeError("list boom") + return await real_list(self) + + monkeypatch.setattr(AsyncTags, "list", flaky_list) + + with pytest.raises(RuntimeError) as exc_info: + await store.delete_tag("release-1") + + msg = str(exc_info.value) + assert "document_meta" in msg + assert "retry delete_tag" in msg + + monkeypatch.undo() + tags = await store.list_tags() + assert set(tags["release-1"].tables) == {"document_meta"} + + await store.delete_tag("release-1") + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_create_tag_cancellation_cleans_up(temp_db_path, monkeypatch): + """Cancellation during per-table tag creation must not leave a partial + tag behind: cleanup runs before the cancellation propagates.""" + async with Store(temp_db_path, create=True) as store: + real_create = AsyncTags.create + calls = {"n": 0} + + async def cancelled_create(self, name: str, version: int) -> None: + calls["n"] += 1 + if calls["n"] == 4: + raise asyncio.CancelledError() + await real_create(self, name, version) + + monkeypatch.setattr(AsyncTags, "create", cancelled_create) + + with pytest.raises(asyncio.CancelledError): + await store.create_tag("broken") + + monkeypatch.undo() + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_create_tag_cleanup_survives_cancellation(temp_db_path, monkeypatch): + """Cancelling create_tag while it cleans up a failed creation does not + interrupt the cleanup: no partial tag remains and the cancellation is + delivered afterwards.""" + async with Store(temp_db_path, create=True) as store: + real_create = AsyncTags.create + real_delete = AsyncTags.delete + create_calls = {"n": 0} + cleanup_started = asyncio.Event() + release = asyncio.Event() + + async def flaky_create(self, name: str, version: int) -> None: + create_calls["n"] += 1 + if create_calls["n"] == 4: + raise RuntimeError("create boom") + await real_create(self, name, version) + + async def slow_delete(self, name: str) -> None: + cleanup_started.set() + await release.wait() + await real_delete(self, name) + + monkeypatch.setattr(AsyncTags, "create", flaky_create) + monkeypatch.setattr(AsyncTags, "delete", slow_delete) + + task = asyncio.create_task(store.create_tag("broken")) + await cleanup_started.wait() + task.cancel() + release.set() + + with pytest.raises(asyncio.CancelledError): + await task + + monkeypatch.undo() + assert await store.list_tags() == {} + + +@pytest.mark.asyncio +async def test_create_tag_cancellation_after_commit_cleans_committed_tag( + temp_db_path, monkeypatch +): + """Cancellation arriving after lance committed a table's tag but before + the attempt recorded it must still clean that table: cleanup sweeps all + tables, relying on the preflight guarantee that the name was unused.""" + async with Store(temp_db_path, create=True) as store: + real_create = AsyncTags.create + calls = {"n": 0} + + async def committing_cancelled_create(self, name: str, version: int) -> None: + calls["n"] += 1 + await real_create(self, name, version) + if calls["n"] == 4: + raise asyncio.CancelledError() + + monkeypatch.setattr(AsyncTags, "create", committing_cancelled_create) + + with pytest.raises(asyncio.CancelledError): + await store.create_tag("broken") + + monkeypatch.undo() + assert await store.list_tags() == {} diff --git a/tests/store/test_time_travel.py b/tests/store/test_time_travel.py deleted file mode 100644 index db0a7e85..00000000 --- a/tests/store/test_time_travel.py +++ /dev/null @@ -1,96 +0,0 @@ -import asyncio -from datetime import UTC, datetime, timedelta - -import pytest - -from haiku.rag.store import ReadOnlyError, Store -from haiku.rag.store.models import Document -from haiku.rag.store.repositories.document import DocumentRepository - - -class TestStoreTimeTravel: - @pytest.mark.asyncio - async def test_store_with_before_is_read_only(self, temp_db_path): - """Store with before parameter is automatically read-only.""" - async with Store(temp_db_path, create=True): - pass - - before = datetime.now(UTC) + timedelta(hours=1) - async with Store(temp_db_path, before=before) as store: - assert store.is_read_only is True - - @pytest.mark.asyncio - async def test_store_before_raises_on_write(self, temp_db_path): - """Store with before parameter raises on write operations.""" - async with Store(temp_db_path, create=True): - pass - - before = datetime.now(UTC) + timedelta(hours=1) - async with Store(temp_db_path, before=before) as store: - with pytest.raises(ReadOnlyError): - store._assert_writable() - - @pytest.mark.asyncio - async def test_store_before_checks_out_historical_state(self, temp_db_path): - """Store with before parameter checks out tables to historical state.""" - async with Store(temp_db_path, create=True) as store: - repo = DocumentRepository(store) - await repo.create(Document(content="First document")) - - versions_after_first = await store.list_table_versions("documents") - latest_version = max(versions_after_first, key=lambda v: v["version"]) - time_after_first = latest_version["timestamp"] - - await asyncio.sleep(0.5) - - await repo.create(Document(content="Second document")) - - versions_after_second = await store.list_table_versions("documents") - assert len(versions_after_second) > len(versions_after_first) - - async with Store(temp_db_path, before=time_after_first) as store: - repo = DocumentRepository(store) - - docs = await repo.list_all(include_content=True) - assert len(docs) == 1 - assert docs[0].content == "First document" - - async with Store(temp_db_path) as store: - repo = DocumentRepository(store) - - docs = await repo.list_all() - assert len(docs) == 2 - - @pytest.mark.asyncio - async def test_store_before_no_version_raises(self, temp_db_path): - """Store with before datetime before any version raises ValueError.""" - async with Store(temp_db_path, create=True): - pass - - before = datetime(2000, 1, 1, tzinfo=UTC) - with pytest.raises(ValueError) as exc_info: - async with Store(temp_db_path, before=before): - pass - assert "No data exists before" in str(exc_info.value) - - @pytest.mark.asyncio - async def test_current_table_versions_returns_versions(self, temp_db_path): - """current_table_versions returns dict of table versions.""" - async with Store(temp_db_path, create=True) as store: - versions = await store.current_table_versions() - - assert "documents" in versions - assert "chunks" in versions - assert "settings" in versions - assert all(isinstance(v, int) for v in versions.values()) - - @pytest.mark.asyncio - async def test_list_table_versions_returns_history(self, temp_db_path): - """list_table_versions returns version history for a table.""" - async with Store(temp_db_path, create=True) as store: - versions = await store.list_table_versions("documents") - - assert len(versions) >= 1 - for v in versions: - assert "version" in v - assert "timestamp" in v diff --git a/tests/test_cli.py b/tests/test_cli.py index fbc1acfb..d16e1ff0 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -84,3 +84,190 @@ class TestCliMigrationError: with pytest.raises(SystemExit) as exc_info: cli_wrapper() assert exc_info.value.code == 1 + + +class TestTagCommands: + def test_tag_round_trip(self, temp_db_path): + db = str(temp_db_path) + + result = runner.invoke(cli, ["init", "--db", db]) + assert result.exit_code == 0 + + result = runner.invoke(cli, ["tag", "create", "release-1", "--db", db]) + assert result.exit_code == 0 + assert "release-1" in result.output + + result = runner.invoke(cli, ["tag", "list", "--db", db]) + assert result.exit_code == 0 + assert "release-1" in result.output + assert "partial" not in result.output + + result = runner.invoke(cli, ["history", "--db", db, "-t", "documents"]) + assert result.exit_code == 0 + assert "release-1" in result.output + + result = runner.invoke(cli, ["tag", "create", "release-1", "--db", db]) + assert result.exit_code == 1 + assert "already exists" in result.output + + result = runner.invoke(cli, ["tag", "delete", "release-1", "--db", db]) + assert result.exit_code == 0 + + result = runner.invoke(cli, ["tag", "list", "--db", db]) + assert result.exit_code == 0 + assert "No tags" in result.output + + result = runner.invoke(cli, ["tag", "delete", "release-1", "--db", db]) + assert result.exit_code == 1 + assert "does not exist" in result.output + + def test_tag_create_rejected_when_migrations_pending(self, temp_db_path): + """A writable tag operation must hit the migration gate and must not + mutate a legacy database (e.g. by creating missing tables).""" + import asyncio + + import lancedb + + from haiku.rag.store.engine import Store + + async def _prepare_legacy_db(): + async with Store(temp_db_path, create=True) as store: + await store.set_haiku_version("0.19.0") + db = await lancedb.connect_async(temp_db_path.absolute()) + await db.drop_table("document_meta") + db.close() + + asyncio.run(_prepare_legacy_db()) + + result = runner.invoke( + cli, ["tag", "create", "release-1", "--db", str(temp_db_path)] + ) + assert result.exit_code == 1 + assert isinstance(result.exception, MigrationRequiredError) + + async def _table_names() -> list[str]: + db = await lancedb.connect_async(temp_db_path.absolute()) + tables = (await db.list_tables()).tables + db.close() + return tables + + assert "document_meta" not in asyncio.run(_table_names()) + + def test_tag_commands_missing_database_exit_nonzero(self, tmp_path): + missing = str(tmp_path / "does_not_exist.lancedb") + for args in ( + ["tag", "create", "r1", "--db", missing], + ["tag", "delete", "r1", "--db", missing], + ["tag", "list", "--db", missing], + ): + result = runner.invoke(cli, args) + assert result.exit_code == 1, args + assert "does not exist" in result.output, args + + def test_tag_create_invalid_name_fails_cleanly(self, temp_db_path): + """lance restricts ref names to alphanumeric, '.', '-', '_'; the CLI + surfaces that as a clean error instead of a traceback.""" + db = str(temp_db_path) + result = runner.invoke(cli, ["init", "--db", db]) + assert result.exit_code == 0 + + result = runner.invoke(cli, ["tag", "create", "[red]release[/red]", "--db", db]) + assert result.exit_code == 1 + assert "Error:" in result.output + assert "Ref characters" in result.output + + +class TestTagRestore: + def test_restore_requires_confirmation_and_decline_changes_nothing( + self, temp_db_path + ): + db = str(temp_db_path) + assert runner.invoke(cli, ["init", "--db", db]).exit_code == 0 + assert runner.invoke(cli, ["tag", "create", "r1", "--db", db]).exit_code == 0 + + result = runner.invoke(cli, ["tag", "restore", "r1", "--db", db], input="n\n") + assert result.exit_code == 1 + assert "live database state" in result.output + assert "Stop all ingestion" in result.output + assert "not transactionally atomic" in result.output + assert "safety tag" in result.output + + result = runner.invoke(cli, ["tag", "list", "--db", db]) + assert "before-restore" not in result.output + + def test_restore_non_interactive_without_yes_fails(self, temp_db_path): + db = str(temp_db_path) + assert runner.invoke(cli, ["init", "--db", db]).exit_code == 0 + assert runner.invoke(cli, ["tag", "create", "r1", "--db", db]).exit_code == 0 + + result = runner.invoke(cli, ["tag", "restore", "r1", "--db", db]) + assert result.exit_code == 1 + + result = runner.invoke(cli, ["tag", "list", "--db", db]) + assert "before-restore" not in result.output + + def test_restore_with_yes(self, temp_db_path): + db = str(temp_db_path) + assert runner.invoke(cli, ["init", "--db", db]).exit_code == 0 + assert runner.invoke(cli, ["tag", "create", "r1", "--db", db]).exit_code == 0 + + result = runner.invoke(cli, ["tag", "restore", "r1", "--yes", "--db", db]) + assert result.exit_code == 0 + assert "Restored database to tag 'r1'" in result.output + assert "before-restore-" in result.output + assert "now live" in result.output + assert "migrate" in result.output + + result = runner.invoke(cli, ["tag", "list", "--db", db]) + assert "before-restore-" in result.output + + def test_restore_missing_tag_errors(self, temp_db_path): + db = str(temp_db_path) + assert runner.invoke(cli, ["init", "--db", db]).exit_code == 0 + + result = runner.invoke(cli, ["tag", "restore", "nope", "--yes", "--db", db]) + assert result.exit_code == 1 + assert "does not exist" in result.output + + def test_restore_partial_tag_errors(self, temp_db_path): + import asyncio + + from haiku.rag.store.engine import Store + + async def _partial_tag(): + async with Store(temp_db_path, create=True) as store: + version = await store.chunks_table.version() + await store.chunks_table.tags.create("stale", version) + + asyncio.run(_partial_tag()) + + result = runner.invoke( + cli, ["tag", "restore", "stale", "--yes", "--db", str(temp_db_path)] + ) + assert result.exit_code == 1 + assert "partial" in result.output + assert "documents" in result.output + + def test_restore_missing_database_exits_nonzero(self, tmp_path): + missing = tmp_path / "does_not_exist.lancedb" + result = runner.invoke( + cli, ["tag", "restore", "r1", "--yes", "--db", str(missing)] + ) + assert result.exit_code == 1 + assert "does not exist" in result.output + + # Without --yes the missing database is reported before the + # confirmation prompt, not after the user confirms. + result = runner.invoke(cli, ["tag", "restore", "r1", "--db", str(missing)]) + assert result.exit_code == 1 + assert "does not exist" in result.output + assert "Continue?" not in result.output + + def test_tag_help_includes_restore(self): + result = runner.invoke(cli, ["tag", "--help"]) + assert result.exit_code == 0 + assert "restore" in result.output + + result = runner.invoke(cli, ["--help"]) + assert "--before" not in result.output + assert "--at" not in result.output diff --git a/tests/test_client.py b/tests/test_client.py index 40884c0a..63956c82 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -2264,3 +2264,30 @@ async def test_rebuild_rechunk_with_url_prefixed_stored_content( assert doc_after is not None assert "example.com" in doc_after.content assert "Stored" in doc_after.content + + +async def test_metadata_only_update_waits_for_write_lock(temp_db_path): + """The metadata-only update path serializes with other writers so it + cannot land inside another writer's critical section (e.g. between + create_tag's version snapshot and its per-table tag creation).""" + import asyncio + + dim = Config.embeddings.model.vector_dim + docling_doc = DoclingDocument(name="d") + docling_doc.add_text(label=DocItemLabel.TEXT, text="body") + + async with HaikuRAG(temp_db_path, create=True) as client: + doc = await client.import_document( + docling_doc, + [Chunk(content="body", embedding=[0.1] * dim, order=0)], + uri="mem://meta", + ) + + async with client.store._write_lock: + task = asyncio.create_task( + client.update_document(document_id=doc.id, metadata={"k": "v"}) + ) + await asyncio.sleep(0.1) + assert not task.done() + updated = await task + assert updated.metadata == {"k": "v"} diff --git a/tests/test_info.py b/tests/test_info.py index 42206fef..a3b3f4c0 100644 --- a/tests/test_info.py +++ b/tests/test_info.py @@ -355,7 +355,77 @@ async def test_app_history_skips_exists_check_for_remote(tmp_path): with patch("haiku.rag.store.engine.Store") as mock_store_cls: mock_store = AsyncMock() mock_store.list_table_versions = AsyncMock(return_value=[]) + mock_store.list_tags = AsyncMock(return_value={}) mock_store_cls.return_value.__aenter__ = AsyncMock(return_value=mock_store) mock_store_cls.return_value.__aexit__ = AsyncMock(return_value=False) await app.history() mock_store_cls.assert_called_once() + + +@pytest.mark.asyncio +async def test_app_tag_rendering_escapes_markup(tmp_path): + """lance forbids markup characters in ref names, but externally created + tags are rendered defensively: markup-looking names must come out as + literal text in tag list and history, not be interpreted by Rich.""" + from rich.console import Console + + from haiku.rag.store.engine import TagInfo + + config = AppConfig( + lancedb=LanceDBConfig( + uri="s3://bucket/path", + storage_options={"endpoint": "http://localhost:9000"}, + ) + ) + app = HaikuRAGApp(db_path=tmp_path / "db.lancedb", config=config) + app.console = Console(record=True, width=200) + + hostile = "[red]release[/red]" + tags = {hostile: TagInfo(tables={"documents": 1}, missing_tables=[])} + + with patch("haiku.rag.store.engine.Store") as mock_store_cls: + mock_store = AsyncMock() + mock_store.list_tags = AsyncMock(return_value=tags) + mock_store.list_table_versions = AsyncMock( + return_value=[{"version": 1, "timestamp": "2026-07-14 10:00:00"}] + ) + mock_store_cls.return_value.__aenter__ = AsyncMock(return_value=mock_store) + mock_store_cls.return_value.__aexit__ = AsyncMock(return_value=False) + + await app.list_tags() + await app.history(table="documents") + + output = app.console.export_text() + assert output.count(hostile) == 2 + + +@pytest.mark.asyncio +async def test_app_history_survives_tag_annotation_failure(tmp_path): + """history degrades to version history without annotations, with a + warning, when aggregate tag loading fails.""" + from rich.console import Console + + config = AppConfig( + lancedb=LanceDBConfig( + uri="s3://bucket/path", + storage_options={"endpoint": "http://localhost:9000"}, + ) + ) + app = HaikuRAGApp(db_path=tmp_path / "db.lancedb", config=config) + app.console = Console(record=True, width=200) + + with patch("haiku.rag.store.engine.Store") as mock_store_cls: + mock_store = AsyncMock() + mock_store.list_tags = AsyncMock(side_effect=RuntimeError("tags boom")) + mock_store.list_table_versions = AsyncMock( + return_value=[{"version": 1, "timestamp": "2026-07-15 10:00:00"}] + ) + mock_store_cls.return_value.__aenter__ = AsyncMock(return_value=mock_store) + mock_store_cls.return_value.__aexit__ = AsyncMock(return_value=False) + + await app.history(table="documents") + + output = app.console.export_text() + assert "v1" in output + assert "2026-07-15 10:00:00" in output + assert "tags boom" in output diff --git a/tests/test_rebuild.py b/tests/test_rebuild.py index 9649a625..c371bae2 100644 --- a/tests/test_rebuild.py +++ b/tests/test_rebuild.py @@ -1049,3 +1049,49 @@ async def test_rebuild_set_embedder_raises_on_vector_dim_mismatch(temp_db_path): with pytest.raises(ConfigMismatchError): async for _ in client.rebuild_database(mode=RebuildMode.SET_EMBEDDER): pass + + +async def test_rebuild_blocks_tag_operations(temp_db_path, monkeypatch): + """rebuild_database holds the rebuild lock for its whole run: tag + operations fail mid-rebuild and work again once it completes.""" + import random + + from docling_core.types.doc.document import DoclingDocument + from docling_core.types.doc.labels import DocItemLabel + + from haiku.rag.embeddings import EmbedderWrapper + from haiku.rag.store.models.chunk import Chunk + + async def fake_embed_documents(self, texts): + result = [] + for t in texts: + random.seed(hash(t) % (2**32)) + result.append([random.random() for _ in range(2560)]) + return result + + monkeypatch.setattr(EmbedderWrapper, "embed_documents", fake_embed_documents) + + dim = Config.embeddings.model.vector_dim + docling_doc = DoclingDocument(name="d") + docling_doc.add_text(label=DocItemLabel.TEXT, text="body") + + async with HaikuRAG(temp_db_path, create=True) as client: + await client.import_document( + docling_doc, + [Chunk(content="body", embedding=[0.1] * dim, order=0)], + uri="mem://rebuild", + ) + + rebuild = client.rebuild_database(mode=RebuildMode.EMBED_ONLY) + await anext(rebuild) + + assert client.store._rebuild_lock.locked() + with pytest.raises(ValueError, match="[Rr]ebuild in progress"): + await client.store.create_tag("mid-rebuild") + + async for _ in rebuild: + pass + + assert not client.store._rebuild_lock.locked() + await client.store.create_tag("post-rebuild") + assert set(await client.store.list_tags()) == {"post-rebuild"} diff --git a/tests/test_utils.py b/tests/test_utils.py index 0d1944b6..4214784b 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -382,79 +382,6 @@ def test_get_package_versions(): assert len(value) > 0 -# --- parse_datetime tests --- - - -def test_parse_datetime_iso8601(): - from haiku.rag.utils import parse_datetime - - dt = parse_datetime("2025-01-15T14:30:00") - assert dt.year == 2025 - assert dt.month == 1 - assert dt.day == 15 - assert dt.hour == 14 - assert dt.minute == 30 - - -def test_parse_datetime_date_only(): - from haiku.rag.utils import parse_datetime - - dt = parse_datetime("2025-01-15") - assert dt.year == 2025 - assert dt.month == 1 - assert dt.day == 15 - - -def test_parse_datetime_with_timezone(): - from haiku.rag.utils import parse_datetime - - dt = parse_datetime("2025-01-15T14:30:00+00:00") - assert dt.year == 2025 - assert dt.tzinfo is not None - - -def test_parse_datetime_invalid(): - from haiku.rag.utils import parse_datetime - - with pytest.raises(ValueError, match="Could not parse datetime"): - parse_datetime("not-a-date") - - -# --- to_utc tests --- - - -def test_to_utc_naive_datetime(): - from datetime import datetime - - from haiku.rag.utils import to_utc - - naive = datetime(2025, 6, 15, 12, 0, 0) - result = to_utc(naive) - assert result.tzinfo is not None - - -def test_to_utc_utc_datetime(): - from datetime import UTC, datetime - - from haiku.rag.utils import to_utc - - utc_dt = datetime(2025, 6, 15, 12, 0, 0, tzinfo=UTC) - result = to_utc(utc_dt) - assert result is utc_dt - - -def test_to_utc_aware_non_utc(): - from datetime import UTC, datetime, timedelta, timezone - - from haiku.rag.utils import to_utc - - eastern = timezone(timedelta(hours=-5)) - aware = datetime(2025, 6, 15, 12, 0, 0, tzinfo=eastern) - result = to_utc(aware) - assert result.tzinfo == UTC - assert result.hour == 17 - - # --- apply_common_settings tests --- diff --git a/tests/test_vacuum_debounce.py b/tests/test_vacuum_debounce.py index 2b4176f8..83a96dbc 100644 --- a/tests/test_vacuum_debounce.py +++ b/tests/test_vacuum_debounce.py @@ -92,3 +92,33 @@ async def test_metadata_refresh_sweep_schedules_vacuum(temp_db_path): source_metadata={"source_revision": "r2", "md5": "same"}, ) assert client._vacuum_dirty is True + + +@pytest.mark.asyncio +async def test_metadata_refresh_waits_for_write_lock(temp_db_path): + """The revision/MD5 short-circuit write serializes with other writers so + it cannot land inside another writer's critical section (e.g. between + create_tag's version snapshot and its per-table tag creation).""" + dim = Config.embeddings.model.vector_dim + async with HaikuRAG(temp_db_path, create=True) as client: + doc = await client.import_document( + _docling_doc("d", "body"), + [Chunk(content="body", embedding=[0.1] * dim, order=0)], + uri="mem://sweep", + metadata={"source_revision": "r1"}, + ) + + async with client.store._write_lock: + task = asyncio.create_task( + _refresh_doc_metadata( + client, + doc, + title=None, + user_metadata={}, + source_metadata={"source_revision": "r2", "md5": "same"}, + ) + ) + await asyncio.sleep(0.1) + assert not task.done() + refreshed = await task + assert refreshed.metadata["source_revision"] == "r2" diff --git a/tests/utils/test_datetime.py b/tests/utils/test_datetime.py deleted file mode 100644 index 6d2f406b..00000000 --- a/tests/utils/test_datetime.py +++ /dev/null @@ -1,85 +0,0 @@ -from datetime import UTC, datetime, timezone - -import pytest - -from haiku.rag.utils import parse_datetime, to_utc - - -class TestParseDateTime: - def test_parse_iso8601_with_timezone(self): - """Parse ISO 8601 datetime with timezone.""" - result = parse_datetime("2025-01-15T14:30:00+00:00") - assert result.year == 2025 - assert result.month == 1 - assert result.day == 15 - assert result.hour == 14 - assert result.minute == 30 - assert result.second == 0 - assert result.tzinfo is not None - - def test_parse_iso8601_without_timezone(self): - """Parse ISO 8601 datetime without timezone (naive).""" - result = parse_datetime("2025-01-15T14:30:00") - assert result.year == 2025 - assert result.month == 1 - assert result.day == 15 - assert result.hour == 14 - assert result.minute == 30 - - def test_parse_date_only(self): - """Parse date-only string as start of day.""" - result = parse_datetime("2025-01-15") - assert result.year == 2025 - assert result.month == 1 - assert result.day == 15 - assert result.hour == 0 - assert result.minute == 0 - assert result.second == 0 - - def test_parse_various_formats(self): - """Parse various datetime formats.""" - # ISO with Z suffix - result = parse_datetime("2025-01-15T14:30:00Z") - assert result.year == 2025 - assert result.month == 1 - assert result.day == 15 - - # With milliseconds - result = parse_datetime("2025-01-15T14:30:00.123") - assert result.microsecond == 123000 - - def test_parse_invalid_raises_value_error(self): - """Invalid datetime string raises ValueError.""" - with pytest.raises(ValueError) as exc_info: - parse_datetime("not-a-datetime") - assert "Could not parse datetime" in str(exc_info.value) - - -class TestToUtc: - def test_naive_datetime_assumes_local_and_converts(self): - """Naive datetime is assumed local and converted to UTC.""" - naive = datetime(2025, 1, 15, 14, 30, 0) - result = to_utc(naive) - assert result.tzinfo == UTC - - def test_utc_datetime_unchanged(self): - """UTC datetime is returned as-is.""" - utc_dt = datetime(2025, 1, 15, 14, 30, 0, tzinfo=UTC) - result = to_utc(utc_dt) - assert result == utc_dt - assert result.tzinfo == UTC - - def test_other_timezone_converts_to_utc(self): - """Datetime with other timezone is converted to UTC.""" - from datetime import timedelta - - # Create a datetime at UTC+5 - tz_plus5 = timezone(timedelta(hours=5)) - dt_plus5 = datetime(2025, 1, 15, 19, 30, 0, tzinfo=tz_plus5) - - result = to_utc(dt_plus5) - - # 19:30 UTC+5 = 14:30 UTC - assert result.tzinfo == UTC - assert result.hour == 14 - assert result.minute == 30 diff --git a/uv.lock b/uv.lock index 7ec3a648..1b87ee69 100644 --- a/uv.lock +++ b/uv.lock @@ -1755,7 +1755,7 @@ requires-dist = [ { name = "httpx", specifier = ">=0.28.1" }, { name = "jinja2", specifier = ">=3.1.0" }, { name = "jsonpatch", specifier = ">=1.33" }, - { name = "lancedb", specifier = "==0.30.2" }, + { name = "lancedb", specifier = "==0.34.0" }, { name = "obstore", marker = "extra == 's3'", specifier = ">=0.9,<0.10" }, { name = "opencv-python-headless", marker = "extra == 'docling'", specifier = ">=4.6.0.66,<5.0.0.0" }, { name = "pathspec", specifier = ">=1.0.4" }, @@ -2252,7 +2252,7 @@ wheels = [ [[package]] name = "lancedb" -version = "0.30.2" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "deprecation" }, @@ -2264,12 +2264,10 @@ dependencies = [ { name = "tqdm" }, ] wheels = [ - { url = "https://files.pythonhosted.org/packages/7f/87/67b23006663be175c396ae8f7c6ac98bfa4728de5b5583016b8b8c54eb14/lancedb-0.30.2-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:3dd8cb9e2e25efb32c088b24b3fbc57f3f24a636f4b8ad4b287b1eb52f6b5075", size = 41720461, upload-time = "2026-03-31T22:42:32.853Z" }, - { url = "https://files.pythonhosted.org/packages/78/68/b3b5f638f8de91de75751414114690cae9c294dc79d9ab2602f4562ed9df/lancedb-0.30.2-cp39-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:f083d50b257f645bd5c4b295d693648ffb37640ce1e9d72f55041b1382f0dbd6", size = 43626135, upload-time = "2026-03-31T22:50:28.577Z" }, - { url = "https://files.pythonhosted.org/packages/ef/d1/ea8b74a8b56dd4925cc9cb9cc23c7d9675708a7f6b33d22136dc7bb34dbc/lancedb-0.30.2-cp39-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3aef5538db9cd82af79c90831035b4d67e9aa182ef73095a1b919caddf9bb7a5", size = 46619289, upload-time = "2026-03-31T22:55:02.242Z" }, - { url = "https://files.pythonhosted.org/packages/74/4b/5bfeacf948cfc3452b286a792dcbbfaf04649ef0820e1d3790d47bf5527e/lancedb-0.30.2-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:8b161cb1da04ae6ad45afe10093cfe4107821d93e7712b50200c435d6f4c8a20", size = 43641193, upload-time = "2026-03-31T22:51:13.63Z" }, - { url = "https://files.pythonhosted.org/packages/28/4c/a51af0ce1d18fd86afa3e8538a81abf5523d24632abe7665ce6795b8009d/lancedb-0.30.2-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:7fabc0f57944fd79ddef62ed8cf4df770654b172b1ad1019a999304fed3169f3", size = 46665361, upload-time = "2026-03-31T22:54:20.282Z" }, - { url = "https://files.pythonhosted.org/packages/88/d0/7e44e8143ac2dae8979ba882cc33d4af7b8da4741fb0361497e69b4a4379/lancedb-0.30.2-cp39-abi3-win_amd64.whl", hash = "sha256:531da53002c1c6fda829afccc8ced3056ef58eb036f09ddb2b94a06877ecc66c", size = 50940681, upload-time = "2026-03-31T23:25:52.35Z" }, + { url = "https://files.pythonhosted.org/packages/df/f7/5262b9aa593f790757163c0165ab0da1dda054758901bea7e4f02c9cb633/lancedb-0.34.0-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:c462f2e6f933cad659fd0179394eaab578acbc9151fe2ef41bc29b36ecca5058", size = 52654213, upload-time = "2026-07-02T17:13:31.102Z" }, + { url = "https://files.pythonhosted.org/packages/69/99/05ea0d32229ebea695193ff20c15d6ecae25785ad82a9d4723d98832a284/lancedb-0.34.0-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:48829e88e708947d0520454ab9e4f8efa35f3e3626469eadd3a6e061b89cb223", size = 55434501, upload-time = "2026-07-02T17:13:34.81Z" }, + { url = "https://files.pythonhosted.org/packages/cd/4e/4325c13d5afa93c466428a5a0f168ad4d96f5eb4a77bbe7c5100d39c9897/lancedb-0.34.0-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:05ba8a5b58e064edfbe5be71b1abf2e411b4eaf295d1a173dcb1a55c5bfb5285", size = 58659359, upload-time = "2026-07-02T17:13:38.424Z" }, + { url = "https://files.pythonhosted.org/packages/d9/5d/8ca165f1386caf6c4d1c515afd52f345b66432264eecfdfb7fd33eefd9af/lancedb-0.34.0-cp39-abi3-win_amd64.whl", hash = "sha256:51cbc11808f9e3332819b9367c975b3a888541447a8e7bea09c57c852a279153", size = 63530726, upload-time = "2026-07-02T17:13:41.612Z" }, ] [[package]]