Merge pull request #505 from ggozad/feat/version-tags-branches
Database tags with restore; remove --before time travel
This commit is contained in:
commit
aa7a4ff91a
29 changed files with 1986 additions and 518 deletions
12
CHANGELOG.md
12
CHANGELOG.md
|
|
@ -1,6 +1,18 @@
|
||||||
# Changelog
|
# Changelog
|
||||||
## [Unreleased]
|
## [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
|
### 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.
|
- `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.
|
||||||
|
|
|
||||||
|
|
@ -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.)
|
- **MCP server** — Expose as tools for AI assistants (Claude Desktop, etc.)
|
||||||
- **Visual grounding** — View chunks highlighted on original page images
|
- **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).
|
- **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
|
- **Inspector** — TUI for browsing documents, chunks, and search results
|
||||||
|
|
||||||
## Installation
|
## Installation
|
||||||
|
|
|
||||||
|
|
@ -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.
|
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.
|
||||||
|
|
|
||||||
68
docs/cli.md
68
docs/cli.md
|
|
@ -7,7 +7,6 @@ The `haiku-rag` CLI provides complete document management functionality.
|
||||||
|
|
||||||
- `--config` - Specify custom configuration file
|
- `--config` - Specify custom configuration file
|
||||||
- `--read-only` - Open database in read-only mode (blocks writes, skips upgrades)
|
- `--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
|
- `--version` / `-v` - Show version and exit
|
||||||
|
|
||||||
Per-command options:
|
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
|
||||||
haiku-rag --config /path/to/config.yaml list --db /path/to/custom.db
|
haiku-rag --config /path/to/config.yaml list --db /path/to/custom.db
|
||||||
haiku-rag --read-only search "query"
|
haiku-rag --read-only search "query"
|
||||||
haiku-rag --before "2025-01-15" search "query"
|
|
||||||
haiku-rag add -h
|
haiku-rag add -h
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -542,35 +540,53 @@ haiku-skills chat --use-entrypoints --skill medic
|
||||||
└── haiku.rag.yaml # Optional config
|
└── 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:
|
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.
|
||||||
|
|
||||||
- **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:
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Query documents as of January 15, 2025
|
# Tag the current state, e.g. at deploy time or after an ingestion run
|
||||||
haiku-rag --before "2025-01-15" list
|
haiku-rag tag create release-1
|
||||||
|
|
||||||
# Search historical state
|
# List tags with the versions they point to
|
||||||
haiku-rag --before "2025-01-15T14:30:00" search "machine learning"
|
haiku-rag tag list
|
||||||
|
|
||||||
# Ask questions against historical data
|
# Delete a tag, releasing its versions for cleanup
|
||||||
haiku-rag --before "2025-01-15" ask "What documents existed?"
|
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`
|
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.
|
||||||
- Date only: `2025-01-15` (interpreted as start of day)
|
|
||||||
|
|
||||||
!!! note
|
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.
|
||||||
Time travel mode automatically enables read-only mode. You cannot modify the database while viewing historical state.
|
|
||||||
|
### 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-<timestamp>`) 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
|
### Version History
|
||||||
|
|
||||||
|
|
@ -587,20 +603,18 @@ haiku-rag history --table documents
|
||||||
haiku-rag history --limit 10
|
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
|
Version History
|
||||||
|
|
||||||
documents
|
documents
|
||||||
v5: 2025-01-15 14:30:00
|
v5: 2025-01-15 14:30:00 <- release-1
|
||||||
v4: 2025-01-14 10:00:00
|
v4: 2025-01-14 10:00:00
|
||||||
v3: 2025-01-13 09:15:00
|
v3: 2025-01-13 09:15:00
|
||||||
|
|
||||||
chunks
|
chunks
|
||||||
v8: 2025-01-15 14:30:00
|
v8: 2025-01-15 14:30:00 <- release-1
|
||||||
v7: 2025-01-14 10:00:00
|
v7: 2025-01-14 10:00:00
|
||||||
...
|
...
|
||||||
```
|
```
|
||||||
|
|
||||||
Use the timestamps from `history` to construct `--before` queries.
|
|
||||||
|
|
|
||||||
|
|
@ -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.
|
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
|
### Rebuilding the Database
|
||||||
|
|
||||||
```python
|
```python
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,10 @@
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.markdown import Markdown
|
from rich.markdown import Markdown
|
||||||
|
from rich.markup import escape
|
||||||
from rich.progress import (
|
from rich.progress import (
|
||||||
BarColumn,
|
BarColumn,
|
||||||
DownloadColumn,
|
DownloadColumn,
|
||||||
|
|
@ -22,6 +22,7 @@ from haiku.rag.store.models.chunk import SearchType
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.store.engine import Store
|
||||||
from haiku.rag.store.models import SearchResult
|
from haiku.rag.store.models import SearchResult
|
||||||
from haiku.rag.utils import format_bytes, format_citations_rich
|
from haiku.rag.utils import format_bytes, format_citations_rich
|
||||||
|
|
||||||
|
|
@ -34,12 +35,10 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path: Path,
|
db_path: Path,
|
||||||
config: AppConfig = Config,
|
config: AppConfig = Config,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
):
|
):
|
||||||
self.db_path = db_path
|
self.db_path = db_path
|
||||||
self.config = config
|
self.config = config
|
||||||
self.read_only = read_only
|
self.read_only = read_only
|
||||||
self.before = before
|
|
||||||
self.console = Console()
|
self.console = Console()
|
||||||
|
|
||||||
from haiku.rag.store.engine import ConnectionMode
|
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
|
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
|
# Basic: show path/URI
|
||||||
self.console.print("[bold]haiku.rag database info[/bold]")
|
self.console.print("[bold]haiku.rag database info[/bold]")
|
||||||
self.console.print(
|
self.console.print(
|
||||||
|
|
@ -288,7 +282,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
skip_migration_check=True,
|
skip_migration_check=True,
|
||||||
before=self.before,
|
|
||||||
) as store:
|
) as store:
|
||||||
tables = [
|
tables = [
|
||||||
"documents",
|
"documents",
|
||||||
|
|
@ -307,6 +300,14 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
|
|
||||||
self.console.print("[bold]Version History[/bold]")
|
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:
|
for table_name in tables:
|
||||||
versions = await store.list_table_versions(table_name)
|
versions = await store.list_table_versions(table_name)
|
||||||
|
|
||||||
|
|
@ -316,6 +317,12 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
if limit:
|
if limit:
|
||||||
versions = versions[: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]")
|
self.console.print(f"\n[bold cyan]{table_name}[/bold cyan]")
|
||||||
|
|
||||||
if not versions:
|
if not versions:
|
||||||
|
|
@ -325,16 +332,105 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
for v in versions:
|
for v in versions:
|
||||||
version_num = v["version"]
|
version_num = v["version"]
|
||||||
timestamp = v["timestamp"]
|
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(
|
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 def list_documents(self, filter: str | None = None):
|
||||||
async with HaikuRAG(
|
async with HaikuRAG(
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
documents = await self.client.list_documents(filter=filter)
|
documents = await self.client.list_documents(filter=filter)
|
||||||
for doc in documents:
|
for doc in documents:
|
||||||
|
|
@ -347,7 +443,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
doc = await self.client.create_document(
|
doc = await self.client.create_document(
|
||||||
text, title=title, metadata=metadata
|
text, title=title, metadata=metadata
|
||||||
|
|
@ -364,7 +459,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
result = await self.client.create_document_from_source(
|
result = await self.client.create_document_from_source(
|
||||||
source, title=title, metadata=metadata
|
source, title=title, metadata=metadata
|
||||||
|
|
@ -386,7 +480,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
doc = await self.client.get_document_by_id(doc_id)
|
doc = await self.client.get_document_by_id(doc_id)
|
||||||
if doc is None:
|
if doc is None:
|
||||||
|
|
@ -399,7 +492,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
deleted = await self.client.delete_document(doc_id)
|
deleted = await self.client.delete_document(doc_id)
|
||||||
if deleted:
|
if deleted:
|
||||||
|
|
@ -443,7 +535,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
results = await self.client.search(
|
results = await self.client.search(
|
||||||
search_input,
|
search_input,
|
||||||
|
|
@ -465,7 +556,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
chunk = await self.client.get_chunk_by_id(chunk_id)
|
chunk = await self.client.get_chunk_by_id(chunk_id)
|
||||||
if not chunk:
|
if not chunk:
|
||||||
|
|
@ -509,7 +599,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
answer, citations = await self.client.ask(question, filter=filter)
|
answer, citations = await self.client.ask(question, filter=filter)
|
||||||
|
|
||||||
|
|
@ -537,7 +626,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=self.before,
|
|
||||||
) as self.client:
|
) as self.client:
|
||||||
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
||||||
self.console.print()
|
self.console.print()
|
||||||
|
|
@ -561,7 +649,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
config=self.config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as client:
|
) as client:
|
||||||
if mode == RebuildMode.SET_EMBEDDER:
|
if mode == RebuildMode.SET_EMBEDDER:
|
||||||
async for _ in client.rebuild_database(mode=mode):
|
async for _ in client.rebuild_database(mode=mode):
|
||||||
|
|
@ -605,7 +692,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
config=self.config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as client:
|
) as client:
|
||||||
await client.vacuum()
|
await client.vacuum()
|
||||||
self.console.print("[bold green]Vacuum completed successfully.[/bold green]")
|
self.console.print("[bold green]Vacuum completed successfully.[/bold green]")
|
||||||
|
|
@ -634,7 +720,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
config=self.config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
) as client:
|
) as client:
|
||||||
row_count = await client.store.chunks_table.count_rows()
|
row_count = await client.store.chunks_table.count_rows()
|
||||||
self.console.print(f"Chunks in database: {row_count}")
|
self.console.print(f"Chunks in database: {row_count}")
|
||||||
|
|
@ -802,7 +887,6 @@ class HaikuRAGApp: # pragma: no cover
|
||||||
self.db_path,
|
self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
):
|
):
|
||||||
server = create_mcp_server(
|
server = create_mcp_server(
|
||||||
self.db_path, config=self.config, read_only=self.read_only
|
self.db_path, config=self.config, read_only=self.read_only
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,9 @@
|
||||||
from datetime import datetime
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
def run_chat(
|
def run_chat(
|
||||||
db_path: Path | None = None,
|
db_path: Path | None = None,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
model: str | None = None,
|
model: str | None = None,
|
||||||
skills: list[str] | None = None,
|
skills: list[str] | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -14,7 +12,6 @@ def run_chat(
|
||||||
Args:
|
Args:
|
||||||
db_path: Path to the LanceDB database. If None, uses default from config.
|
db_path: Path to the LanceDB database. If None, uses default from config.
|
||||||
read_only: Whether to open the database in read-only mode.
|
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.
|
model: Model to use for the chat.
|
||||||
skills: Skills to enable ("rag", "analysis"). Defaults to ["rag"].
|
skills: Skills to enable ("rag", "analysis"). Defaults to ["rag"].
|
||||||
"""
|
"""
|
||||||
|
|
@ -55,7 +52,6 @@ def run_chat(
|
||||||
db_path,
|
db_path,
|
||||||
skills=skill_list,
|
skills=skill_list,
|
||||||
read_only=read_only,
|
read_only=read_only,
|
||||||
before=before,
|
|
||||||
model=model or get_model(config.qa.model, config),
|
model=model or get_model(config.qa.model, config),
|
||||||
)
|
)
|
||||||
app.run()
|
app.run()
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,6 @@ import asyncio
|
||||||
import json
|
import json
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
from datetime import datetime
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
|
@ -85,14 +84,12 @@ class ChatApp(App):
|
||||||
db_path: Path,
|
db_path: Path,
|
||||||
skills: list[Skill],
|
skills: list[Skill],
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
model: str | None = None,
|
model: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.db_path = db_path
|
self.db_path = db_path
|
||||||
self._skills = skills
|
self._skills = skills
|
||||||
self.read_only = read_only
|
self.read_only = read_only
|
||||||
self.before = before
|
|
||||||
self._model = model
|
self._model = model
|
||||||
self.client: HaikuRAG | None = None
|
self.client: HaikuRAG | None = None
|
||||||
self.config = get_config()
|
self.config = get_config()
|
||||||
|
|
@ -150,7 +147,6 @@ class ChatApp(App):
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
)
|
)
|
||||||
await self.client.__aenter__()
|
await self.client.__aenter__()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,6 @@ import asyncio
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
import warnings
|
import warnings
|
||||||
from datetime import datetime
|
|
||||||
from importlib.metadata import version
|
from importlib.metadata import version
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
@ -47,7 +46,6 @@ def cli():
|
||||||
|
|
||||||
# Module-level flags set by callback
|
# Module-level flags set by callback
|
||||||
_read_only: bool = False
|
_read_only: bool = False
|
||||||
_before: datetime | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def create_app(db: Path | None = None) -> HaikuRAGApp: # pragma: no cover
|
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()
|
config = get_config()
|
||||||
db_path = db if db else config.storage.data_dir / "haiku.rag.lancedb"
|
db_path = db if db else config.storage.data_dir / "haiku.rag.lancedb"
|
||||||
return HaikuRAGApp(
|
return HaikuRAGApp(db_path=db_path, config=config, read_only=_read_only)
|
||||||
db_path=db_path, config=config, read_only=_read_only, before=_before
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def check_version(): # pragma: no cover
|
async def check_version(): # pragma: no cover
|
||||||
|
|
@ -102,28 +98,10 @@ def main(
|
||||||
"--read-only",
|
"--read-only",
|
||||||
help="Open database in read-only mode",
|
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"""
|
"""haiku.rag CLI - Vector database RAG system"""
|
||||||
global _read_only, _before
|
global _read_only
|
||||||
_read_only = 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
|
# Load config from --config, local folder, or default directory
|
||||||
config_path = find_config_file(cli_path=config)
|
config_path = find_config_file(cli_path=config)
|
||||||
if config_path:
|
if config_path:
|
||||||
|
|
@ -635,6 +613,98 @@ def history( # pragma: no cover
|
||||||
asyncio.run(app.history(table=table, limit=limit))
|
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")
|
@_cli.command("download-models", help="Download Docling and Ollama models per config")
|
||||||
def download_models_cmd(): # pragma: no cover
|
def download_models_cmd(): # pragma: no cover
|
||||||
app = HaikuRAGApp(db_path=Path(), config=get_config(), read_only=True)
|
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
|
raise typer.Exit(1) from e
|
||||||
|
|
||||||
db_path = db if db else get_config().storage.data_dir / "haiku.rag.lancedb"
|
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")
|
@_cli.command("chat", help="Launch interactive chat TUI for conversational RAG")
|
||||||
|
|
@ -692,7 +762,6 @@ def chat( # pragma: no cover
|
||||||
run_chat(
|
run_chat(
|
||||||
db_path,
|
db_path,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
before=_before,
|
|
||||||
model=model,
|
model=model,
|
||||||
skills=skills,
|
skills=skills,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,6 @@ import logging
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import tempfile
|
import tempfile
|
||||||
from collections.abc import AsyncGenerator, Sequence
|
from collections.abc import AsyncGenerator, Sequence
|
||||||
from datetime import datetime
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from functools import cached_property
|
from functools import cached_property
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
@ -72,7 +71,6 @@ class HaikuRAG:
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
create: bool = False,
|
create: bool = False,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
):
|
):
|
||||||
"""Initialize the RAG client with a database path.
|
"""Initialize the RAG client with a database path.
|
||||||
|
|
||||||
|
|
@ -82,8 +80,6 @@ class HaikuRAG:
|
||||||
skip_validation: Whether to skip configuration validation on database load.
|
skip_validation: Whether to skip configuration validation on database load.
|
||||||
create: Whether to create the database if it doesn't exist.
|
create: Whether to create the database if it doesn't exist.
|
||||||
read_only: Whether to open the database in read-only mode.
|
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
|
self._config = config
|
||||||
if db_path is None:
|
if db_path is None:
|
||||||
|
|
@ -93,7 +89,6 @@ class HaikuRAG:
|
||||||
self._skip_validation = skip_validation
|
self._skip_validation = skip_validation
|
||||||
self._create = create
|
self._create = create
|
||||||
self._read_only = read_only
|
self._read_only = read_only
|
||||||
self._before = before
|
|
||||||
self._vacuum_tasks: set[asyncio.Task] = set()
|
self._vacuum_tasks: set[asyncio.Task] = set()
|
||||||
self._last_vacuum_at: float | None = None
|
self._last_vacuum_at: float | None = None
|
||||||
self._vacuum_dirty = False
|
self._vacuum_dirty = False
|
||||||
|
|
@ -125,7 +120,6 @@ class HaikuRAG:
|
||||||
skip_validation=self._skip_validation,
|
skip_validation=self._skip_validation,
|
||||||
create=self._create,
|
create=self._create,
|
||||||
read_only=self._read_only,
|
read_only=self._read_only,
|
||||||
before=self._before,
|
|
||||||
)
|
)
|
||||||
# If _initialize fails mid-way (e.g. migration check raises after
|
# If _initialize fails mid-way (e.g. migration check raises after
|
||||||
# connect), close the store so we don't leak the LanceDB connection —
|
# connect), close the store so we don't leak the LanceDB connection —
|
||||||
|
|
|
||||||
|
|
@ -395,7 +395,8 @@ async def _refresh_doc_metadata(
|
||||||
updated = True
|
updated = True
|
||||||
|
|
||||||
if updated:
|
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.
|
# Reclaim the document_meta churn from rolling source_revision sweeps.
|
||||||
# The vacuum is debounced, and document_meta is tiny, so this is cheap.
|
# The vacuum is debounced, and document_meta is tiny, so this is cheap.
|
||||||
if client._config.storage.auto_vacuum:
|
if client._config.storage.auto_vacuum:
|
||||||
|
|
@ -864,7 +865,8 @@ async def update_document(
|
||||||
existing_doc.uri = uri
|
existing_doc.uri = uri
|
||||||
|
|
||||||
if content is None and chunks is None and docling_document is None:
|
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:
|
if client._config.storage.auto_vacuum:
|
||||||
client._schedule_vacuum()
|
client._schedule_vacuum()
|
||||||
return updated
|
return updated
|
||||||
|
|
|
||||||
|
|
@ -71,15 +71,30 @@ async def rebuild_database(
|
||||||
"""Rebuild the database with the specified mode.
|
"""Rebuild the database with the specified mode.
|
||||||
|
|
||||||
Yields the ID of each document as it is processed.
|
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
|
from haiku.rag.client import RebuildMode
|
||||||
|
|
||||||
if mode is None:
|
if mode is None:
|
||||||
mode = RebuildMode.FULL
|
mode = RebuildMode.FULL
|
||||||
|
|
||||||
if mode == RebuildMode.SET_EMBEDDER:
|
async with client.store._rebuild_lock:
|
||||||
await _set_embedder(client)
|
if mode == RebuildMode.SET_EMBEDDER:
|
||||||
return
|
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
|
# Resolve any leftover staging/marker tables from a previously
|
||||||
# interrupted rebuild. Returns True only when phase 1 was already
|
# interrupted rebuild. Returns True only when phase 1 was already
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,3 @@
|
||||||
from datetime import datetime
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
|
@ -67,13 +66,10 @@ class InspectorApp(App):
|
||||||
Binding("c", "show_context", "Context", show=True),
|
Binding("c", "show_context", "Context", show=True),
|
||||||
]
|
]
|
||||||
|
|
||||||
def __init__(
|
def __init__(self, db_path: Path, read_only: bool = False):
|
||||||
self, db_path: Path, read_only: bool = False, before: datetime | None = None
|
|
||||||
):
|
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.db_path = db_path
|
self.db_path = db_path
|
||||||
self.read_only = read_only
|
self.read_only = read_only
|
||||||
self.before = before
|
|
||||||
self.client: HaikuRAG | None = None
|
self.client: HaikuRAG | None = None
|
||||||
|
|
||||||
def compose(self) -> "ComposeResult":
|
def compose(self) -> "ComposeResult":
|
||||||
|
|
@ -91,7 +87,6 @@ class InspectorApp(App):
|
||||||
db_path=self.db_path,
|
db_path=self.db_path,
|
||||||
config=config,
|
config=config,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
before=self.before,
|
|
||||||
)
|
)
|
||||||
await self.client.__aenter__()
|
await self.client.__aenter__()
|
||||||
|
|
||||||
|
|
@ -234,18 +229,16 @@ class InspectorApp(App):
|
||||||
def run_inspector(
|
def run_inspector(
|
||||||
db_path: Path | None = None,
|
db_path: Path | None = None,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Run the inspector TUI.
|
"""Run the inspector TUI.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
db_path: Path to the LanceDB database. If None, uses default from config.
|
db_path: Path to the LanceDB database. If None, uses default from config.
|
||||||
read_only: Whether to open the database in read-only mode.
|
read_only: Whether to open the database in read-only mode.
|
||||||
before: Query database as it existed before this datetime.
|
|
||||||
"""
|
"""
|
||||||
config = get_config()
|
config = get_config()
|
||||||
if db_path is None:
|
if db_path is None:
|
||||||
db_path = config.storage.data_dir / "haiku.rag.lancedb"
|
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()
|
app.run()
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,9 @@
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
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 enum import Enum
|
||||||
from importlib import metadata
|
from importlib import metadata
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
@ -187,6 +189,68 @@ REQUIRED_TABLES: tuple[str, ...] = (
|
||||||
"settings",
|
"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:
|
async def get_database_stats(db: lancedb.AsyncConnection) -> dict:
|
||||||
"""Collect stats for every haiku.rag table on the connection.
|
"""Collect stats for every haiku.rag table on the connection.
|
||||||
|
|
@ -347,19 +411,19 @@ class Store:
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
create: bool = False,
|
create: bool = False,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
before: datetime | None = None,
|
|
||||||
skip_migration_check: bool = False,
|
skip_migration_check: bool = False,
|
||||||
):
|
):
|
||||||
self.db_path: Path = db_path
|
self.db_path: Path = db_path
|
||||||
self._config = config
|
self._config = config
|
||||||
self._before = before
|
self._read_only = read_only
|
||||||
# Time-travel mode is always read-only
|
|
||||||
self._read_only = read_only or (before is not None)
|
|
||||||
self._create = create
|
self._create = create
|
||||||
self._skip_validation = skip_validation
|
self._skip_validation = skip_validation
|
||||||
self._skip_migration_check = skip_migration_check
|
self._skip_migration_check = skip_migration_check
|
||||||
self._vacuum_lock = asyncio.Lock()
|
self._vacuum_lock = asyncio.Lock()
|
||||||
self._write_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
|
self._is_new_db = False
|
||||||
|
|
||||||
# Check if database exists (for local filesystem only)
|
# Check if database exists (for local filesystem only)
|
||||||
|
|
@ -409,10 +473,6 @@ class Store:
|
||||||
# pending, before creating any newly-introduced table.
|
# pending, before creating any newly-introduced table.
|
||||||
await self._init_tables(is_new_db)
|
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.
|
# Set version for new databases.
|
||||||
if is_new_db and not self._read_only:
|
if is_new_db and not self._read_only:
|
||||||
await self._set_initial_version()
|
await self._set_initial_version()
|
||||||
|
|
@ -474,6 +534,13 @@ class Store:
|
||||||
if self._read_only:
|
if self._read_only:
|
||||||
raise ReadOnlyError("Cannot modify database in read-only mode")
|
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:
|
async def vacuum(self, retention_seconds: int | None = None) -> None:
|
||||||
"""Optimize and clean up old versions across all tables to reduce disk usage.
|
"""Optimize and clean up old versions across all tables to reduce disk usage.
|
||||||
|
|
||||||
|
|
@ -487,6 +554,8 @@ class Store:
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ReadOnlyError: If the store is in read-only mode.
|
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()
|
self._assert_writable()
|
||||||
|
|
||||||
|
|
@ -497,25 +566,54 @@ class Store:
|
||||||
if self._vacuum_lock.locked():
|
if self._vacuum_lock.locked():
|
||||||
return
|
return
|
||||||
|
|
||||||
async with self._vacuum_lock:
|
async with self._vacuum_lock, self._write_lock:
|
||||||
try:
|
try:
|
||||||
# Evaluate config at runtime to allow dynamic changes
|
# Evaluate config at runtime to allow dynamic changes
|
||||||
if retention_seconds is None:
|
if retention_seconds is None:
|
||||||
retention_seconds = self._config.storage.vacuum_retention_seconds
|
retention_seconds = self._config.storage.vacuum_retention_seconds
|
||||||
# Perform maintenance per table using optimize() with configurable retention
|
# Perform maintenance per table using optimize() with configurable retention
|
||||||
retention = timedelta(seconds=retention_seconds)
|
retention = timedelta(seconds=retention_seconds)
|
||||||
for table in [
|
for table in self._tables().values():
|
||||||
self.documents_table,
|
await table.optimize(
|
||||||
self.document_meta_table,
|
cleanup_older_than=await self._tag_safe_retention(
|
||||||
self.chunks_table,
|
table, retention
|
||||||
self.document_items_table,
|
)
|
||||||
self.settings_table,
|
)
|
||||||
]:
|
except OSError as e:
|
||||||
await table.optimize(cleanup_older_than=retention)
|
# Resource errors (e.g. disk pressure) skip the pass; lance
|
||||||
except (RuntimeError, OSError) as e:
|
# errors surface as RuntimeError and must not be swallowed —
|
||||||
# Handle resource errors gracefully
|
# a silently skipped cleanup hides tag-interaction bugs.
|
||||||
logger.debug(f"Vacuum skipped due to resource constraints: {e}")
|
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
|
@property
|
||||||
def _connection_mode(self) -> ConnectionMode:
|
def _connection_mode(self) -> ConnectionMode:
|
||||||
return ConnectionMode.from_config(self._config)
|
return ConnectionMode.from_config(self._config)
|
||||||
|
|
@ -788,15 +886,19 @@ class Store:
|
||||||
if hasattr(self, "db"):
|
if hasattr(self, "db"):
|
||||||
self.db.close()
|
self.db.close()
|
||||||
|
|
||||||
|
def _tables(self) -> dict[str, lancedb.AsyncTable]:
|
||||||
|
"""Map every haiku.rag table name to its open AsyncTable."""
|
||||||
|
return {
|
||||||
|
"documents": self.documents_table,
|
||||||
|
"document_meta": self.document_meta_table,
|
||||||
|
"chunks": self.chunks_table,
|
||||||
|
"document_items": self.document_items_table,
|
||||||
|
"settings": self.settings_table,
|
||||||
|
}
|
||||||
|
|
||||||
async def current_table_versions(self) -> dict[str, int]:
|
async def current_table_versions(self) -> dict[str, int]:
|
||||||
"""Capture current versions of key tables for rollback using LanceDB's API."""
|
"""Capture current versions of key tables for rollback using LanceDB's API."""
|
||||||
return {
|
return {name: await table.version() for name, table in self._tables().items()}
|
||||||
"documents": await self.documents_table.version(),
|
|
||||||
"document_meta": await self.document_meta_table.version(),
|
|
||||||
"chunks": await self.chunks_table.version(),
|
|
||||||
"document_items": await self.document_items_table.version(),
|
|
||||||
"settings": await self.settings_table.version(),
|
|
||||||
}
|
|
||||||
|
|
||||||
async def restore_table_versions(self, versions: dict[str, int]) -> bool:
|
async def restore_table_versions(self, versions: dict[str, int]) -> bool:
|
||||||
"""Restore tables to the provided versions using LanceDB's API.
|
"""Restore tables to the provided versions using LanceDB's API.
|
||||||
|
|
@ -805,74 +907,243 @@ class Store:
|
||||||
ReadOnlyError: If the store is in read-only mode.
|
ReadOnlyError: If the store is in read-only mode.
|
||||||
"""
|
"""
|
||||||
self._assert_writable()
|
self._assert_writable()
|
||||||
await self.documents_table.restore(int(versions["documents"]))
|
for name, table in self._tables().items():
|
||||||
await self.document_meta_table.restore(int(versions["document_meta"]))
|
await table.restore(int(versions[name]))
|
||||||
await self.chunks_table.restore(int(versions["chunks"]))
|
|
||||||
await self.document_items_table.restore(int(versions["document_items"]))
|
|
||||||
await self.settings_table.restore(int(versions["settings"]))
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def _checkout_tables_before(self, before: datetime) -> None:
|
async def create_tag(self, name: str) -> None:
|
||||||
"""Checkout all tables to their state at or before the given datetime.
|
"""Tag the current version of every table with the given name.
|
||||||
|
|
||||||
Args:
|
Serializes with client writes via the write lock so a write cannot
|
||||||
before: The datetime to checkout to
|
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:
|
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.
|
self._assert_writable()
|
||||||
# Convert 'before' to naive local time for comparison.
|
self._assert_not_rebuilding()
|
||||||
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
|
|
||||||
|
|
||||||
tables = [
|
async with self._rebuild_lock, self._write_lock:
|
||||||
("documents", self.documents_table),
|
await self._create_tag_locked(name)
|
||||||
("document_meta", self.document_meta_table),
|
|
||||||
("chunks", self.chunks_table),
|
async def _create_tag_locked(self, name: str) -> None:
|
||||||
("document_items", self.document_items_table),
|
"""Create a tag on every table; the caller must hold the write lock."""
|
||||||
("settings", self.settings_table),
|
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 self.current_table_versions()
|
||||||
versions = await table.list_versions()
|
try:
|
||||||
# Find the latest version at or before the target datetime
|
for table_name, table in tables.items():
|
||||||
# Versions are sorted by version number, not timestamp, so we need to check all
|
await table.tags.create(name, versions[table_name])
|
||||||
best_version = None
|
except BaseException as exc:
|
||||||
best_timestamp = None
|
# 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:
|
async def _delete_tag_locked(self, name: str) -> tuple[bool, list[str]]:
|
||||||
# LanceDB version timestamps are naive datetime objects in local time
|
"""Delete the tag from every table that has it; the caller must
|
||||||
v_timestamp = v["timestamp"]
|
hold the write lock.
|
||||||
# Make sure it's naive for comparison
|
|
||||||
if v_timestamp.tzinfo is not None:
|
|
||||||
v_timestamp = v_timestamp.replace(tzinfo=None)
|
|
||||||
|
|
||||||
if v_timestamp <= before_local:
|
Returns whether the tag was found anywhere and the tables where
|
||||||
if best_timestamp is None or v_timestamp > best_timestamp:
|
listing or deletion failed.
|
||||||
best_version = v["version"]
|
"""
|
||||||
best_timestamp = v_timestamp
|
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:
|
async def list_tags(self) -> dict[str, TagInfo]:
|
||||||
# Find the earliest version to report in error message
|
"""Aggregate per-table tags into database-level tags.
|
||||||
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."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Checkout to the found version
|
Returns:
|
||||||
await table.checkout(best_version)
|
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]]:
|
async def list_table_versions(self, table_name: str) -> list[dict[str, Any]]:
|
||||||
"""List version history for a table.
|
"""List version history for a table.
|
||||||
|
|
@ -884,14 +1155,7 @@ class Store:
|
||||||
Returns:
|
Returns:
|
||||||
List of version info dicts with "version" and "timestamp" keys
|
List of version info dicts with "version" and "timestamp" keys
|
||||||
"""
|
"""
|
||||||
table_map = {
|
table = self._tables().get(table_name)
|
||||||
"documents": self.documents_table,
|
|
||||||
"document_meta": self.document_meta_table,
|
|
||||||
"chunks": self.chunks_table,
|
|
||||||
"document_items": self.document_items_table,
|
|
||||||
"settings": self.settings_table,
|
|
||||||
}
|
|
||||||
table = table_map.get(table_name)
|
|
||||||
if table is None:
|
if table is None:
|
||||||
raise ValueError(f"Unknown table: {table_name}")
|
raise ValueError(f"Unknown table: {table_name}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -157,10 +157,12 @@ async def _apply_extract_picture_bytes(store: Store) -> None:
|
||||||
schema=_V0_45_0_ITEMS_SCHEMA,
|
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 (
|
await (
|
||||||
store.document_items_table.merge_insert(["document_id", "self_ref"])
|
store.document_items_table.merge_insert(["document_id", "self_ref"])
|
||||||
.when_matched_update_all()
|
.when_matched_update_all()
|
||||||
.when_not_matched_insert_all()
|
|
||||||
.execute(new_records)
|
.execute(new_records)
|
||||||
)
|
)
|
||||||
wrote_items = True
|
wrote_items = True
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,9 @@
|
||||||
import math
|
import math
|
||||||
import sys
|
import sys
|
||||||
from datetime import UTC, datetime
|
|
||||||
from importlib import metadata
|
from importlib import metadata
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, cast
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
|
|
||||||
from dateutil import parser as dateutil_parser
|
|
||||||
from packaging.version import Version, parse
|
from packaging.version import Version, parse
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|
@ -38,55 +36,6 @@ def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||||
return dot_product / (norm1 * norm2)
|
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(
|
def apply_common_settings(
|
||||||
settings: Any | None,
|
settings: Any | None,
|
||||||
settings_class: type[Any],
|
settings_class: type[Any],
|
||||||
|
|
|
||||||
|
|
@ -28,7 +28,7 @@ dependencies = [
|
||||||
"jinja2>=3.1.0",
|
"jinja2>=3.1.0",
|
||||||
"jsonpatch>=1.33",
|
"jsonpatch>=1.33",
|
||||||
"fastmcp>=3.3.0",
|
"fastmcp>=3.3.0",
|
||||||
"lancedb==0.30.2",
|
"lancedb==0.34.0",
|
||||||
"pathspec>=1.0.4",
|
"pathspec>=1.0.4",
|
||||||
"pydantic>=2.12.5",
|
"pydantic>=2.12.5",
|
||||||
"pydantic-ai-slim[openai,logfire,ag-ui]>=1.100.0",
|
"pydantic-ai-slim[openai,logfire,ag-ui]>=1.100.0",
|
||||||
|
|
|
||||||
|
|
@ -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
|
"""The full round-trip: ingest a file, build a sync_state-shaped snapshot
|
||||||
from document.metadata, hand it to FSSource.discover() — must see
|
from document.metadata, hand it to FSSource.discover() — must see
|
||||||
UNCHANGED, not UPSERT. This is exactly what the periodic poller does."""
|
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")
|
file_path.write_text("hello")
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
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
|
assert doc.uri is not None
|
||||||
snapshot = {doc.uri: doc.metadata["source_revision"]}
|
snapshot = {doc.uri: doc.metadata["source_revision"]}
|
||||||
|
|
||||||
src = FSSource(root=tmp_path)
|
src = FSSource(root=docs_dir)
|
||||||
kinds: list[SourceEventKind] = []
|
kinds: list[SourceEventKind] = []
|
||||||
async for event in src.discover(since=snapshot):
|
async for event in src.discover(since=snapshot):
|
||||||
kinds.append(event.kind)
|
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):
|
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
|
"""Counterpart to the unchanged test: a file modified after ingest still
|
||||||
triggers UPSERT. Ensures the round-trip doesn't accidentally over-skip."""
|
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")
|
file_path.write_text("hello")
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
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.
|
# on any sane filesystem, but assert anyway to make the intent explicit.
|
||||||
assert str(file_path.stat().st_mtime_ns) != doc.metadata["source_revision"]
|
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] = []
|
kinds: list[SourceEventKind] = []
|
||||||
async for event in src.discover(since=snapshot):
|
async for event in src.discover(since=snapshot):
|
||||||
kinds.append(event.kind)
|
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
|
"""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
|
fetch context to each child, so the provider sees the configured source id
|
||||||
rather than an ad-hoc fs: identity."""
|
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] = []
|
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)
|
seen_source_ids.append(source_id)
|
||||||
return {"collection": 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:
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
docs = await client.create_document_from_source(
|
docs = await client.create_document_from_source(
|
||||||
tmp_path,
|
docs_dir,
|
||||||
sources=[source],
|
sources=[source],
|
||||||
source_id="docs",
|
source_id="docs",
|
||||||
metadata_provider=Provider(),
|
metadata_provider=Provider(),
|
||||||
|
|
|
||||||
408
tests/store/test_restore.py
Normal file
408
tests/store/test_restore.py
Normal file
|
|
@ -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
|
||||||
534
tests/store/test_tags.py
Normal file
534
tests/store/test_tags.py
Normal file
|
|
@ -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() == {}
|
||||||
|
|
@ -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
|
|
||||||
|
|
@ -84,3 +84,190 @@ class TestCliMigrationError:
|
||||||
with pytest.raises(SystemExit) as exc_info:
|
with pytest.raises(SystemExit) as exc_info:
|
||||||
cli_wrapper()
|
cli_wrapper()
|
||||||
assert exc_info.value.code == 1
|
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
|
||||||
|
|
|
||||||
|
|
@ -2264,3 +2264,30 @@ async def test_rebuild_rechunk_with_url_prefixed_stored_content(
|
||||||
assert doc_after is not None
|
assert doc_after is not None
|
||||||
assert "example.com" in doc_after.content
|
assert "example.com" in doc_after.content
|
||||||
assert "Stored" 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"}
|
||||||
|
|
|
||||||
|
|
@ -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:
|
with patch("haiku.rag.store.engine.Store") as mock_store_cls:
|
||||||
mock_store = AsyncMock()
|
mock_store = AsyncMock()
|
||||||
mock_store.list_table_versions = AsyncMock(return_value=[])
|
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.__aenter__ = AsyncMock(return_value=mock_store)
|
||||||
mock_store_cls.return_value.__aexit__ = AsyncMock(return_value=False)
|
mock_store_cls.return_value.__aexit__ = AsyncMock(return_value=False)
|
||||||
await app.history()
|
await app.history()
|
||||||
mock_store_cls.assert_called_once()
|
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
|
||||||
|
|
|
||||||
|
|
@ -1049,3 +1049,49 @@ async def test_rebuild_set_embedder_raises_on_vector_dim_mismatch(temp_db_path):
|
||||||
with pytest.raises(ConfigMismatchError):
|
with pytest.raises(ConfigMismatchError):
|
||||||
async for _ in client.rebuild_database(mode=RebuildMode.SET_EMBEDDER):
|
async for _ in client.rebuild_database(mode=RebuildMode.SET_EMBEDDER):
|
||||||
pass
|
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"}
|
||||||
|
|
|
||||||
|
|
@ -382,79 +382,6 @@ def test_get_package_versions():
|
||||||
assert len(value) > 0
|
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 ---
|
# --- apply_common_settings tests ---
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -92,3 +92,33 @@ async def test_metadata_refresh_sweep_schedules_vacuum(temp_db_path):
|
||||||
source_metadata={"source_revision": "r2", "md5": "same"},
|
source_metadata={"source_revision": "r2", "md5": "same"},
|
||||||
)
|
)
|
||||||
assert client._vacuum_dirty is True
|
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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
||||||
14
uv.lock
14
uv.lock
|
|
@ -1755,7 +1755,7 @@ requires-dist = [
|
||||||
{ name = "httpx", specifier = ">=0.28.1" },
|
{ name = "httpx", specifier = ">=0.28.1" },
|
||||||
{ name = "jinja2", specifier = ">=3.1.0" },
|
{ name = "jinja2", specifier = ">=3.1.0" },
|
||||||
{ name = "jsonpatch", specifier = ">=1.33" },
|
{ 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 = "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 = "opencv-python-headless", marker = "extra == 'docling'", specifier = ">=4.6.0.66,<5.0.0.0" },
|
||||||
{ name = "pathspec", specifier = ">=1.0.4" },
|
{ name = "pathspec", specifier = ">=1.0.4" },
|
||||||
|
|
@ -2252,7 +2252,7 @@ wheels = [
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "lancedb"
|
name = "lancedb"
|
||||||
version = "0.30.2"
|
version = "0.34.0"
|
||||||
source = { registry = "https://pypi.org/simple" }
|
source = { registry = "https://pypi.org/simple" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "deprecation" },
|
{ name = "deprecation" },
|
||||||
|
|
@ -2264,12 +2264,10 @@ dependencies = [
|
||||||
{ name = "tqdm" },
|
{ name = "tqdm" },
|
||||||
]
|
]
|
||||||
wheels = [
|
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/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/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/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/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/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/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/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" },
|
||||||
{ 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" },
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue