Merge pull request #598 from ggozad/chore/db-entry-refactor
Replace lancedb.uri with lancedb.databases; name every database; remove HAIKU_RAG_DB and DB_PATH
This commit is contained in:
commit
75547f1cbd
71 changed files with 1287 additions and 862 deletions
35
CHANGELOG.md
35
CHANGELOG.md
|
|
@ -2,6 +2,18 @@
|
||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
### Removed
|
||||||
|
|
||||||
|
- `lancedb.uri`. Write `lancedb.databases: {NAME: <location>}`; a config carrying
|
||||||
|
`uri` fails to load with that message. Configurations generated by
|
||||||
|
`init-config` through 0.81 carry `uri: ""` and must drop the key.
|
||||||
|
- `HAIKU_RAG_DB`. Capabilities cover the databases the configuration places, or
|
||||||
|
the `db_path` argument.
|
||||||
|
- `DB_PATH` in the `app/` backend and `examples/custom_agent_agui.py`. Both load
|
||||||
|
the configuration as the CLI does (`HAIKU_RAG_CONFIG_PATH`, `./haiku.rag.yaml`,
|
||||||
|
the platform directory); the compose files set
|
||||||
|
`HAIKU_RAG_CONFIG_PATH=/app/haiku.rag.yaml` and mount `DB_VOLUME` at `/data`.
|
||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
- `qa.max_searches` counts search units: searches a model emits in one
|
- `qa.max_searches` counts search units: searches a model emits in one
|
||||||
|
|
@ -10,6 +22,29 @@
|
||||||
- Searches in one model response deduplicate their results: evidence a sibling
|
- Searches in one model response deduplicate their results: evidence a sibling
|
||||||
search already showed collapses to a reference line, and a picture attaches
|
search already showed collapses to a reference line, and a picture attaches
|
||||||
once per response.
|
once per response.
|
||||||
|
- Every database has a name: the key in `lancedb.databases`, or the path's stem
|
||||||
|
for `--db PATH`, `db_path=` and the default database, which is the entry
|
||||||
|
`haiku.rag` under `storage.data_dir` and selectable by that name.
|
||||||
|
`SearchResult.source`, `Document.source` and `Citation.source` carry it on
|
||||||
|
every value a database produces.
|
||||||
|
- `db_path=` beside a configured `lancedb.databases` raises
|
||||||
|
`AmbiguousDatabaseError` (`HaikuRAG`, `create_capability`, `create_mcp_server`,
|
||||||
|
`Sandbox`). `haiku-rag --db PATH` and `haiku-ingester --db PATH` open that path
|
||||||
|
whatever is configured.
|
||||||
|
- `DatabaseRef(name, location, given)` replaces `DatabaseRef(name, uri, db_path)`;
|
||||||
|
`DatabaseScope.at(path)` added; `locate_database` returns `Path | str`.
|
||||||
|
`IngesterApp(config, scope)` takes a resolved scope in place of `db_path`.
|
||||||
|
`SingleDatabaseSession(ref, config)` replaces `SingleDatabaseSession(db_path,
|
||||||
|
config, source=)`. `Store.db_path` is `None` for a database behind a URI.
|
||||||
|
- Opening a configured or default database that does not exist raises
|
||||||
|
`SourceUnavailableError` naming the database and the remedy (`haiku-rag init`
|
||||||
|
or `create=True`), where the default database raised `FileNotFoundError` with
|
||||||
|
its path. A database given as a path still raises `FileNotFoundError`.
|
||||||
|
- `Store(location, config)`, `connect_lancedb(location, config)`,
|
||||||
|
`gather_database_info(location, config)` and `run_doctor(config, location, ...)`
|
||||||
|
take the database location, a path or a URI. `ConnectionMode.of(location)`
|
||||||
|
replaces `ConnectionMode.from_config`. `DatabaseRef.connection()` and
|
||||||
|
`default_db_path` removed.
|
||||||
|
|
||||||
## [0.81.0] - 2026-09-01
|
## [0.81.0] - 2026-09-01
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,9 @@
|
||||||
ANTHROPIC_API_KEY=your-anthropic-key
|
ANTHROPIC_API_KEY=your-anthropic-key
|
||||||
OPENAI_API_KEY=your-openai-key
|
OPENAI_API_KEY=your-openai-key
|
||||||
|
|
||||||
# Database path
|
# Host path of the LanceDB database, mounted at /data where haiku.rag.yaml
|
||||||
DB_PATH=/path/to/your/haiku.rag.lancedb
|
# places it
|
||||||
|
DB_VOLUME=./data/haiku.rag.lancedb
|
||||||
|
|
||||||
# Optional: Ollama base URL (if using local models)
|
# Optional: Ollama base URL (if using local models)
|
||||||
# Use host.docker.internal to reach Ollama running on the host machine
|
# Use host.docker.internal to reach Ollama running on the host machine
|
||||||
|
|
|
||||||
|
|
@ -40,7 +40,8 @@ A conversational RAG interface built with [CopilotKit](https://copilotkit.ai/) a
|
||||||
|
|
||||||
| Variable | Description | Required |
|
| Variable | Description | Required |
|
||||||
|----------|-------------|----------|
|
|----------|-------------|----------|
|
||||||
| `DB_PATH` | Path to your haiku.rag LanceDB database | Yes |
|
| `DB_VOLUME` | Host path of the LanceDB database the compose files mount at `/data`, where `haiku.rag.yaml` places it (default `./data/haiku.rag.lancedb`) | No |
|
||||||
|
| `HAIKU_RAG_CONFIG_PATH` | The configuration file; the compose files set it to the mounted `/app/haiku.rag.yaml` | No |
|
||||||
| `ANTHROPIC_API_KEY` | Anthropic API key | One LLM key required |
|
| `ANTHROPIC_API_KEY` | Anthropic API key | One LLM key required |
|
||||||
| `OPENAI_API_KEY` | OpenAI API key | One LLM key required |
|
| `OPENAI_API_KEY` | OpenAI API key | One LLM key required |
|
||||||
| `OLLAMA_BASE_URL` | Ollama server URL (default: `http://host.docker.internal:11434`) | For local models |
|
| `OLLAMA_BASE_URL` | Ollama server URL (default: `http://host.docker.internal:11434`) | For local models |
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,7 @@
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ag_ui.core import EventType, StateSnapshotEvent
|
from ag_ui.core import EventType, StateSnapshotEvent
|
||||||
|
|
@ -26,8 +24,8 @@ from haiku.rag.capabilities.policy import (
|
||||||
)
|
)
|
||||||
from haiku.rag.capabilities.rag import AGENT_PREAMBLE, RAGState, create_capability
|
from haiku.rag.capabilities.rag import AGENT_PREAMBLE, RAGState, create_capability
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import load_yaml_config
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config import get_config
|
||||||
from haiku.rag.telemetry import configure as configure_telemetry
|
from haiku.rag.telemetry import configure as configure_telemetry
|
||||||
from haiku.rag.utils import get_model
|
from haiku.rag.utils import get_model
|
||||||
|
|
||||||
|
|
@ -40,19 +38,23 @@ logging.basicConfig(
|
||||||
)
|
)
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Load config
|
# The configuration places the database. This app serves one.
|
||||||
config_path = Path("/app/haiku.rag.yaml")
|
config = get_config()
|
||||||
if config_path.exists():
|
scope = DatabaseScope.resolve(config)
|
||||||
yaml_data = load_yaml_config(config_path)
|
if scope.covers_multiple:
|
||||||
config = AppConfig.model_validate(yaml_data)
|
raise SystemExit(
|
||||||
else:
|
f"lancedb.databases names {', '.join(scope.names)}; this app serves one "
|
||||||
config = AppConfig()
|
"database: configure exactly one entry"
|
||||||
|
)
|
||||||
|
[database] = scope.databases
|
||||||
|
|
||||||
# Get DB path from environment
|
|
||||||
db_path_str = os.getenv("DB_PATH", "haiku_rag.lancedb")
|
|
||||||
db_path = Path(db_path_str)
|
|
||||||
|
|
||||||
logger.info(f"Database path: {db_path}")
|
def _database_exists() -> bool:
|
||||||
|
"""A database behind a URI has no path to check."""
|
||||||
|
return database.db_path is None or database.db_path.exists()
|
||||||
|
|
||||||
|
|
||||||
|
logger.info(f"Database: {database.name} at {database.location}")
|
||||||
logger.info(f"QA Provider: {config.qa.model.provider}, Model: {config.qa.model.name}")
|
logger.info(f"QA Provider: {config.qa.model.provider}, Model: {config.qa.model.name}")
|
||||||
|
|
||||||
# Only HaikuRAG client is a singleton (expensive to create)
|
# Only HaikuRAG client is a singleton (expensive to create)
|
||||||
|
|
@ -71,7 +73,7 @@ async def get_client() -> HaikuRAG:
|
||||||
if _client is None:
|
if _client is None:
|
||||||
async with _client_lock:
|
async with _client_lock:
|
||||||
if _client is None:
|
if _client is None:
|
||||||
client = HaikuRAG(db_path=db_path, config=config, create=True)
|
client = HaikuRAG(config=config, create=True)
|
||||||
await client.__aenter__()
|
await client.__aenter__()
|
||||||
_client = client
|
_client = client
|
||||||
return _client
|
return _client
|
||||||
|
|
@ -82,7 +84,7 @@ class AppDeps:
|
||||||
state: dict[str, Any] = field(default_factory=dict)
|
state: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
capability = create_capability(db_path=db_path, config=config, defer_loading=False)
|
capability = create_capability(config=config, defer_loading=False)
|
||||||
|
|
||||||
agent = Agent(
|
agent = Agent(
|
||||||
get_model(config.qa.model, config),
|
get_model(config.qa.model, config),
|
||||||
|
|
@ -138,15 +140,15 @@ async def health_check(_: Request) -> JSONResponse:
|
||||||
"status": "healthy",
|
"status": "healthy",
|
||||||
"qa_provider": config.qa.model.provider,
|
"qa_provider": config.qa.model.provider,
|
||||||
"qa_model": config.qa.model.name,
|
"qa_model": config.qa.model.name,
|
||||||
"db_path": str(db_path),
|
"db_path": str(database.location),
|
||||||
"db_exists": db_path.exists(),
|
"db_exists": _database_exists(),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
async def list_documents(_: Request) -> JSONResponse:
|
async def list_documents(_: Request) -> JSONResponse:
|
||||||
"""List all documents in the database."""
|
"""List all documents in the database."""
|
||||||
if not db_path.exists():
|
if not _database_exists():
|
||||||
return JSONResponse({"documents": [], "error": "Database not found"})
|
return JSONResponse({"documents": [], "error": "Database not found"})
|
||||||
|
|
||||||
client = await get_client()
|
client = await get_client()
|
||||||
|
|
@ -162,11 +164,11 @@ async def list_documents(_: Request) -> JSONResponse:
|
||||||
|
|
||||||
async def db_info(_: Request) -> JSONResponse:
|
async def db_info(_: Request) -> JSONResponse:
|
||||||
"""Get database info and statistics."""
|
"""Get database info and statistics."""
|
||||||
if not db_path.exists():
|
if not _database_exists():
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
{
|
{
|
||||||
"exists": False,
|
"exists": False,
|
||||||
"path": str(db_path),
|
"path": str(database.location),
|
||||||
"documents": 0,
|
"documents": 0,
|
||||||
"chunks": 0,
|
"chunks": 0,
|
||||||
}
|
}
|
||||||
|
|
@ -180,7 +182,7 @@ async def db_info(_: Request) -> JSONResponse:
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
{
|
{
|
||||||
"exists": True,
|
"exists": True,
|
||||||
"path": str(db_path),
|
"path": str(database.location),
|
||||||
"documents": stats["documents"].get("num_rows", 0),
|
"documents": stats["documents"].get("num_rows", 0),
|
||||||
"chunks": stats["chunks"].get("num_rows", 0),
|
"chunks": stats["chunks"].get("num_rows", 0),
|
||||||
"documents_bytes": stats["documents"].get("total_bytes", 0),
|
"documents_bytes": stats["documents"].get("total_bytes", 0),
|
||||||
|
|
@ -214,7 +216,7 @@ async def visualize_chunk(request: Request) -> JSONResponse:
|
||||||
if isinstance(parsed, list):
|
if isinstance(parsed, list):
|
||||||
refs = [str(x) for x in parsed]
|
refs = [str(x) for x in parsed]
|
||||||
|
|
||||||
if not db_path.exists():
|
if not _database_exists():
|
||||||
return JSONResponse({"error": "Database not found"}, status_code=404)
|
return JSONResponse({"error": "Database not found"}, status_code=404)
|
||||||
|
|
||||||
client = await get_client()
|
client = await get_client()
|
||||||
|
|
|
||||||
|
|
@ -11,13 +11,14 @@ services:
|
||||||
ports:
|
ports:
|
||||||
- "127.0.0.1:8001:8000"
|
- "127.0.0.1:8001:8000"
|
||||||
environment:
|
environment:
|
||||||
- DB_PATH=/data
|
- HAIKU_RAG_CONFIG_PATH=/app/haiku.rag.yaml
|
||||||
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY:-}
|
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY:-}
|
||||||
- OPENAI_API_KEY=${OPENAI_API_KEY:-}
|
- OPENAI_API_KEY=${OPENAI_API_KEY:-}
|
||||||
- OLLAMA_BASE_URL=${OLLAMA_BASE_URL:-http://host.docker.internal:11434}
|
- OLLAMA_BASE_URL=${OLLAMA_BASE_URL:-http://host.docker.internal:11434}
|
||||||
- LOGFIRE_TOKEN=${LOGFIRE_TOKEN:-}
|
- LOGFIRE_TOKEN=${LOGFIRE_TOKEN:-}
|
||||||
volumes:
|
volumes:
|
||||||
- ${DB_PATH:-./data/haiku.rag.lancedb}:/data
|
# haiku.rag.yaml places the database at /data.
|
||||||
|
- ${DB_VOLUME:-./data/haiku.rag.lancedb}:/data
|
||||||
- ./backend:/app/src:ro
|
- ./backend:/app/src:ro
|
||||||
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro
|
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro
|
||||||
extra_hosts:
|
extra_hosts:
|
||||||
|
|
|
||||||
|
|
@ -7,13 +7,14 @@ services:
|
||||||
ports:
|
ports:
|
||||||
- "127.0.0.1:8001:8000"
|
- "127.0.0.1:8001:8000"
|
||||||
environment:
|
environment:
|
||||||
- DB_PATH=/data
|
- HAIKU_RAG_CONFIG_PATH=/app/haiku.rag.yaml
|
||||||
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY:-}
|
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY:-}
|
||||||
- OPENAI_API_KEY=${OPENAI_API_KEY:-}
|
- OPENAI_API_KEY=${OPENAI_API_KEY:-}
|
||||||
- OLLAMA_BASE_URL=${OLLAMA_BASE_URL:-http://host.docker.internal:11434}
|
- OLLAMA_BASE_URL=${OLLAMA_BASE_URL:-http://host.docker.internal:11434}
|
||||||
- LOGFIRE_TOKEN=${LOGFIRE_TOKEN:-}
|
- LOGFIRE_TOKEN=${LOGFIRE_TOKEN:-}
|
||||||
volumes:
|
volumes:
|
||||||
- ${DB_PATH:-./data/haiku.rag.lancedb}:/data
|
# haiku.rag.yaml places the database at /data.
|
||||||
|
- ${DB_VOLUME:-./data/haiku.rag.lancedb}:/data
|
||||||
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro
|
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro
|
||||||
extra_hosts:
|
extra_hosts:
|
||||||
- "host.docker.internal:host-gateway"
|
- "host.docker.internal:host-gateway"
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,12 @@
|
||||||
# haiku.rag configuration for the chat app
|
# haiku.rag configuration for the chat app
|
||||||
# Copy to haiku.rag.yaml and customize as needed
|
# Copy to haiku.rag.yaml and customize as needed
|
||||||
|
|
||||||
|
# The database. The compose files mount DB_VOLUME (default
|
||||||
|
# ./data/haiku.rag.lancedb) at /data.
|
||||||
|
lancedb:
|
||||||
|
databases:
|
||||||
|
haiku.rag: /data
|
||||||
|
|
||||||
# QA model configuration
|
# QA model configuration
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
|
|
|
||||||
12
docs/apps.md
12
docs/apps.md
|
|
@ -40,8 +40,8 @@ Create a `.env` file in the `app/` directory:
|
||||||
ANTHROPIC_API_KEY=your-anthropic-key
|
ANTHROPIC_API_KEY=your-anthropic-key
|
||||||
OPENAI_API_KEY=your-openai-key
|
OPENAI_API_KEY=your-openai-key
|
||||||
|
|
||||||
# Database path
|
# Host path of the LanceDB database, mounted at /data
|
||||||
DB_PATH=/path/to/your/haiku.rag.lancedb
|
DB_VOLUME=/path/to/your/haiku.rag.lancedb
|
||||||
|
|
||||||
# Optional: Ollama base URL (if using local models)
|
# Optional: Ollama base URL (if using local models)
|
||||||
OLLAMA_BASE_URL=http://localhost:11434
|
OLLAMA_BASE_URL=http://localhost:11434
|
||||||
|
|
@ -50,16 +50,22 @@ OLLAMA_BASE_URL=http://localhost:11434
|
||||||
LOGFIRE_TOKEN=your-logfire-token
|
LOGFIRE_TOKEN=your-logfire-token
|
||||||
```
|
```
|
||||||
|
|
||||||
For full configuration, mount a `haiku.rag.yaml` file:
|
The mounted `haiku.rag.yaml` places the database at `/data` and configures the models; the compose files point `HAIKU_RAG_CONFIG_PATH` at it:
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
# app/haiku.rag.yaml
|
# app/haiku.rag.yaml
|
||||||
|
lancedb:
|
||||||
|
databases:
|
||||||
|
haiku.rag: /data
|
||||||
|
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: anthropic
|
provider: anthropic
|
||||||
name: claude-sonnet-4-20250514
|
name: claude-sonnet-4-20250514
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Outside compose, the backend loads its configuration like the CLI: `HAIKU_RAG_CONFIG_PATH`, then `./haiku.rag.yaml`, then the platform directory.
|
||||||
|
|
||||||
## API endpoints
|
## API endpoints
|
||||||
|
|
||||||
| Endpoint | Method | Description |
|
| Endpoint | Method | Description |
|
||||||
|
|
|
||||||
|
|
@ -145,12 +145,6 @@ Applications serving AG-UI should adapt the agent with Pydantic AI's `AGUIAdapte
|
||||||
|
|
||||||
## Database Selection
|
## Database Selection
|
||||||
|
|
||||||
RAG and analysis capabilities select databases in this order:
|
RAG and analysis capabilities cover the databases the configuration places: [`lancedb.databases`](../configuration/storage.md#multiple-databases), or with nothing configured the default database `haiku.rag` under `storage.data_dir`. The `db_path` argument places one database where the configuration places none; beside `lancedb.databases` it raises `AmbiguousDatabaseError`.
|
||||||
|
|
||||||
1. The `db_path` argument.
|
|
||||||
2. `HAIKU_RAG_DB`.
|
|
||||||
3. [`lancedb.databases`](../configuration/storage.md#multiple-databases), which selects the full configured set.
|
|
||||||
4. [`lancedb.uri`](../configuration/storage.md#changing-the-default-database-path), which selects one database.
|
|
||||||
5. `config.storage.data_dir / "haiku.rag.lancedb"`.
|
|
||||||
|
|
||||||
Passing a client through `rag=` bypasses this selection. The capability uses the databases covered by that client and does not close it.
|
Passing a client through `rag=` bypasses this selection. The capability uses the databases covered by that client and does not close it.
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,7 @@ The `haiku-rag` CLI provides complete document management functionality.
|
||||||
|
|
||||||
Per-command options:
|
Per-command options:
|
||||||
|
|
||||||
- `--db` - Specify custom database path
|
- `--db` - Open the database at this path, named by its stem, whatever the configuration places
|
||||||
- `-h` - Show help for specific command
|
- `-h` - Show help for specific command
|
||||||
|
|
||||||
Example:
|
Example:
|
||||||
|
|
|
||||||
|
|
@ -85,10 +85,9 @@ ingester:
|
||||||
delete_orphans: true
|
delete_orphans: true
|
||||||
|
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: "" # Empty for local, or db://, s3://, az://, gs://
|
databases: {} # Name-to-location map; empty places haiku.rag under data_dir
|
||||||
api_key: ""
|
api_key: "" # LanceDB Cloud (db://) credentials
|
||||||
region: ""
|
region: ""
|
||||||
databases: {} # Name-to-location map to search multiple at once; excludes uri
|
|
||||||
|
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
|
|
|
||||||
|
|
@ -68,18 +68,19 @@ Vacuum also folds new rows into the full-text index. Search stays correct withou
|
||||||
|
|
||||||
This is an upstream limitation rather than a `haiku.rag` setting. Compaction bounds itself by row count instead of bytes, and LanceDB's async API exposes no batch size or fragment target to override it. Tracked at [lancedb/lancedb#2325](https://github.com/lancedb/lancedb/issues/2325). The requirement above will drop once compaction batches by bytes.
|
This is an upstream limitation rather than a `haiku.rag` setting. Compaction bounds itself by row count instead of bytes, and LanceDB's async API exposes no batch size or fragment target to override it. Tracked at [lancedb/lancedb#2325](https://github.com/lancedb/lancedb/issues/2325). The requirement above will drop once compaction batches by bytes.
|
||||||
|
|
||||||
### Changing the Default Database Path
|
### Placing the Database
|
||||||
|
|
||||||
`storage.data_dir` holds the default database, always called `haiku.rag.lancedb`. To put the database somewhere else for every command, give `lancedb.uri` a local path:
|
`lancedb.databases` maps a name to a location, a local path or a URI, and is the one way to place databases. With nothing configured, the database is the entry `haiku.rag` at `<storage.data_dir>/haiku.rag.lancedb`. To put one database somewhere else, name it:
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: /data/notes.lancedb
|
databases:
|
||||||
|
notes: /data/notes.lancedb
|
||||||
```
|
```
|
||||||
|
|
||||||
An explicit `--db PATH` overrides `lancedb.uri` for that invocation.
|
The name is what `source` carries in search results, citations and documents, and what `--db-name` and `sources` select. The default database answers to `haiku.rag`.
|
||||||
|
|
||||||
This places one database without naming it. Its `source` is `None` in search results, citations and documents, since only [`lancedb.databases`](#multiple-databases) assigns the names that carry provenance. A path here changes where the database lives, not what it is called.
|
An explicit `--db PATH` on the command line opens that database instead, named by the path's stem, whatever is configured. From Python, `db_path` places the database only where the configuration places none: beside `lancedb.databases` it raises `AmbiguousDatabaseError`.
|
||||||
|
|
||||||
A value with no scheme is a local path wherever it is configured, so `haiku-rag init` creates it and every command that opens an existing database requires it to exist. A mistyped path fails rather than becoming a new empty database.
|
A value with no scheme is a local path wherever it is configured, so `haiku-rag init` creates it and every command that opens an existing database requires it to exist. A mistyped path fails rather than becoming a new empty database.
|
||||||
|
|
||||||
|
|
@ -109,28 +110,31 @@ async with HaikuRAG(create=True) as client:
|
||||||
|
|
||||||
The [default location](index.md#configuration-file-locations) is platform-specific (e.g., `~/Library/Application Support/haiku.rag/` on macOS).
|
The [default location](index.md#configuration-file-locations) is platform-specific (e.g., `~/Library/Application Support/haiku.rag/` on macOS).
|
||||||
|
|
||||||
Opening a nonexistent unnamed local database raises `FileNotFoundError`, naming its path. This prevents accidental database creation from typos or misconfigured paths. A database named in `lancedb.databases` raises `SourceUnavailableError` instead, naming the database and not its location.
|
Opening a nonexistent local database given as a path raises `FileNotFoundError`, naming the path. This prevents accidental database creation from typos or misconfigured paths. A configured or default database raises `SourceUnavailableError` instead, naming the database and not its location.
|
||||||
|
|
||||||
## Remote Storage
|
## Remote Storage
|
||||||
|
|
||||||
For remote storage, use the `lancedb` settings with various backends:
|
For remote storage, give the database a URI as its location. Credentials and storage options are connection settings, shared by every database in the configuration:
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
# LanceDB Cloud
|
# LanceDB Cloud
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: db://your-database-name
|
databases:
|
||||||
|
papers: db://your-database-name
|
||||||
api_key: your-api-key
|
api_key: your-api-key
|
||||||
region: us-west-2 # optional
|
region: us-west-2 # optional
|
||||||
|
|
||||||
# Amazon S3
|
# Amazon S3
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: s3://my-bucket/my-table
|
databases:
|
||||||
|
papers: s3://my-bucket/my-table
|
||||||
storage_options:
|
storage_options:
|
||||||
region: us-east-1
|
region: us-east-1
|
||||||
|
|
||||||
# Amazon S3 with explicit credentials
|
# Amazon S3 with explicit credentials
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: s3://my-bucket/my-table
|
databases:
|
||||||
|
papers: s3://my-bucket/my-table
|
||||||
storage_options:
|
storage_options:
|
||||||
aws_access_key_id: YOUR_ACCESS_KEY
|
aws_access_key_id: YOUR_ACCESS_KEY
|
||||||
aws_secret_access_key: YOUR_SECRET_KEY
|
aws_secret_access_key: YOUR_SECRET_KEY
|
||||||
|
|
@ -138,7 +142,8 @@ lancedb:
|
||||||
|
|
||||||
# S3-compatible (SeaweedFS, Tigris, etc.)
|
# S3-compatible (SeaweedFS, Tigris, etc.)
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: s3://my-bucket/my-table
|
databases:
|
||||||
|
papers: s3://my-bucket/my-table
|
||||||
storage_options:
|
storage_options:
|
||||||
endpoint: http://localhost:8333
|
endpoint: http://localhost:8333
|
||||||
aws_access_key_id: YOUR_ACCESS_KEY
|
aws_access_key_id: YOUR_ACCESS_KEY
|
||||||
|
|
@ -148,21 +153,24 @@ lancedb:
|
||||||
|
|
||||||
# Azure Blob Storage
|
# Azure Blob Storage
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: az://my-container/my-table
|
databases:
|
||||||
|
papers: az://my-container/my-table
|
||||||
|
|
||||||
# Google Cloud Storage
|
# Google Cloud Storage
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: gs://my-bucket/my-table
|
databases:
|
||||||
|
papers: gs://my-bucket/my-table
|
||||||
|
|
||||||
# HDFS
|
# HDFS
|
||||||
lancedb:
|
lancedb:
|
||||||
uri: hdfs://namenode:port/path/to/table
|
databases:
|
||||||
|
papers: hdfs://namenode:port/path/to/table
|
||||||
```
|
```
|
||||||
|
|
||||||
- **LanceDB Cloud** (`db://`): Requires `api_key` and `region`. Table optimization and indexing are managed server-side.
|
- **LanceDB Cloud** (`db://`): Requires `api_key` and `region`. Table optimization and indexing are managed server-side.
|
||||||
- **Object storage** (`s3://`, `gs://`, `az://`, `hdfs://`): Uses `storage_options` for credentials and endpoint configuration. Authentication can also be provided via environment variables (`AWS_ACCESS_KEY_ID`, `AWS_SECRET_ACCESS_KEY`, etc.) or cloud provider SDK defaults (AWS CLI, Azure CLI, gcloud).
|
- **Object storage** (`s3://`, `gs://`, `az://`, `hdfs://`): Uses `storage_options` for credentials and endpoint configuration. Authentication can also be provided via environment variables (`AWS_ACCESS_KEY_ID`, `AWS_SECRET_ACCESS_KEY`, etc.) or cloud provider SDK defaults (AWS CLI, Azure CLI, gcloud).
|
||||||
- **S3-compatible stores** (MinIO, Tigris, etc.): Set `endpoint` in `storage_options`. When using `http://` endpoints, also set `allow_http: "true"`.
|
- **S3-compatible stores** (MinIO, Tigris, etc.): Set `endpoint` in `storage_options`. When using `http://` endpoints, also set `allow_http: "true"`.
|
||||||
- **Local path** (no scheme): `uri` also takes a local path, which is how the default database is pointed elsewhere. See [Changing the Default Database Path](#changing-the-default-database-path).
|
- **Local path** (no scheme): a location without a scheme is a local path. See [Placing the Database](#placing-the-database).
|
||||||
|
|
||||||
The `storage_options` keys are case-insensitive and passed directly to the underlying object store library. Available keys depend on the backend. See the [LanceDB storage docs](https://lancedb.com/docs/storage/) for details.
|
The `storage_options` keys are case-insensitive and passed directly to the underlying object store library. Available keys depend on the backend. See the [LanceDB storage docs](https://lancedb.com/docs/storage/) for details.
|
||||||
|
|
||||||
|
|
@ -192,7 +200,7 @@ The recommended layout for production is "different buckets, same account, separ
|
||||||
|
|
||||||
Each process picks up its own credentials from the AWS default chain (env vars, IAM instance role, AWS profile), so no credentials are hard-coded in the configuration files.
|
Each process picks up its own credentials from the AWS default chain (env vars, IAM instance role, AWS profile), so no credentials are hard-coded in the configuration files.
|
||||||
|
|
||||||
`haiku-ingester` writes the database the configuration places, so a `lancedb.uri` needs no further option. `--db PATH` overrides it. When `lancedb.databases` contains more than one database the ingester has no way to name which it writes, and refuses to start with `AmbiguousDatabaseError`: give each database its own ingester process, each with a configuration naming a single database, or select one with `--db PATH`. A one-entry mapping is unambiguous and is accepted.
|
`haiku-ingester` writes the database the configuration places, so a one-entry `lancedb.databases` needs no further option. `--db PATH` overrides it. When `lancedb.databases` contains more than one database the ingester has no way to name which it writes, and refuses to start with `AmbiguousDatabaseError`: give each database its own ingester process, each with a configuration naming a single database, or select one with `--db PATH`.
|
||||||
|
|
||||||
## Multiple Databases
|
## Multiple Databases
|
||||||
|
|
||||||
|
|
@ -206,7 +214,7 @@ lancedb:
|
||||||
notes: /data/notes.lancedb
|
notes: /data/notes.lancedb
|
||||||
```
|
```
|
||||||
|
|
||||||
A location can be a URI or local path. `databases` and `uri` are mutually exclusive.
|
A location can be a URI or local path.
|
||||||
|
|
||||||
Results, documents, and citations use the configured name as `source`. An unavailable configured database raises `SourceUnavailableError`, which names the database and not its location, so a location never travels in an error a consumer might render or log. A migration, configuration or read-only failure keeps its own type, with the database named in the message. Commands that report on a database, such as `info`, still show where it is.
|
Results, documents, and citations use the configured name as `source`. An unavailable configured database raises `SourceUnavailableError`, which names the database and not its location, so a location never travels in an error a consumer might render or log. A migration, configuration or read-only failure keeps its own type, with the database named in the message. Commands that report on a database, such as `info`, still show where it is.
|
||||||
|
|
||||||
|
|
@ -227,7 +235,7 @@ results = await client.search("query") # every database
|
||||||
results = await client.search("query", sources=["papers"]) # one of them
|
results = await client.search("query", sources=["papers"]) # one of them
|
||||||
```
|
```
|
||||||
|
|
||||||
Candidates are combined into one ranked list with the configured reranker, or by cosine similarity to the query when reranking is disabled, with within-database rank breaking ties (full-text-only searches order by retrieval score). `SearchResult.source`, `Citation.source`, and `Document.source` contain the configured database name. The name is retained when a client covers only one named database. Databases configured through `lancedb.uri` are unnamed, so their `source` is `None`.
|
Candidates are combined into one ranked list with the configured reranker, or by cosine similarity to the query when reranking is disabled, with within-database rank breaking ties (full-text-only searches order by retrieval score). `SearchResult.source`, `Citation.source`, and `Document.source` carry the database name, for a set and for one database alike.
|
||||||
|
|
||||||
The CLI labels results and citations only when the operation spans multiple databases. A command already narrowed with `--db-name` does not repeat the name on every result.
|
The CLI labels results and citations only when the operation spans multiple databases. A command already narrowed with `--db-name` does not repeat the name on every result.
|
||||||
|
|
||||||
|
|
@ -283,7 +291,7 @@ haiku-rag --db-name papers list # one of them
|
||||||
haiku-rag --db-name papers migrate
|
haiku-rag --db-name papers migrate
|
||||||
```
|
```
|
||||||
|
|
||||||
`--db-name` selects an entry from `lancedb.databases`, including remote entries. `--db` selects a local path and overrides the configured location. A single-database command requires one of these options when multiple databases are configured. A configured set of one is selected automatically.
|
`--db-name` selects an entry from `lancedb.databases`, including remote entries, and `haiku.rag` when nothing is configured. `--db` opens a local path, named by its stem, whatever is configured. A single-database command requires one of these options when multiple databases are configured. A configured set of one is selected automatically.
|
||||||
|
|
||||||
Each database is created, migrated and vacuumed on its own:
|
Each database is created, migrated and vacuumed on its own:
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -35,7 +35,8 @@ snapshot is only meaningful while this process is the only writer.
|
||||||
## Storage
|
## Storage
|
||||||
|
|
||||||
LanceDB is embedded, so there is no server. The same code runs against a local
|
LanceDB is embedded, so there is no server. The same code runs against a local
|
||||||
directory, S3, GCS, Azure or LanceDB Cloud by changing `lancedb.uri`.
|
directory, S3, GCS, Azure or LanceDB Cloud by changing a database's location in
|
||||||
|
`lancedb.databases`.
|
||||||
|
|
||||||
Tables are versioned. Vacuum collapses old versions on a retention window, and
|
Tables are versioned. Vacuum collapses old versions on a retention window, and
|
||||||
[tags](cli.md) name a state across all tables so a database can be restored to
|
[tags](cli.md) name a state across all tables so a database can be restored to
|
||||||
|
|
|
||||||
|
|
@ -27,7 +27,7 @@ async with HaikuRAG("path/to/database.lancedb", read_only=True) as client:
|
||||||
`async with` is the lifecycle. A caller that owns the client some other way releases it with `await client.aclose()`, which does the same work for every client shape. `client.close()` closes the connection to one database and nothing else, since draining the background vacuum and releasing the embedder and reranker are awaitable; it refuses a client covering several.
|
`async with` is the lifecycle. A caller that owns the client some other way releases it with `await client.aclose()`, which does the same work for every client shape. `client.close()` closes the connection to one database and nothing else, since draining the background vacuum and releasing the embedder and reranker are awaitable; it refuses a client covering several.
|
||||||
|
|
||||||
!!! note
|
!!! note
|
||||||
Databases must be explicitly created with `create=True` or via `haiku-rag init` before use. Opening a nonexistent unnamed local database raises `FileNotFoundError`, naming its path; one named in `lancedb.databases` raises `SourceUnavailableError`, which names the database rather than its location.
|
Databases must be explicitly created with `create=True` or via `haiku-rag init` before use. Opening a nonexistent local database given as a path raises `FileNotFoundError`, naming the path; a configured or default database raises `SourceUnavailableError`, which names the database rather than its location. A path beside a configured `lancedb.databases` raises `AmbiguousDatabaseError`.
|
||||||
|
|
||||||
!!! note
|
!!! note
|
||||||
Read-only mode is useful for safely accessing databases without risk of modification. It blocks all write operations and downgrades an embedding provider/name mismatch to a warning instead of raising `ConfigMismatchError`.
|
Read-only mode is useful for safely accessing databases without risk of modification. It blocks all write operations and downgrades an embedding provider/name mismatch to a warning instead of raising `ConfigMismatchError`.
|
||||||
|
|
@ -268,8 +268,8 @@ On the constructor `sources=[]` means something else. Passing `sources` alongsid
|
||||||
|
|
||||||
```python
|
```python
|
||||||
client.covers_multiple # whether the client covers more than one database
|
client.covers_multiple # whether the client covers more than one database
|
||||||
client.source_names # configured names, in order
|
client.source_names # database names, in order; known before the client opens
|
||||||
client.source # one configured name, or None for a set or unnamed database
|
client.source # the one database's name, or None for a set
|
||||||
|
|
||||||
owner = await client.reader_for("papers") # the client reading that database
|
owner = await client.reader_for("papers") # the client reading that database
|
||||||
papers, wiki = await client.clients_for(["papers", "wiki"])
|
papers, wiki = await client.clients_for(["papers", "wiki"])
|
||||||
|
|
@ -277,6 +277,17 @@ papers, wiki = await client.clients_for(["papers", "wiki"])
|
||||||
|
|
||||||
`reader_for` and `clients_for` open databases lazily and return borrowed clients. They remain valid while the covering client is open and inherit its read-only mode. The covering client owns and closes their database sessions.
|
`reader_for` and `clients_for` open databases lazily and return borrowed clients. They remain valid while the covering client is open and inherit its read-only mode. The covering client owns and closes their database sessions.
|
||||||
|
|
||||||
|
To learn what a configuration covers without opening anything, resolve it:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from haiku.rag.client import DatabaseScope
|
||||||
|
|
||||||
|
for ref in DatabaseScope.resolve(config).databases:
|
||||||
|
print(ref.name, ref.location) # "haiku.rag", Path(".../haiku.rag.lancedb") when nothing is configured
|
||||||
|
```
|
||||||
|
|
||||||
|
`DatabaseScope.resolve` is pure: it reads the configuration and classifies each location as a local path or a URI.
|
||||||
|
|
||||||
### Filtering Search Results
|
### Filtering Search Results
|
||||||
|
|
||||||
Filter search results to only include chunks from documents matching specific criteria:
|
Filter search results to only include chunks from documents matching specific criteria:
|
||||||
|
|
|
||||||
|
|
@ -49,6 +49,13 @@ async def evaluate_dataset(
|
||||||
if document_filter is not None:
|
if document_filter is not None:
|
||||||
console.print(f"Document filter: {document_filter}", style="dim")
|
console.print(f"Document filter: {document_filter}", style="dim")
|
||||||
|
|
||||||
|
if db_path is not None and config.lancedb.databases:
|
||||||
|
raise ValueError(
|
||||||
|
"--db PATH places the database where the configuration places none, "
|
||||||
|
f"and this configuration names {', '.join(config.lancedb.databases)} "
|
||||||
|
"in lancedb.databases. Drop --db to evaluate the configured set."
|
||||||
|
)
|
||||||
|
|
||||||
if not skip_db:
|
if not skip_db:
|
||||||
if spec.uses_configured_databases(config, db_path):
|
if spec.uses_configured_databases(config, db_path):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|
@ -157,7 +164,11 @@ def run(
|
||||||
config: Path | None = typer.Option(
|
config: Path | None = typer.Option(
|
||||||
None, "--config", help="Path to haiku.rag YAML config file."
|
None, "--config", help="Path to haiku.rag YAML config file."
|
||||||
),
|
),
|
||||||
db: Path | None = typer.Option(None, "--db", help="Override the database path."),
|
db: Path | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--db",
|
||||||
|
help="Database path, where the configuration places no database.",
|
||||||
|
),
|
||||||
skip_db: bool = typer.Option(
|
skip_db: bool = typer.Option(
|
||||||
False, "--skip-db", help="Skip updating the evaluation db."
|
False, "--skip-db", help="Skip updating the evaluation db."
|
||||||
),
|
),
|
||||||
|
|
|
||||||
|
|
@ -44,7 +44,7 @@ class CapabilityRunResult:
|
||||||
cited_uris: list[str] = field(default_factory=list)
|
cited_uris: list[str] = field(default_factory=list)
|
||||||
cited_chunk_ids: list[str] = field(default_factory=list)
|
cited_chunk_ids: list[str] = field(default_factory=list)
|
||||||
# The database each cited chunk came from, in the order they were cited.
|
# The database each cited chunk came from, in the order they were cited.
|
||||||
# Empty string where the database is unnamed.
|
# Empty string for a citation built without a source.
|
||||||
cited_sources: list[str] = field(default_factory=list)
|
cited_sources: list[str] = field(default_factory=list)
|
||||||
searched_uris: list[str] = field(default_factory=list)
|
searched_uris: list[str] = field(default_factory=list)
|
||||||
n_searches: int = 0
|
n_searches: int = 0
|
||||||
|
|
|
||||||
|
|
@ -89,9 +89,9 @@ class DatasetSpec:
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Whether `lancedb.databases` places the databases to evaluate over.
|
"""Whether `lancedb.databases` places the databases to evaluate over.
|
||||||
|
|
||||||
A path names one database and wins over the configuration, both when it
|
`--db PATH` places the database where the configuration places none;
|
||||||
comes from `--db` and when the client resolves it. True for a mapping of
|
`evaluate_dataset` refuses the two together. True for a mapping of one,
|
||||||
one, which is a configured database like any other and keeps its name.
|
which is a configured database like any other and keeps its name.
|
||||||
"""
|
"""
|
||||||
return bool(config.lancedb.databases) and override_path is None
|
return bool(config.lancedb.databases) and override_path is None
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -711,6 +711,37 @@ class TestEvaluateDatasetJudgeModel:
|
||||||
assert mock_qa.call_args[1]["judge_model"] is custom_judge
|
assert mock_qa.call_args[1]["judge_model"] is custom_judge
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_db_path_beside_configured_databases_is_refused(tmp_path) -> None:
|
||||||
|
"""`--db` places the database where the configuration places none; beside
|
||||||
|
`lancedb.databases` the run refuses before touching anything."""
|
||||||
|
from haiku.rag.config.models import LanceDBConfig
|
||||||
|
|
||||||
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(databases={"alpha": str(tmp_path / "a.lancedb")})
|
||||||
|
)
|
||||||
|
spec = DatasetSpec(
|
||||||
|
key="test",
|
||||||
|
db_filename="test.lancedb",
|
||||||
|
document_loader=lambda: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
|
||||||
|
document_mapper=lambda doc: None,
|
||||||
|
qa_loader=lambda: [], # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
|
||||||
|
qa_case_builder=lambda idx, doc: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="alpha"):
|
||||||
|
await evaluate_dataset(
|
||||||
|
spec=spec,
|
||||||
|
config=config,
|
||||||
|
skip_db=True,
|
||||||
|
skip_retrieval=True,
|
||||||
|
skip_qa=True,
|
||||||
|
limit=None,
|
||||||
|
name=None,
|
||||||
|
db_path=tmp_path / "other.lancedb",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestExperimentMetadataTargets:
|
class TestExperimentMetadataTargets:
|
||||||
def test_default_target_is_rag_capability(self) -> None:
|
def test_default_target_is_rag_capability(self) -> None:
|
||||||
result = build_experiment_metadata(
|
result = build_experiment_metadata(
|
||||||
|
|
|
||||||
|
|
@ -485,8 +485,8 @@ def test_records_the_database_each_citation_came_from():
|
||||||
assert result.cited_sources == ["alpha", "beta", "alpha"]
|
assert result.cited_sources == ["alpha", "beta", "alpha"]
|
||||||
|
|
||||||
|
|
||||||
def test_an_unnamed_database_records_no_source():
|
def test_a_hand_built_citation_without_a_source_records_an_empty_string():
|
||||||
"""One database names nothing: the field holds an empty string."""
|
"""A citation built without a source is recorded as an empty string."""
|
||||||
from haiku.rag.capabilities._base import EvidenceState
|
from haiku.rag.capabilities._base import EvidenceState
|
||||||
from haiku.rag.capabilities.ledger import CapabilityEvidenceRecord
|
from haiku.rag.capabilities.ledger import CapabilityEvidenceRecord
|
||||||
from haiku.rag.store.models.citation import Citation
|
from haiku.rag.store.models.citation import Citation
|
||||||
|
|
|
||||||
|
|
@ -26,8 +26,8 @@ uv run python examples/custom_agent.py /path/to/db.lancedb
|
||||||
|
|
||||||
**Script:** `custom_agent_agui.py`
|
**Script:** `custom_agent_agui.py`
|
||||||
|
|
||||||
A Starlette app that adapts a native RAG-capable agent to AG-UI.
|
A Starlette app that adapts a native RAG-capable agent to AG-UI. The configuration places the database (`HAIKU_RAG_CONFIG_PATH`, or `./haiku.rag.yaml`):
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
DB_PATH=/path/to/db.lancedb uv run uvicorn examples.custom_agent_agui:app --reload --port 8000
|
uv run uvicorn examples.custom_agent_agui:app --reload --port 8000
|
||||||
```
|
```
|
||||||
|
|
|
||||||
|
|
@ -8,13 +8,13 @@ Requirements:
|
||||||
- An Anthropic API key (for the QA model) or adjust the model below
|
- An Anthropic API key (for the QA model) or adjust the model below
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
DB_PATH=/path/to/db.lancedb uv run uvicorn examples.custom_agent_agui:app --reload --port 8000
|
uv run uvicorn examples.custom_agent_agui:app --reload --port 8000
|
||||||
|
|
||||||
|
The configuration places the database (HAIKU_RAG_CONFIG_PATH, or
|
||||||
|
./haiku.rag.yaml).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
|
||||||
import sys
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ag_ui.core import EventType, StateSnapshotEvent
|
from ag_ui.core import EventType, StateSnapshotEvent
|
||||||
|
|
@ -30,14 +30,7 @@ from haiku.rag.capabilities.compaction import create_capability as compaction
|
||||||
from haiku.rag.capabilities.policy import create_capability as citation_policy
|
from haiku.rag.capabilities.policy import create_capability as citation_policy
|
||||||
from haiku.rag.capabilities.rag import RAGState, create_capability
|
from haiku.rag.capabilities.rag import RAGState, create_capability
|
||||||
|
|
||||||
db_path = os.environ.get("DB_PATH")
|
capability = create_capability(defer_loading=False)
|
||||||
if not db_path:
|
|
||||||
print(
|
|
||||||
"Set DB_PATH environment variable to your haiku.rag database", file=sys.stderr
|
|
||||||
)
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
capability = create_capability(db_path=Path(db_path), defer_loading=False)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,4 @@
|
||||||
import logging
|
import logging
|
||||||
from functools import cached_property
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
|
@ -64,21 +63,10 @@ class HaikuRAGApp:
|
||||||
[ref] = self.scope.databases
|
[ref] = self.scope.databases
|
||||||
return ref
|
return ref
|
||||||
|
|
||||||
@cached_property
|
|
||||||
def _connection(self) -> "tuple[AppConfig, Path]":
|
|
||||||
"""How to open the one database this command works on, directly.
|
|
||||||
|
|
||||||
Derived per database from its configured location.
|
|
||||||
"""
|
|
||||||
from haiku.rag.client.session import default_db_path
|
|
||||||
|
|
||||||
config, db_path = self._one.connection(self.config)
|
|
||||||
return config, db_path or default_db_path(config)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def _store_config(self) -> AppConfig:
|
def _location(self) -> "Path | str":
|
||||||
"""The configuration for opening the one database directly."""
|
"""Where the one database this command works on is."""
|
||||||
return self._connection[0]
|
return self._one.location
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def _is_local(self) -> bool:
|
def _is_local(self) -> bool:
|
||||||
|
|
@ -91,17 +79,14 @@ class HaikuRAGApp:
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def _path(self) -> Path:
|
def _path(self) -> Path:
|
||||||
"""The path of the one database this command works on.
|
"""The path of the one local database this command works on."""
|
||||||
|
assert self._one.db_path is not None
|
||||||
A database behind a URI has none of its own, and the default stands in:
|
return self._one.db_path
|
||||||
the URI in `_store_config` is what decides where it connects.
|
|
||||||
"""
|
|
||||||
return self._connection[1]
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def display_path(self) -> "Path | str":
|
def display_path(self) -> "Path | str":
|
||||||
"""What a one-database command calls the database it opened."""
|
"""What a one-database command calls the database it opened."""
|
||||||
return self._one.db_path or self._one.uri
|
return self._one.location
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def database_missing(self) -> bool:
|
def database_missing(self) -> bool:
|
||||||
|
|
@ -140,7 +125,7 @@ class HaikuRAGApp:
|
||||||
self.console.print("[red]Database path does not exist.[/red]")
|
self.console.print("[red]Database path does not exist.[/red]")
|
||||||
return
|
return
|
||||||
|
|
||||||
info = await gather_database_info(self._store_config, self._path)
|
info = await gather_database_info(self._location, self.config)
|
||||||
|
|
||||||
if not info.exists:
|
if not info.exists:
|
||||||
self.console.print(
|
self.console.print(
|
||||||
|
|
@ -282,8 +267,8 @@ class HaikuRAGApp:
|
||||||
cm = status if status is not None else nullcontext()
|
cm = status if status is not None else nullcontext()
|
||||||
with cm:
|
with cm:
|
||||||
report = await run_doctor(
|
report = await run_doctor(
|
||||||
self._store_config,
|
self.config,
|
||||||
self._path,
|
self._location,
|
||||||
dict(os.environ),
|
dict(os.environ),
|
||||||
duplicates_out=duplicates_out,
|
duplicates_out=duplicates_out,
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
|
|
@ -340,8 +325,8 @@ class HaikuRAGApp:
|
||||||
return
|
return
|
||||||
|
|
||||||
async with Store(
|
async with Store(
|
||||||
self._path,
|
self._location,
|
||||||
config=self._store_config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
skip_migration_check=True,
|
skip_migration_check=True,
|
||||||
|
|
@ -415,15 +400,15 @@ class HaikuRAGApp:
|
||||||
"""
|
"""
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
|
|
||||||
return Store(self._path, config=self._store_config, read_only=self.read_only)
|
return Store(self._location, config=self.config, read_only=self.read_only)
|
||||||
|
|
||||||
def _tag_read_store(self) -> "Store":
|
def _tag_read_store(self) -> "Store":
|
||||||
"""Read-only store for tag inspection; works on old or drifted DBs."""
|
"""Read-only store for tag inspection; works on old or drifted DBs."""
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
|
|
||||||
return Store(
|
return Store(
|
||||||
self._path,
|
self._location,
|
||||||
config=self._store_config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
skip_migration_check=True,
|
skip_migration_check=True,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
|
|
@ -760,8 +745,8 @@ class HaikuRAGApp:
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
|
|
||||||
async with Store(
|
async with Store(
|
||||||
self._path,
|
self._location,
|
||||||
config=self._store_config,
|
config=self.config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
skip_migration_check=True,
|
skip_migration_check=True,
|
||||||
read_only=self.read_only,
|
read_only=self.read_only,
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,4 @@
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
|
||||||
from dataclasses import dataclass, field, replace
|
from dataclasses import dataclass, field, replace
|
||||||
from difflib import get_close_matches
|
from difflib import get_close_matches
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
@ -99,12 +98,7 @@ def _nearest_known_id(chunk_id: str, known_ids: list[str]) -> str:
|
||||||
|
|
||||||
|
|
||||||
def resolve_scope(db_path: Path | str | None, config: AppConfig) -> DatabaseScope:
|
def resolve_scope(db_path: Path | str | None, config: AppConfig) -> DatabaseScope:
|
||||||
"""The databases a capability covers, resolved once at its entry point.
|
"""The databases a capability covers, resolved once at its entry point."""
|
||||||
|
|
||||||
``HAIKU_RAG_DB`` is read here and nowhere else.
|
|
||||||
"""
|
|
||||||
if db_path is None and (env_db := os.environ.get("HAIKU_RAG_DB")):
|
|
||||||
db_path = Path(env_db).expanduser()
|
|
||||||
return DatabaseScope.resolve(config, database_path=db_path)
|
return DatabaseScope.resolve(config, database_path=db_path)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -42,14 +42,8 @@ def run_chat(
|
||||||
config.qa.model = model_config
|
config.qa.model = model_config
|
||||||
config.analysis.model = model_config
|
config.analysis.model = model_config
|
||||||
|
|
||||||
# The capabilities read the databases the scope covers, not what the
|
# The app opens the scope and lends that client to the capabilities, which
|
||||||
# configuration names: a `--db PATH` selection is outside the
|
# read what `--db PATH` or `--db-name NAME` selected.
|
||||||
# configuration, and a `--db-name NAME` selection is narrower than it.
|
|
||||||
if scope.covers_multiple:
|
|
||||||
capability_config, capability_db_path = config, None
|
|
||||||
else:
|
|
||||||
capability_config, capability_db_path = scope.databases[0].connection(config)
|
|
||||||
|
|
||||||
enabled = capabilities or ["rag"]
|
enabled = capabilities or ["rag"]
|
||||||
capability_list = []
|
capability_list = []
|
||||||
defer_loading = len(enabled) > 1
|
defer_loading = len(enabled) > 1
|
||||||
|
|
@ -68,8 +62,7 @@ def run_chat(
|
||||||
|
|
||||||
capability_list.append(
|
capability_list.append(
|
||||||
create_capability(
|
create_capability(
|
||||||
db_path=capability_db_path,
|
config=config,
|
||||||
config=capability_config,
|
|
||||||
defer_loading=defer_loading,
|
defer_loading=defer_loading,
|
||||||
vision=driving_model.vision,
|
vision=driving_model.vision,
|
||||||
)
|
)
|
||||||
|
|
@ -80,8 +73,7 @@ def run_chat(
|
||||||
|
|
||||||
capability_list.append(
|
capability_list.append(
|
||||||
create_capability(
|
create_capability(
|
||||||
db_path=capability_db_path,
|
config=config,
|
||||||
config=capability_config,
|
|
||||||
defer_loading=defer_loading,
|
defer_loading=defer_loading,
|
||||||
vision=driving_model.vision,
|
vision=driving_model.vision,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -148,10 +148,12 @@ class ChatApp(App):
|
||||||
# a client whose __aenter__ failed.
|
# a client whose __aenter__ failed.
|
||||||
await client.__aenter__()
|
await client.__aenter__()
|
||||||
self.client = client
|
self.client = client
|
||||||
# Lent to the capabilities: already the databases they were built for,
|
# Lent to the capabilities, with the scope it covers: one connection
|
||||||
# and one connection per database however many capabilities read it.
|
# per database however many capabilities read it, and the analysis
|
||||||
|
# sandbox is built over the same selection.
|
||||||
for capability in self._capabilities:
|
for capability in self._capabilities:
|
||||||
capability.borrowed_rag = client
|
capability.borrowed_rag = client
|
||||||
|
capability.scope = self.scope
|
||||||
|
|
||||||
self._agent = Agent(
|
self._agent = Agent(
|
||||||
self._model,
|
self._model,
|
||||||
|
|
@ -425,7 +427,8 @@ class ChatApp(App):
|
||||||
|
|
||||||
def on_document_filter_modal_filter_changed(self, event: Any) -> None:
|
def on_document_filter_modal_filter_changed(self, event: Any) -> None:
|
||||||
"""Scope the conversation to the selection: the filter carries the ids,
|
"""Scope the conversation to the selection: the filter carries the ids,
|
||||||
and `sources` restricts the search to the databases the selection names.
|
and over a set `sources` restricts the search to the databases the
|
||||||
|
selection names. One database needs no narrowing by source.
|
||||||
"""
|
"""
|
||||||
from haiku.rag.tools.filters import build_document_id_filter
|
from haiku.rag.tools.filters import build_document_id_filter
|
||||||
|
|
||||||
|
|
@ -434,10 +437,9 @@ class ChatApp(App):
|
||||||
doc_filter = build_document_id_filter(
|
doc_filter = build_document_id_filter(
|
||||||
sorted({doc_id for _, doc_id in event.selected})
|
sorted({doc_id for _, doc_id in event.selected})
|
||||||
)
|
)
|
||||||
selected_sources = {source for source, _ in event.selected}
|
selected_sources = sorted({source for source, _ in event.selected if source})
|
||||||
sources: list[str] | None = None
|
covers_multiple = self.client is not None and self.client.covers_multiple
|
||||||
if selected_sources and None not in selected_sources:
|
sources = selected_sources if covers_multiple and selected_sources else None
|
||||||
sources = sorted(s for s in selected_sources if s is not None)
|
|
||||||
for namespace, state_type in (
|
for namespace, state_type in (
|
||||||
(RAG_STATE_NAMESPACE, RAGState),
|
(RAG_STATE_NAMESPACE, RAGState),
|
||||||
(ANALYSIS_STATE_NAMESPACE, AnalysisState),
|
(ANALYSIS_STATE_NAMESPACE, AnalysisState),
|
||||||
|
|
|
||||||
|
|
@ -24,15 +24,18 @@ class DocumentCheckbox(Checkbox):
|
||||||
self.doc_id = doc_id
|
self.doc_id = doc_id
|
||||||
|
|
||||||
|
|
||||||
def _labelled(docs) -> list[tuple[str, str | None, str]]:
|
def _labelled(
|
||||||
"""Each document's label, database and id, sorted by label. The database is
|
docs, *, name_database: bool = False
|
||||||
named alongside the title, which a title alone does not say. Labels are
|
) -> list[tuple[str, str | None, str]]:
|
||||||
escaped: titles and database names are data, not Textual markup."""
|
"""Each document's label, database and id, sorted by label. Across several
|
||||||
|
databases the database is named alongside the title, which a title alone
|
||||||
|
does not say. Labels are escaped: titles and database names are data, not
|
||||||
|
Textual markup."""
|
||||||
rows = [
|
rows = [
|
||||||
(
|
(
|
||||||
escape(
|
escape(
|
||||||
f"{doc.title or doc.uri or doc.id}"
|
f"{doc.title or doc.uri or doc.id}"
|
||||||
+ (f" ({doc.source})" if doc.source else "")
|
+ (f" ({doc.source})" if name_database and doc.source else "")
|
||||||
),
|
),
|
||||||
doc.source,
|
doc.source,
|
||||||
doc.id,
|
doc.id,
|
||||||
|
|
@ -221,7 +224,9 @@ class DocumentFilterModal(ModalScreen):
|
||||||
DocumentCheckbox(
|
DocumentCheckbox(
|
||||||
label, source, doc_id, value=(source, doc_id) in self._selected
|
label, source, doc_id, value=(source, doc_id) in self._selected
|
||||||
)
|
)
|
||||||
for label, source, doc_id in _labelled(docs)
|
for label, source, doc_id in _labelled(
|
||||||
|
docs, name_database=self.client.covers_multiple
|
||||||
|
)
|
||||||
]
|
]
|
||||||
if boxes:
|
if boxes:
|
||||||
await filter_list.mount_all(boxes)
|
await filter_list.mount_all(boxes)
|
||||||
|
|
|
||||||
|
|
@ -90,9 +90,10 @@ def resolve_scope(
|
||||||
"""The databases a command works on, resolved once.
|
"""The databases a command works on, resolved once.
|
||||||
|
|
||||||
The CLI decides only what it alone knows: that `--db` and `--db-name` are
|
The CLI decides only what it alone knows: that `--db` and `--db-name` are
|
||||||
the same thing said twice, and whether this command can read more than one.
|
the same thing said twice, that a human typing `--db PATH` means that
|
||||||
Everything else — an unknown name, a `lancedb.uri`, the default location —
|
database whatever is configured, and whether this command can read more
|
||||||
is `DatabaseScope.resolve`'s to answer.
|
than one. Everything else — an unknown name, the default location — is
|
||||||
|
`DatabaseScope.resolve`'s to answer.
|
||||||
"""
|
"""
|
||||||
from haiku.rag.client.scope import DatabaseScope
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
|
||||||
|
|
@ -100,9 +101,12 @@ def resolve_scope(
|
||||||
raise AmbiguousDatabaseError(
|
raise AmbiguousDatabaseError(
|
||||||
"pass --db or --db-name, not both: they name the same thing"
|
"pass --db or --db-name, not both: they name the same thing"
|
||||||
)
|
)
|
||||||
scope = DatabaseScope.resolve(
|
if db is not None:
|
||||||
get_config(), database_name=_db_name, database_path=db
|
try:
|
||||||
)
|
return DatabaseScope.at(db)
|
||||||
|
except ValueError as error:
|
||||||
|
raise typer.BadParameter(str(error), param_hint="--db") from error
|
||||||
|
scope = DatabaseScope.resolve(get_config(), database_name=_db_name)
|
||||||
if scope.covers_multiple and not covers_set:
|
if scope.covers_multiple and not covers_set:
|
||||||
raise AmbiguousDatabaseError(
|
raise AmbiguousDatabaseError(
|
||||||
f"lancedb.databases names {', '.join(sorted(scope.names))}; this "
|
f"lancedb.databases names {', '.join(sorted(scope.names))}; this "
|
||||||
|
|
|
||||||
|
|
@ -21,7 +21,6 @@ from haiku.rag.client.session import (
|
||||||
FederatedSession,
|
FederatedSession,
|
||||||
SingleDatabaseSession,
|
SingleDatabaseSession,
|
||||||
aclose_quietly,
|
aclose_quietly,
|
||||||
default_db_path,
|
|
||||||
)
|
)
|
||||||
from haiku.rag.config import AppConfig, get_config
|
from haiku.rag.config import AppConfig, get_config
|
||||||
from haiku.rag.converters import get_converter
|
from haiku.rag.converters import get_converter
|
||||||
|
|
@ -129,25 +128,22 @@ class HaikuRAG:
|
||||||
"""Initialize the RAG client with a database path.
|
"""Initialize the RAG client with a database path.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
db_path: Path or string path to the database. When omitted, resolves
|
db_path: Path or string path to the database, named by its stem.
|
||||||
``lancedb.databases``, then ``lancedb.uri``, then the default
|
Valid where the configuration places no database; beside
|
||||||
path under ``storage.data_dir``.
|
``lancedb.databases`` it raises ``AmbiguousDatabaseError``.
|
||||||
|
When omitted, the configured databases are covered, or the
|
||||||
|
default database ``haiku.rag`` under ``storage.data_dir``.
|
||||||
config: Configuration to use. Defaults to the current global config.
|
config: Configuration to use. Defaults to the current global config.
|
||||||
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.
|
||||||
sources: Names from ``config.lancedb.databases`` this client covers,
|
sources: Names of the databases this client covers, None for all of
|
||||||
None for all of them. Only that setting names databases, so a
|
them. Rejected alongside ``db_path``, which says the same thing
|
||||||
name raises when ``lancedb.uri`` placed the database, and is
|
|
||||||
rejected alongside ``db_path``, which says the same thing
|
|
||||||
another way. ``[]`` raises too: a client over no database can do
|
another way. ``[]`` raises too: a client over no database can do
|
||||||
nothing, unlike ``sources=[]`` on a search, which is a selection
|
nothing, unlike ``sources=[]`` on a search, which is a selection
|
||||||
of nothing to search.
|
of nothing to search.
|
||||||
"""
|
"""
|
||||||
self._configured = config if config is not None else get_config()
|
self._configured = config if config is not None else get_config()
|
||||||
# What the caller configured, kept intact: entering derives a
|
|
||||||
# single-database configuration from it, and every re-entry derives
|
|
||||||
# from the configured set.
|
|
||||||
self._config = self._configured
|
self._config = self._configured
|
||||||
self._requested_db_path = Path(db_path) if db_path is not None else None
|
self._requested_db_path = Path(db_path) if db_path is not None else None
|
||||||
if self._requested_db_path is not None and sources is not None:
|
if self._requested_db_path is not None and sources is not None:
|
||||||
|
|
@ -169,24 +165,31 @@ class HaikuRAG:
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def covers_multiple(self) -> bool:
|
def covers_multiple(self) -> bool:
|
||||||
"""Whether this client reads from more than one database."""
|
"""Whether this client reads from more than one database.
|
||||||
return isinstance(self._session, FederatedSession)
|
|
||||||
|
Known before the client enters: coverage is a fact of the resolved
|
||||||
|
scope.
|
||||||
|
"""
|
||||||
|
if self._session is not None:
|
||||||
|
return isinstance(self._session, FederatedSession)
|
||||||
|
return self._resolve_scope().covers_multiple
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def source_names(self) -> tuple[str, ...]:
|
def source_names(self) -> tuple[str, ...]:
|
||||||
"""The configured databases this client covers, in configured order.
|
"""The databases this client covers, by name, in configured order.
|
||||||
|
|
||||||
A single database contributes its own name, or nothing where the
|
Known before the client enters: coverage is a fact of the resolved
|
||||||
configuration named none.
|
scope.
|
||||||
"""
|
"""
|
||||||
if isinstance(self._session, FederatedSession):
|
if isinstance(self._session, FederatedSession):
|
||||||
return self._session.names
|
return self._session.names
|
||||||
return () if self.source is None else (self.source,)
|
if isinstance(self._session, SingleDatabaseSession):
|
||||||
|
return (self._session.source,)
|
||||||
|
return self._resolve_scope().names
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def source(self) -> str | None:
|
def source(self) -> str | None:
|
||||||
"""The configured database this client reads, or None while covering a
|
"""The database this client reads, or None while covering a set."""
|
||||||
set or reading a database the configuration did not name."""
|
|
||||||
if isinstance(self._session, SingleDatabaseSession):
|
if isinstance(self._session, SingleDatabaseSession):
|
||||||
return self._session.source
|
return self._session.source
|
||||||
return None
|
return None
|
||||||
|
|
@ -339,15 +342,12 @@ class HaikuRAG:
|
||||||
return self
|
return self
|
||||||
|
|
||||||
[ref] = scope.databases
|
[ref] = scope.databases
|
||||||
self._config, db_path = ref.connection(self._configured)
|
|
||||||
|
|
||||||
self._session = await SingleDatabaseSession(
|
self._session = await SingleDatabaseSession(
|
||||||
db_path if db_path is not None else default_db_path(self._config),
|
ref,
|
||||||
self._config,
|
self._config,
|
||||||
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,
|
||||||
source=ref.name,
|
|
||||||
).open()
|
).open()
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
@ -410,9 +410,8 @@ class HaikuRAG:
|
||||||
|
|
||||||
`lender` is the client that opened it, whose reranker this one borrows.
|
`lender` is the client that opened it, whose reranker this one borrows.
|
||||||
"""
|
"""
|
||||||
client = cls(
|
client = cls(config=session.config, read_only=session.read_only)
|
||||||
session.db_path, config=session.config, read_only=session.read_only
|
client._scope = DatabaseScope((session.ref,))
|
||||||
)
|
|
||||||
client._session = session
|
client._session = session
|
||||||
client._owns_session = False
|
client._owns_session = False
|
||||||
client._lender = lender
|
client._lender = lender
|
||||||
|
|
@ -858,7 +857,7 @@ class HaikuRAG:
|
||||||
if unknown:
|
if unknown:
|
||||||
raise UnknownDatabaseError(
|
raise UnknownDatabaseError(
|
||||||
f"unknown database(s) {', '.join(sorted(set(unknown)))}; this "
|
f"unknown database(s) {', '.join(sorted(set(unknown)))}; this "
|
||||||
f"client covers {', '.join(sorted(covered)) or 'a single unnamed database'}"
|
f"client covers {', '.join(sorted(covered))}"
|
||||||
)
|
)
|
||||||
|
|
||||||
async def clients_covering(
|
async def clients_covering(
|
||||||
|
|
@ -882,8 +881,8 @@ class HaikuRAG:
|
||||||
return []
|
return []
|
||||||
if sources != [self.source]:
|
if sources != [self.source]:
|
||||||
raise UnknownDatabaseError(
|
raise UnknownDatabaseError(
|
||||||
f"unknown database(s) {', '.join(sources) or '(none)'}; this "
|
f"unknown database(s) {', '.join(sources)}; this client covers "
|
||||||
f"client covers {self.source or 'a single unnamed database'}"
|
f"{self.source}"
|
||||||
)
|
)
|
||||||
return [self]
|
return [self]
|
||||||
|
|
||||||
|
|
@ -905,9 +904,6 @@ class HaikuRAG:
|
||||||
if not await self.clients_covering(sources):
|
if not await self.clients_covering(sources):
|
||||||
return []
|
return []
|
||||||
results = await search(self, query, limit, search_type, filter, include_images)
|
results = await search(self, query, limit, search_type, filter, include_images)
|
||||||
# A database named in config keeps its name even when it is the only one
|
|
||||||
# this client covers. Only an unnamed `lancedb.uri` database leaves
|
|
||||||
# source unset.
|
|
||||||
for result in results:
|
for result in results:
|
||||||
result.source = self.source
|
result.source = self.source
|
||||||
return results
|
return results
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,4 @@
|
||||||
|
from collections.abc import Mapping
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
@ -8,50 +9,60 @@ from haiku.rag.store.exceptions import (
|
||||||
)
|
)
|
||||||
from haiku.rag.utils import locate_database
|
from haiku.rag.utils import locate_database
|
||||||
|
|
||||||
|
DEFAULT_DATABASE_FILENAME = "haiku.rag.lancedb"
|
||||||
|
|
||||||
|
|
||||||
|
def database_name(path: Path) -> str:
|
||||||
|
"""The name a database at `path` answers to: the path's stem."""
|
||||||
|
if not path.stem:
|
||||||
|
raise ValueError(f"a database at {path} has no name: the path has no stem")
|
||||||
|
return path.stem
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class DatabaseRef:
|
class DatabaseRef:
|
||||||
"""A resolved database location, and the configured name it answers to.
|
"""A resolved database: the name it answers to, and where it is.
|
||||||
|
|
||||||
Exactly one of ``uri`` and ``db_path`` is set. ``name`` is the key from
|
``name`` is the key from ``lancedb.databases``, or the stem of a path the
|
||||||
``lancedb.databases``, and the only identity that leaves the configuration:
|
caller gave. It is the only identity that leaves the configuration: it
|
||||||
it travels in results, citations and errors, where a location must not.
|
travels in results, citations and errors, where a location must not.
|
||||||
None where nothing names the database.
|
``location`` is a local path, or a URI. ``given`` marks a path the caller
|
||||||
|
gave, whose errors may name it: the caller already knows where it is.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name: str | None
|
name: str
|
||||||
uri: str
|
location: Path | str
|
||||||
db_path: Path | None
|
given: bool = False
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
if bool(self.uri) == (self.db_path is not None):
|
if not self.name.strip():
|
||||||
|
raise ValueError(f"a database at {self.location} has no name")
|
||||||
|
if isinstance(self.location, str):
|
||||||
|
if not self.location.strip():
|
||||||
|
raise ValueError(f"database {self.name!r} has no location")
|
||||||
|
object.__setattr__(self, "location", locate_database(self.location))
|
||||||
|
if self.given and not isinstance(self.location, Path):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"a database is either a URI or a local path: "
|
f"database {self.name!r} is given as a path, and {self.location} "
|
||||||
f"got uri={self.uri!r} and db_path={self.db_path!r}"
|
"is a URI"
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def at(cls, path: Path | str, *, name: str | None = None) -> "DatabaseRef":
|
def at(cls, path: Path | str) -> "DatabaseRef":
|
||||||
"""A database at a path the caller named, taken as given."""
|
"""A database at a path the caller named, taken as given."""
|
||||||
return cls(name=name, uri="", db_path=Path(path))
|
path = Path(path)
|
||||||
|
return cls(name=database_name(path), location=path, given=True)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def configured(cls, name: str | None, location: str) -> "DatabaseRef":
|
def configured(cls, name: str, location: str | Path) -> "DatabaseRef":
|
||||||
"""A database the configuration placed, by ``lancedb.uri`` or an entry in
|
"""A database the configuration placed. A location carrying a scheme is
|
||||||
``lancedb.databases``. A location carrying a scheme is a URI, anything
|
a URI, anything else a local path."""
|
||||||
else a local path."""
|
return cls(name=name, location=location)
|
||||||
uri, db_path = locate_database(location)
|
|
||||||
return cls(name=name, uri=uri, db_path=db_path)
|
|
||||||
|
|
||||||
def connection(self, config: AppConfig) -> tuple[AppConfig, Path | None]:
|
@property
|
||||||
"""The configuration and path to open this one database with.
|
def db_path(self) -> Path | None:
|
||||||
|
"""The local path, or None for a database behind a URI."""
|
||||||
A copy: the caller's configuration still names whatever set it named.
|
return self.location if isinstance(self.location, Path) else None
|
||||||
"""
|
|
||||||
one = config.model_copy(deep=True)
|
|
||||||
one.lancedb.databases = {}
|
|
||||||
one.lancedb.uri = self.uri
|
|
||||||
return one, self.db_path
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
|
|
@ -59,10 +70,7 @@ class DatabaseScope:
|
||||||
"""The databases an operation covers.
|
"""The databases an operation covers.
|
||||||
|
|
||||||
Resolved once, from configuration plus at most one selector, then passed
|
Resolved once, from configuration plus at most one selector, then passed
|
||||||
down. Never empty.
|
down. Never empty. Nothing here reads the environment.
|
||||||
|
|
||||||
Nothing here reads the environment: ``HAIKU_RAG_DB`` is the capability entry
|
|
||||||
point's to honour.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
databases: tuple[DatabaseRef, ...]
|
databases: tuple[DatabaseRef, ...]
|
||||||
|
|
@ -71,6 +79,14 @@ class DatabaseScope:
|
||||||
if not self.databases:
|
if not self.databases:
|
||||||
raise ValueError("a scope covers at least one database")
|
raise ValueError("a scope covers at least one database")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def at(cls, path: Path | str) -> "DatabaseScope":
|
||||||
|
"""One database at a path the caller named, whatever is configured.
|
||||||
|
|
||||||
|
The CLI's ``--db``: a human typing a path means that database.
|
||||||
|
"""
|
||||||
|
return cls((DatabaseRef.at(path),))
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def resolve(
|
def resolve(
|
||||||
cls,
|
cls,
|
||||||
|
|
@ -81,9 +97,10 @@ class DatabaseScope:
|
||||||
) -> "DatabaseScope":
|
) -> "DatabaseScope":
|
||||||
"""The databases named by `config` and at most one selector.
|
"""The databases named by `config` and at most one selector.
|
||||||
|
|
||||||
A path names one database that nothing calls anything; a name selects one
|
The configuration places databases: ``lancedb.databases``, or where it
|
||||||
of the configured set and keeps its name. With no selector the configured
|
names none, the default database under ``storage.data_dir`` as the entry
|
||||||
set is covered in configuration order, a set of one included.
|
``haiku.rag``. A name selects one of them. A path places a database
|
||||||
|
where the configuration places none, and is refused beside one it does.
|
||||||
"""
|
"""
|
||||||
if database_name is not None and database_path is not None:
|
if database_name is not None and database_path is not None:
|
||||||
raise AmbiguousDatabaseError(
|
raise AmbiguousDatabaseError(
|
||||||
|
|
@ -91,33 +108,37 @@ class DatabaseScope:
|
||||||
"pass one of them"
|
"pass one of them"
|
||||||
)
|
)
|
||||||
|
|
||||||
declared = config.lancedb.databases
|
configured = config.lancedb.databases
|
||||||
|
|
||||||
if database_path is not None:
|
if database_path is not None:
|
||||||
return cls((DatabaseRef.at(database_path),))
|
if configured:
|
||||||
|
raise AmbiguousDatabaseError(
|
||||||
|
"a database path and lancedb.databases both place the "
|
||||||
|
f"database: db_path={Path(database_path)} and databases "
|
||||||
|
f"name {', '.join(sorted(configured))}; pass one of them"
|
||||||
|
)
|
||||||
|
return cls.at(database_path)
|
||||||
|
|
||||||
|
declared: Mapping[str, str | Path] = configured or {
|
||||||
|
"haiku.rag": config.storage.data_dir / DEFAULT_DATABASE_FILENAME
|
||||||
|
}
|
||||||
|
|
||||||
if database_name is not None:
|
if database_name is not None:
|
||||||
if database_name not in declared:
|
if database_name not in declared:
|
||||||
raise UnknownDatabaseError(
|
raise UnknownDatabaseError(
|
||||||
f"unknown database {database_name!r}; lancedb.databases names "
|
f"unknown database {database_name!r}; the databases are "
|
||||||
f"{', '.join(sorted(declared)) or 'nothing'}"
|
f"{', '.join(sorted(declared))}"
|
||||||
)
|
)
|
||||||
return cls(
|
return cls(
|
||||||
(DatabaseRef.configured(database_name, declared[database_name]),)
|
(DatabaseRef.configured(database_name, declared[database_name]),)
|
||||||
)
|
)
|
||||||
|
|
||||||
if declared:
|
return cls(
|
||||||
return cls(
|
tuple(
|
||||||
tuple(
|
DatabaseRef.configured(name, location)
|
||||||
DatabaseRef.configured(name, location)
|
for name, location in declared.items()
|
||||||
for name, location in declared.items()
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
)
|
||||||
if config.lancedb.uri:
|
|
||||||
return cls((DatabaseRef.configured(None, config.lancedb.uri),))
|
|
||||||
|
|
||||||
return cls((DatabaseRef.at(config.storage.data_dir / "haiku.rag.lancedb"),))
|
|
||||||
|
|
||||||
def select(self, names: list[str]) -> "DatabaseScope":
|
def select(self, names: list[str]) -> "DatabaseScope":
|
||||||
"""The databases in this scope named by `names`, in the order given.
|
"""The databases in this scope named by `names`, in the order given.
|
||||||
|
|
@ -128,7 +149,7 @@ class DatabaseScope:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"sources=[] selects no database; pass None for all of them"
|
"sources=[] selects no database; pass None for all of them"
|
||||||
)
|
)
|
||||||
by_name = {ref.name: ref for ref in self.databases if ref.name is not None}
|
by_name = {ref.name: ref for ref in self.databases}
|
||||||
missing = [name for name in names if name not in by_name]
|
missing = [name for name in names if name not in by_name]
|
||||||
if missing:
|
if missing:
|
||||||
raise UnknownDatabaseError(
|
raise UnknownDatabaseError(
|
||||||
|
|
@ -144,5 +165,5 @@ class DatabaseScope:
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def names(self) -> tuple[str, ...]:
|
def names(self) -> tuple[str, ...]:
|
||||||
"""The configured names covered, in order. Empty where none is named."""
|
"""The names of the databases covered, in order."""
|
||||||
return tuple(ref.name for ref in self.databases if ref.name is not None)
|
return tuple(ref.name for ref in self.databases)
|
||||||
|
|
|
||||||
|
|
@ -43,36 +43,29 @@ async def aclose_quietly(closeable: Any, what: str) -> None:
|
||||||
logger.debug("Closing the %s failed on teardown", what, exc_info=True)
|
logger.debug("Closing the %s failed on teardown", what, exc_info=True)
|
||||||
|
|
||||||
|
|
||||||
def default_db_path(config: AppConfig) -> Path:
|
|
||||||
"""Where a database lives when its location names no path."""
|
|
||||||
return config.storage.data_dir / "haiku.rag.lancedb"
|
|
||||||
|
|
||||||
|
|
||||||
class SingleDatabaseSession:
|
class SingleDatabaseSession:
|
||||||
"""One database: its store, its repositories, and their lifecycle.
|
"""One database: its store, its repositories, and their lifecycle.
|
||||||
|
|
||||||
Everything that needs a store lives here, so nothing above has to ask whether
|
Everything that needs a store lives here, so nothing above has to ask whether
|
||||||
it has one. ``source`` is the configured name this database answers to, or
|
it has one. Built from the resolved reference: ``source`` is the name it
|
||||||
None where nothing names it.
|
answers to, and the store receives its location.
|
||||||
|
|
||||||
``db_path``, ``config``, ``read_only`` and ``source`` are readable: a client
|
``ref``, ``config``, ``read_only`` and ``source`` are readable: a client
|
||||||
borrowing this session reports them as its own.
|
borrowing this session reports them as its own.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
db_path: Path | str,
|
ref: DatabaseRef,
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
*,
|
*,
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
create: bool = False,
|
create: bool = False,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
source: str | None = None,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
self.db_path = db_path
|
self.ref = ref
|
||||||
self.config = config
|
self.config = config
|
||||||
self.read_only = read_only
|
self.read_only = read_only
|
||||||
self.source = source
|
|
||||||
self._skip_validation = skip_validation
|
self._skip_validation = skip_validation
|
||||||
self._create = create
|
self._create = create
|
||||||
self._vacuum_tasks: set[asyncio.Task] = set()
|
self._vacuum_tasks: set[asyncio.Task] = set()
|
||||||
|
|
@ -80,19 +73,25 @@ class SingleDatabaseSession:
|
||||||
self._vacuum_dirty = False
|
self._vacuum_dirty = False
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def location(self) -> Path | str:
|
def source(self) -> str:
|
||||||
"""Configured URI or local path for this database.
|
return self.ref.name
|
||||||
|
|
||||||
Not `db_path`, which is a placeholder where a URI holds the database.
|
@property
|
||||||
"""
|
def location(self) -> Path | str:
|
||||||
return self.config.lancedb.uri or self.db_path
|
"""Where this database is: its path, or its URI."""
|
||||||
|
return self.ref.location
|
||||||
|
|
||||||
|
@property
|
||||||
|
def db_path(self) -> Path | None:
|
||||||
|
"""The local path, or None for a database behind a URI."""
|
||||||
|
return self.ref.db_path
|
||||||
|
|
||||||
async def open(self) -> "SingleDatabaseSession":
|
async def open(self) -> "SingleDatabaseSession":
|
||||||
"""Connect, validate, and build the repositories."""
|
"""Connect, validate, and build the repositories."""
|
||||||
failure: str | None = None
|
failure: str | None = None
|
||||||
try:
|
try:
|
||||||
self.store = Store(
|
self.store = Store(
|
||||||
self.db_path,
|
self.location,
|
||||||
config=self.config,
|
config=self.config,
|
||||||
skip_validation=self._skip_validation,
|
skip_validation=self._skip_validation,
|
||||||
create=self._create,
|
create=self._create,
|
||||||
|
|
@ -107,20 +106,22 @@ class SingleDatabaseSession:
|
||||||
raise
|
raise
|
||||||
except _NAMEABLE_FAILURES as error:
|
except _NAMEABLE_FAILURES as error:
|
||||||
# The message keeps its remedy and gains the database's name.
|
# The message keeps its remedy and gains the database's name.
|
||||||
if self.source is None:
|
if self.ref.given:
|
||||||
raise
|
raise
|
||||||
raise type(error)(f"database {self.source!r}: {error}") from error
|
raise type(error)(f"database {self.source!r}: {error}") from error
|
||||||
except Exception as error:
|
except Exception as error:
|
||||||
# Without a name there is nothing to report in the location's place.
|
# A path the caller gave may be named: the caller knows it already.
|
||||||
if self.source is None:
|
if self.ref.given:
|
||||||
raise
|
raise
|
||||||
failure = type(error).__name__
|
failure = (
|
||||||
|
"does not exist; create it with `haiku-rag init` or `create=True`"
|
||||||
|
if isinstance(error, FileNotFoundError)
|
||||||
|
else f"could not be opened: {type(error).__name__}"
|
||||||
|
)
|
||||||
if failure is not None:
|
if failure is not None:
|
||||||
# Raised outside the handler: the exception carries neither a cause
|
# Raised outside the handler: the exception carries neither a cause
|
||||||
# nor a location-bearing context.
|
# nor a location-bearing context.
|
||||||
raise SourceUnavailableError(
|
raise SourceUnavailableError(f"database {self.source!r} {failure}")
|
||||||
f"database {self.source!r} could not be opened: {failure}"
|
|
||||||
)
|
|
||||||
self.document_repository = DocumentRepository(self.store)
|
self.document_repository = DocumentRepository(self.store)
|
||||||
self.chunk_repository = ChunkRepository(self.store)
|
self.chunk_repository = ChunkRepository(self.store)
|
||||||
self.document_item_repository = DocumentItemRepository(self.store)
|
self.document_item_repository = DocumentItemRepository(self.store)
|
||||||
|
|
@ -266,9 +267,7 @@ class FederatedSession:
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._refs: dict[str, DatabaseRef] = {
|
self._refs: dict[str, DatabaseRef] = {ref.name: ref for ref in scope.databases}
|
||||||
ref.name: ref for ref in scope.databases if ref.name is not None
|
|
||||||
}
|
|
||||||
self._config = config
|
self._config = config
|
||||||
self._skip_validation = skip_validation
|
self._skip_validation = skip_validation
|
||||||
self._read_only = read_only
|
self._read_only = read_only
|
||||||
|
|
@ -309,14 +308,11 @@ class FederatedSession:
|
||||||
|
|
||||||
Registered here because a cancelled `gather` discards its results.
|
Registered here because a cancelled `gather` discards its results.
|
||||||
"""
|
"""
|
||||||
ref = self._refs[name]
|
|
||||||
one, db_path = ref.connection(self._config)
|
|
||||||
self._sessions[name] = await SingleDatabaseSession(
|
self._sessions[name] = await SingleDatabaseSession(
|
||||||
db_path if db_path is not None else default_db_path(one),
|
self._refs[name],
|
||||||
one,
|
self._config,
|
||||||
skip_validation=self._skip_validation,
|
skip_validation=self._skip_validation,
|
||||||
read_only=self._read_only,
|
read_only=self._read_only,
|
||||||
source=ref.name,
|
|
||||||
).open()
|
).open()
|
||||||
|
|
||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
|
|
|
||||||
|
|
@ -102,12 +102,12 @@ class LanceDBConfig(ConfigModel):
|
||||||
The cache sizes are per process, since the session is shared across
|
The cache sizes are per process, since the session is shared across
|
||||||
connections.
|
connections.
|
||||||
|
|
||||||
`databases` maps a name to a location, for searching multiple at once. The
|
`databases` maps a name to a location, a local path or a URI, and is the one
|
||||||
name is what results and citations carry, so a location never leaves the
|
way to place databases. The name is what results and citations carry, so a
|
||||||
configuration. Mutually exclusive with `uri`.
|
location never leaves the configuration. Empty means the default database,
|
||||||
|
`haiku.rag`, under `storage.data_dir`.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
uri: str = ""
|
|
||||||
api_key: str = ""
|
api_key: str = ""
|
||||||
region: str = ""
|
region: str = ""
|
||||||
storage_options: dict[str, str] = Field(default_factory=dict)
|
storage_options: dict[str, str] = Field(default_factory=dict)
|
||||||
|
|
@ -116,13 +116,23 @@ class LanceDBConfig(ConfigModel):
|
||||||
index_cache_size_bytes: int | None = Field(default=None, ge=0)
|
index_cache_size_bytes: int | None = Field(default=None, ge=0)
|
||||||
metadata_cache_size_bytes: int | None = Field(default=None, ge=0)
|
metadata_cache_size_bytes: int | None = Field(default=None, ge=0)
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="before")
|
||||||
def _one_way_of_naming_databases(self) -> "LanceDBConfig":
|
@classmethod
|
||||||
if self.uri and self.databases:
|
def _uri_names_its_replacement(cls, data: Any) -> Any:
|
||||||
|
if isinstance(data, dict) and "uri" in data:
|
||||||
|
if str(data["uri"]).strip():
|
||||||
|
raise ValueError(
|
||||||
|
"lancedb.uri was removed; write lancedb.databases: {NAME: "
|
||||||
|
f"{data['uri']!r}}} instead"
|
||||||
|
)
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"lancedb.uri and lancedb.databases are mutually exclusive: "
|
"lancedb.uri was removed; remove the empty key. With no "
|
||||||
"use uri for one unnamed location, or databases for named ones"
|
"lancedb.databases the database is haiku.rag under storage.data_dir"
|
||||||
)
|
)
|
||||||
|
return data
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def _every_database_is_named_and_placed(self) -> "LanceDBConfig":
|
||||||
for name, location in self.databases.items():
|
for name, location in self.databases.items():
|
||||||
# A blank name is falsy, so source routing reads it as absent; a
|
# A blank name is falsy, so source routing reads it as absent; a
|
||||||
# blank location resolves to the working directory.
|
# blank location resolves to the working directory.
|
||||||
|
|
|
||||||
|
|
@ -1080,7 +1080,7 @@ async def run_provider_checks(
|
||||||
|
|
||||||
async def run_doctor(
|
async def run_doctor(
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
db_path: Path,
|
location: Path | str,
|
||||||
environ: dict[str, str],
|
environ: dict[str, str],
|
||||||
duplicates_out: Path | None = None,
|
duplicates_out: Path | None = None,
|
||||||
on_progress: Callable[[str], None] | None = None,
|
on_progress: Callable[[str], None] | None = None,
|
||||||
|
|
@ -1092,7 +1092,7 @@ async def run_doctor(
|
||||||
"""
|
"""
|
||||||
notify = on_progress or (lambda _label: None)
|
notify = on_progress or (lambda _label: None)
|
||||||
notify("Inspecting tables")
|
notify("Inspecting tables")
|
||||||
db = await connect_lancedb(config, db_path)
|
db = await connect_lancedb(location, config)
|
||||||
stats = await get_database_stats(db)
|
stats = await get_database_stats(db)
|
||||||
|
|
||||||
results: list[CheckResult] = []
|
results: list[CheckResult] = []
|
||||||
|
|
@ -1110,7 +1110,7 @@ async def run_doctor(
|
||||||
missing = [name for name in REQUIRED_TABLES if not stats[name]["exists"]]
|
missing = [name for name in REQUIRED_TABLES if not stats[name]["exists"]]
|
||||||
if not missing:
|
if not missing:
|
||||||
async with Store(
|
async with Store(
|
||||||
db_path,
|
location,
|
||||||
config=config,
|
config=config,
|
||||||
skip_validation=True,
|
skip_validation=True,
|
||||||
read_only=True,
|
read_only=True,
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,5 @@
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
from fastapi import APIRouter, Depends, HTTPException, status
|
||||||
|
|
||||||
from haiku.rag.client.session import default_db_path
|
|
||||||
from haiku.rag.ingester.api.server import APIState, get_state
|
from haiku.rag.ingester.api.server import APIState, get_state
|
||||||
from haiku.rag.store.info import DatabaseInfo, gather_database_info
|
from haiku.rag.store.info import DatabaseInfo, gather_database_info
|
||||||
|
|
||||||
|
|
@ -21,5 +20,4 @@ async def database(state: APIState = Depends(get_state)) -> DatabaseInfo:
|
||||||
detail="database not configured",
|
detail="database not configured",
|
||||||
)
|
)
|
||||||
[ref] = state.scope.databases
|
[ref] = state.scope.databases
|
||||||
one, db_path = ref.connection(state.config)
|
return await gather_database_info(ref.location, state.config)
|
||||||
return await gather_database_info(one, db_path or default_db_path(one))
|
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,6 @@ import signal
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
@ -22,6 +21,8 @@ from haiku.rag.ingester.workers.retry import RetryPolicy
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||||
|
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_MANIFEST_EXTRA_KEY = "_manifest"
|
_MANIFEST_EXTRA_KEY = "_manifest"
|
||||||
|
|
@ -72,14 +73,14 @@ class IngesterApp:
|
||||||
WorkerPool, and a HaikuRAG client for the worker pool to ingest through.
|
WorkerPool, and a HaikuRAG client for the worker pool to ingest through.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, *, config: AppConfig, db_path: Path | None = None):
|
def __init__(self, *, config: AppConfig, scope: "DatabaseScope | None" = None):
|
||||||
|
"""The ingester over the database `scope` covers, or the one the
|
||||||
|
configuration places when no scope is handed in."""
|
||||||
from haiku.rag.client.scope import DatabaseScope
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.store.exceptions import AmbiguousDatabaseError
|
from haiku.rag.store.exceptions import AmbiguousDatabaseError
|
||||||
|
|
||||||
self._config = config
|
self._config = config
|
||||||
# `--db` is an explicit override; None leaves placement to the
|
self._scope = scope if scope is not None else DatabaseScope.resolve(config)
|
||||||
# configuration.
|
|
||||||
self._scope = DatabaseScope.resolve(config, database_path=db_path)
|
|
||||||
if self._scope.covers_multiple:
|
if self._scope.covers_multiple:
|
||||||
raise AmbiguousDatabaseError(
|
raise AmbiguousDatabaseError(
|
||||||
"haiku-ingester writes one database, and lancedb.databases "
|
"haiku-ingester writes one database, and lancedb.databases "
|
||||||
|
|
|
||||||
|
|
@ -42,6 +42,7 @@ from haiku.rag.store.exceptions import ( # noqa: E402
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.ingester.app import BatchProgress, BatchProgressCallback
|
from haiku.rag.ingester.app import BatchProgress, BatchProgressCallback
|
||||||
|
|
||||||
_cli = typer.Typer(
|
_cli = typer.Typer(
|
||||||
|
|
@ -218,6 +219,19 @@ def _load_manifest(path: Path) -> BatchManifest:
|
||||||
return BatchManifest.model_validate(data)
|
return BatchManifest.model_validate(data)
|
||||||
|
|
||||||
|
|
||||||
|
def _scope_for(db: Path | None) -> "DatabaseScope | None":
|
||||||
|
"""`--db PATH` is the operator's explicit override: that database, whatever
|
||||||
|
is configured. None leaves placement to the configuration."""
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
|
||||||
|
if db is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return DatabaseScope.at(db)
|
||||||
|
except ValueError as error:
|
||||||
|
raise typer.BadParameter(str(error), param_hint="--db") from error
|
||||||
|
|
||||||
|
|
||||||
@_cli.command("serve")
|
@_cli.command("serve")
|
||||||
def serve(
|
def serve(
|
||||||
db: Path | None = typer.Option(
|
db: Path | None = typer.Option(
|
||||||
|
|
@ -259,7 +273,7 @@ def serve(
|
||||||
app_config.ingester.api.port = port
|
app_config.ingester.api.port = port
|
||||||
if root_path is not None:
|
if root_path is not None:
|
||||||
app_config.ingester.api.root_path = root_path
|
app_config.ingester.api.root_path = root_path
|
||||||
app = IngesterApp(config=app_config, db_path=db)
|
app = IngesterApp(config=app_config, scope=_scope_for(db))
|
||||||
asyncio.run(app.serve(api=not no_api))
|
asyncio.run(app.serve(api=not no_api))
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -318,7 +332,7 @@ async def _run_batch(
|
||||||
) -> None:
|
) -> None:
|
||||||
from haiku.rag.ingester.app import IngesterApp
|
from haiku.rag.ingester.app import IngesterApp
|
||||||
|
|
||||||
app = IngesterApp(config=app_config, db_path=db_path)
|
app = IngesterApp(config=app_config, scope=_scope_for(db_path))
|
||||||
if dry_run:
|
if dry_run:
|
||||||
report = await app.run_batch_dry_run()
|
report = await app.run_batch_dry_run()
|
||||||
if report.failed_sweeps:
|
if report.failed_sweeps:
|
||||||
|
|
|
||||||
|
|
@ -20,13 +20,12 @@ async def database_lines(client: "HaikuRAG") -> list[str]:
|
||||||
Reported through the connection the client already holds. A failure becomes
|
Reported through the connection the client already holds. A failure becomes
|
||||||
a line of the report, and the other databases still report.
|
a line of the report, and the other databases still report.
|
||||||
"""
|
"""
|
||||||
from haiku.rag.store.engine import ConnectionMode
|
|
||||||
from haiku.rag.store.info import get_database_stats
|
from haiku.rag.store.info import get_database_stats
|
||||||
|
|
||||||
lines: list[str] = []
|
lines: list[str] = []
|
||||||
db_path = client.store.db_path
|
db_path = client.store.db_path
|
||||||
|
|
||||||
if client.store._connection_mode == ConnectionMode.LOCAL and not db_path.exists():
|
if db_path is not None and not db_path.exists():
|
||||||
return ["[red]Database path does not exist.[/red]"]
|
return ["[red]Database path does not exist.[/red]"]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
|
||||||
|
|
@ -32,9 +32,9 @@ def create_mcp_server(
|
||||||
"""Create an MCP server over one database.
|
"""Create an MCP server over one database.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
db_path: Path to the database file, or None to let `config` place it. A
|
db_path: Path to the database file, where `config` places none; or
|
||||||
path overrides a configured `lancedb.uri`: for a URI-backed
|
None to serve the database the configuration places. Beside
|
||||||
database, pass None.
|
`lancedb.databases` a path raises `AmbiguousDatabaseError`.
|
||||||
config: Configuration to use.
|
config: Configuration to use.
|
||||||
read_only: If True, write tools (add_document_*, delete_document) are not registered.
|
read_only: If True, write tools (add_document_*, delete_document) are not registered.
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -38,11 +38,12 @@ class ConnectionMode(Enum):
|
||||||
OBJECT_STORAGE = "object_storage"
|
OBJECT_STORAGE = "object_storage"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_config(config: AppConfig) -> "ConnectionMode":
|
def of(location: Path | str) -> "ConnectionMode":
|
||||||
uri = config.lancedb.uri
|
"""How a location is connected to: a path is local, `db://` is LanceDB
|
||||||
if not uri:
|
Cloud, any other scheme is object storage."""
|
||||||
|
if isinstance(location, Path) or "://" not in location:
|
||||||
return ConnectionMode.LOCAL
|
return ConnectionMode.LOCAL
|
||||||
if uri.startswith("db://"):
|
if location.startswith("db://"):
|
||||||
return ConnectionMode.CLOUD
|
return ConnectionMode.CLOUD
|
||||||
return ConnectionMode.OBJECT_STORAGE
|
return ConnectionMode.OBJECT_STORAGE
|
||||||
|
|
||||||
|
|
@ -72,8 +73,10 @@ def _session(config: AppConfig) -> lancedb.Session:
|
||||||
|
|
||||||
|
|
||||||
async def connect_lancedb(
|
async def connect_lancedb(
|
||||||
config: AppConfig, db_path: Path | None = None
|
location: Path | str, config: AppConfig
|
||||||
) -> lancedb.AsyncConnection:
|
) -> lancedb.AsyncConnection:
|
||||||
|
"""Connect to the database at `location`, with the connection settings
|
||||||
|
(credentials, storage options, caches, consistency) from `config`."""
|
||||||
interval = config.lancedb.read_consistency_interval_seconds
|
interval = config.lancedb.read_consistency_interval_seconds
|
||||||
kwargs: dict[str, Any] = {
|
kwargs: dict[str, Any] = {
|
||||||
"session": _session(config),
|
"session": _session(config),
|
||||||
|
|
@ -81,22 +84,19 @@ async def connect_lancedb(
|
||||||
timedelta(seconds=interval) if interval is not None else None
|
timedelta(seconds=interval) if interval is not None else None
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
mode = ConnectionMode.from_config(config)
|
mode = ConnectionMode.of(location)
|
||||||
if mode == ConnectionMode.CLOUD:
|
if mode == ConnectionMode.CLOUD:
|
||||||
return await lancedb.connect_async(
|
return await lancedb.connect_async(
|
||||||
uri=config.lancedb.uri,
|
uri=str(location),
|
||||||
api_key=config.lancedb.api_key,
|
api_key=config.lancedb.api_key,
|
||||||
region=config.lancedb.region,
|
region=config.lancedb.region,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
elif mode == ConnectionMode.OBJECT_STORAGE:
|
if mode == ConnectionMode.OBJECT_STORAGE:
|
||||||
if config.lancedb.storage_options:
|
if config.lancedb.storage_options:
|
||||||
kwargs["storage_options"] = config.lancedb.storage_options
|
kwargs["storage_options"] = config.lancedb.storage_options
|
||||||
return await lancedb.connect_async(uri=config.lancedb.uri, **kwargs)
|
return await lancedb.connect_async(uri=str(location), **kwargs)
|
||||||
else:
|
return await lancedb.connect_async(Path(location).absolute(), **kwargs)
|
||||||
if db_path is None:
|
|
||||||
raise ValueError("No lancedb.uri configured and no db_path provided")
|
|
||||||
return await lancedb.connect_async(db_path.absolute(), **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
def _stored_vector_dim(settings: dict) -> int | None:
|
def _stored_vector_dim(settings: dict) -> int | None:
|
||||||
|
|
@ -180,14 +180,24 @@ class TagInfo:
|
||||||
class Store:
|
class Store:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
db_path: Path | str,
|
location: Path | str,
|
||||||
config: AppConfig | None = None,
|
config: AppConfig | None = None,
|
||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
create: bool = False,
|
create: bool = False,
|
||||||
read_only: bool = False,
|
read_only: bool = False,
|
||||||
skip_migration_check: bool = False,
|
skip_migration_check: bool = False,
|
||||||
):
|
):
|
||||||
self.db_path: Path = Path(db_path)
|
"""A store over the database at `location`, a local path or a URI.
|
||||||
|
|
||||||
|
`config` supplies connection settings; where the database is comes
|
||||||
|
from `location` alone.
|
||||||
|
"""
|
||||||
|
self._location: Path | str = location
|
||||||
|
self.db_path: Path | None = (
|
||||||
|
Path(location)
|
||||||
|
if ConnectionMode.of(location) == ConnectionMode.LOCAL
|
||||||
|
else None
|
||||||
|
)
|
||||||
self._config = config if config is not None else get_config()
|
self._config = config if config is not None else get_config()
|
||||||
self._read_only = read_only
|
self._read_only = read_only
|
||||||
self._create = create
|
self._create = create
|
||||||
|
|
@ -200,7 +210,7 @@ class Store:
|
||||||
self._rebuild_lock = asyncio.Lock()
|
self._rebuild_lock = asyncio.Lock()
|
||||||
self._is_new_db = False
|
self._is_new_db = False
|
||||||
|
|
||||||
if self._connection_mode == ConnectionMode.LOCAL:
|
if self.db_path is not None:
|
||||||
if not self.db_path.exists():
|
if not self.db_path.exists():
|
||||||
if not create:
|
if not create:
|
||||||
raise FileNotFoundError(
|
raise FileNotFoundError(
|
||||||
|
|
@ -231,7 +241,7 @@ class Store:
|
||||||
async def _initialize(self):
|
async def _initialize(self):
|
||||||
"""Perform async initialization: connect to LanceDB, init tables, validate."""
|
"""Perform async initialization: connect to LanceDB, init tables, validate."""
|
||||||
self.db: lancedb.AsyncConnection = await connect_lancedb(
|
self.db: lancedb.AsyncConnection = await connect_lancedb(
|
||||||
self._config, self.db_path
|
self.location, self._config
|
||||||
)
|
)
|
||||||
|
|
||||||
# Read once and thread onward: on object storage each of these is a
|
# Read once and thread onward: on object storage each of these is a
|
||||||
|
|
@ -392,9 +402,14 @@ class Store:
|
||||||
needed = datetime.now() - oldest + TAG_RETENTION_MARGIN
|
needed = datetime.now() - oldest + TAG_RETENTION_MARGIN
|
||||||
return max(retention, needed)
|
return max(retention, needed)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def location(self) -> Path | str:
|
||||||
|
"""Where this store connected: a local path, or a URI."""
|
||||||
|
return self._location
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def _connection_mode(self) -> ConnectionMode:
|
def _connection_mode(self) -> ConnectionMode:
|
||||||
return ConnectionMode.from_config(self._config)
|
return ConnectionMode.of(self._location)
|
||||||
|
|
||||||
async def _ensure_vector_index(self) -> None:
|
async def _ensure_vector_index(self) -> None:
|
||||||
"""Create or rebuild vector index on chunks table.
|
"""Create or rebuild vector index on chunks table.
|
||||||
|
|
|
||||||
|
|
@ -96,15 +96,15 @@ class DatabaseInfo(BaseModel):
|
||||||
packages: dict[str, str] = Field(default_factory=dict)
|
packages: dict[str, str] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
async def gather_database_info(config: AppConfig, db_path: Path) -> DatabaseInfo:
|
async def gather_database_info(location: Path | str, config: AppConfig) -> DatabaseInfo:
|
||||||
"""Collect read-only database state without going through Store, so a
|
"""Collect read-only database state without going through Store, so a
|
||||||
database missing tables (e.g. pre-migration) still reports what it can."""
|
database missing tables (e.g. pre-migration) still reports what it can."""
|
||||||
from haiku.rag.store.upgrades import get_pending_upgrades
|
from haiku.rag.store.upgrades import get_pending_upgrades
|
||||||
from haiku.rag.utils import get_package_versions
|
from haiku.rag.utils import get_package_versions
|
||||||
|
|
||||||
display_path = config.lancedb.uri or str(db_path)
|
display_path = str(location)
|
||||||
|
|
||||||
db = await connect_lancedb(config, db_path)
|
db = await connect_lancedb(location, config)
|
||||||
stats = await get_database_stats(db)
|
stats = await get_database_stats(db)
|
||||||
|
|
||||||
if not any(entry["exists"] for entry in stats.values()):
|
if not any(entry["exists"] for entry in stats.values()):
|
||||||
|
|
|
||||||
|
|
@ -146,10 +146,10 @@ class SearchResult(BaseModel):
|
||||||
include the metadata of any other chunks merged with it. Never part of
|
include the metadata of any other chunks merged with it. Never part of
|
||||||
``format_for_agent`` output.
|
``format_for_agent`` output.
|
||||||
|
|
||||||
``source`` names the configured database a result came from: the name from
|
``source`` names the database a result came from: the name from
|
||||||
``lancedb.databases``, never a path or URI, so a location cannot travel in a
|
``lancedb.databases`` or a path's stem, never a path or URI, so a location
|
||||||
result, a citation or a log. It is None only where no database is named, as
|
cannot travel in a result, a citation or a log. Every result a search
|
||||||
with the single ``lancedb.uri``.
|
produces carries it; None only on a result built by hand.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
content: str
|
content: str
|
||||||
|
|
|
||||||
|
|
@ -24,9 +24,9 @@ class Citation(BaseModel):
|
||||||
``chunk_ids`` lists the ids of all chunks whose expansion ranges merged
|
``chunk_ids`` lists the ids of all chunks whose expansion ranges merged
|
||||||
into the cited result (always includes ``chunk_id``).
|
into the cited result (always includes ``chunk_id``).
|
||||||
|
|
||||||
``source`` names the configured database the cited chunk came from: the name
|
``source`` names the database the cited chunk came from: the name from
|
||||||
from ``lancedb.databases``, never a path or URI. It is None only where no
|
``lancedb.databases`` or a path's stem, never a path or URI. None only on a
|
||||||
database is named, as with the single ``lancedb.uri``.
|
citation resolved from a hand-built result.
|
||||||
|
|
||||||
``doc_item_refs`` are the ``self_ref`` values of every item in the cited
|
``doc_item_refs`` are the ``self_ref`` values of every item in the cited
|
||||||
content — the exact items the model saw. Visual grounding resolves bounding
|
content — the exact items the model saw. Visual grounding resolves bounding
|
||||||
|
|
|
||||||
|
|
@ -14,10 +14,10 @@ class Document(BaseModel):
|
||||||
"""
|
"""
|
||||||
Represents a document with an ID, content, and metadata.
|
Represents a document with an ID, content, and metadata.
|
||||||
|
|
||||||
``source`` names the configured database a document came from: the name
|
``source`` names the database a document came from: the name from
|
||||||
from ``lancedb.databases``, never a path or URI. It is None where no
|
``lancedb.databases`` or a path's stem, never a path or URI. Every document
|
||||||
database is named, as with the single ``lancedb.uri``, and is never
|
a database returns carries it; it is never persisted, and None only on a
|
||||||
persisted.
|
document built by hand.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
id: str | None = None
|
id: str | None = None
|
||||||
|
|
|
||||||
|
|
@ -77,8 +77,11 @@ async def _apply_split_document_meta(store: Store) -> None:
|
||||||
Exception
|
Exception
|
||||||
): # pragma: no cover - defensive; stats() failure shouldn't block the split
|
): # pragma: no cover - defensive; stats() failure shouldn't block the split
|
||||||
live_bytes = 0
|
live_bytes = 0
|
||||||
free_bytes = shutil.disk_usage(store.db_path).free
|
# A database behind a URI has no local disk to run out of.
|
||||||
if live_bytes and free_bytes < live_bytes:
|
free_bytes = (
|
||||||
|
shutil.disk_usage(store.db_path).free if store.db_path is not None else None
|
||||||
|
)
|
||||||
|
if live_bytes and free_bytes is not None and free_bytes < live_bytes:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Skipping post-migration vacuum: need ~%.2f GB free to compact the "
|
"Skipping post-migration vacuum: need ~%.2f GB free to compact the "
|
||||||
"documents table, have %.2f GB. Run `haiku-rag vacuum` once you have "
|
"documents table, have %.2f GB. Run `haiku-rag vacuum` once you have "
|
||||||
|
|
|
||||||
|
|
@ -549,17 +549,15 @@ def raise_missing_extra(module: str, extra: str, exc: ModuleNotFoundError) -> No
|
||||||
) from exc
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
def locate_database(location: str) -> tuple[str, Path | None]:
|
def locate_database(location: str) -> Path | str:
|
||||||
"""Split a configured location into (uri, db_path).
|
"""A configured location as a URI, or as a local path.
|
||||||
|
|
||||||
A value with a scheme is a `lancedb.uri`; anything else is a local path.
|
A value with a scheme is a URI, which `ConnectionMode` opens without the
|
||||||
`ConnectionMode` classifies a `uri` as object storage and opens it without
|
existence check a local database gets; anything else is a local path.
|
||||||
the existence check a local database gets, so a local path never travels
|
|
||||||
as one.
|
|
||||||
"""
|
"""
|
||||||
if "://" in location:
|
if "://" in location:
|
||||||
return location, None
|
return location
|
||||||
return "", Path(location)
|
return Path(location)
|
||||||
|
|
||||||
|
|
||||||
def get_default_data_dir() -> Path:
|
def get_default_data_dir() -> Path:
|
||||||
|
|
|
||||||
|
|
@ -98,14 +98,10 @@ def _placed(capability) -> "Path | None":
|
||||||
return ref.db_path
|
return ref.db_path
|
||||||
|
|
||||||
|
|
||||||
def test_capability_factories_resolve_environment_and_defaults(
|
def test_capability_factories_resolve_defaults(temp_db_path, monkeypatch):
|
||||||
temp_db_path, monkeypatch
|
"""The configuration places the database; the environment plays no part."""
|
||||||
):
|
|
||||||
config = AppConfig()
|
config = AppConfig()
|
||||||
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
|
||||||
assert _placed(create_rag(config=config)) == temp_db_path
|
|
||||||
|
|
||||||
monkeypatch.delenv("HAIKU_RAG_DB")
|
|
||||||
assert _placed(create_rag(config=config)) == (
|
assert _placed(create_rag(config=config)) == (
|
||||||
config.storage.data_dir / "haiku.rag.lancedb"
|
config.storage.data_dir / "haiku.rag.lancedb"
|
||||||
)
|
)
|
||||||
|
|
@ -124,30 +120,28 @@ class TestACapabilityFollowsTheConfiguredLocation:
|
||||||
"""A capability nobody handed a client opens one for itself, at the
|
"""A capability nobody handed a client opens one for itself, at the
|
||||||
database the configuration places."""
|
database the configuration places."""
|
||||||
|
|
||||||
def _config(self, tmp_path, uri: str) -> AppConfig:
|
def _config(self, tmp_path, location: str) -> AppConfig:
|
||||||
from haiku.rag.config.models import LanceDBConfig, StorageConfig
|
from haiku.rag.config.models import LanceDBConfig, StorageConfig
|
||||||
|
|
||||||
return AppConfig(
|
return AppConfig(
|
||||||
lancedb=LanceDBConfig(uri=uri),
|
lancedb=LanceDBConfig(databases={"notes": location}),
|
||||||
storage=StorageConfig(data_dir=tmp_path / "elsewhere"),
|
storage=StorageConfig(data_dir=tmp_path / "elsewhere"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_a_configured_uri_is_left_to_the_client(self, tmp_path):
|
def test_a_configured_location_is_the_capability_scope(self, tmp_path):
|
||||||
"""A path overrides a configured location, so the capability passes
|
|
||||||
None and the client resolves the configured URI."""
|
|
||||||
located = tmp_path / "notes.lancedb"
|
located = tmp_path / "notes.lancedb"
|
||||||
for factory in (create_rag, create_analysis):
|
for factory in (create_rag, create_analysis):
|
||||||
[local] = factory(
|
[local] = factory(
|
||||||
config=self._config(tmp_path, str(located))
|
config=self._config(tmp_path, str(located))
|
||||||
).scope.databases
|
).scope.databases
|
||||||
assert local == DatabaseRef.configured(None, str(located))
|
assert local == DatabaseRef("notes", located)
|
||||||
|
|
||||||
remote = self._config(tmp_path, "s3://bucket/one.lancedb")
|
remote = self._config(tmp_path, "s3://bucket/one.lancedb")
|
||||||
[ref] = factory(config=remote).scope.databases
|
[ref] = factory(config=remote).scope.databases
|
||||||
assert ref == DatabaseRef(None, "s3://bucket/one.lancedb", None)
|
assert ref == DatabaseRef("notes", "s3://bucket/one.lancedb")
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_it_opens_the_database_the_uri_places(self, tmp_path):
|
async def test_it_opens_the_database_the_configuration_places(self, tmp_path):
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
located = tmp_path / "notes.lancedb"
|
located = tmp_path / "notes.lancedb"
|
||||||
|
|
@ -162,19 +156,15 @@ class TestACapabilityFollowsTheConfiguredLocation:
|
||||||
finally:
|
finally:
|
||||||
await capability._close()
|
await capability._close()
|
||||||
|
|
||||||
def test_an_explicit_path_still_overrides_the_configured_uri(self, tmp_path):
|
def test_a_path_beside_the_configured_placement_is_refused(self, tmp_path):
|
||||||
|
from haiku.rag.store.exceptions import AmbiguousDatabaseError
|
||||||
|
|
||||||
config = self._config(tmp_path, str(tmp_path / "notes.lancedb"))
|
config = self._config(tmp_path, str(tmp_path / "notes.lancedb"))
|
||||||
chosen = tmp_path / "chosen.lancedb"
|
chosen = tmp_path / "chosen.lancedb"
|
||||||
|
|
||||||
assert _placed(create_rag(db_path=chosen, config=config)) == chosen
|
for factory in (create_rag, create_analysis):
|
||||||
|
with pytest.raises(AmbiguousDatabaseError, match="notes"):
|
||||||
def test_the_environment_still_overrides_the_configured_uri(
|
factory(db_path=chosen, config=config)
|
||||||
self, tmp_path, monkeypatch
|
|
||||||
):
|
|
||||||
config = self._config(tmp_path, "s3://bucket/one.lancedb")
|
|
||||||
monkeypatch.setenv("HAIKU_RAG_DB", str(tmp_path / "from-env.lancedb"))
|
|
||||||
|
|
||||||
assert _placed(create_rag(config=config)) == tmp_path / "from-env.lancedb"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -206,15 +196,15 @@ def test_domain_preamble_is_added_to_capability_instructions(temp_db_path):
|
||||||
|
|
||||||
|
|
||||||
def _single_database_client() -> AsyncMock:
|
def _single_database_client() -> AsyncMock:
|
||||||
"""A stand-in for a client covering one unnamed database.
|
"""A stand-in for a client covering one database.
|
||||||
|
|
||||||
`covers_multiple`, `source` and `clients_covering` answer as one unnamed
|
`covers_multiple`, `source` and `clients_covering` answer as one database
|
||||||
database does; a bare AsyncMock answers every attribute with a truthy Mock.
|
does; a bare AsyncMock answers every attribute with a truthy Mock.
|
||||||
"""
|
"""
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.source = None
|
client.source = "test"
|
||||||
client.clients_covering.return_value = [client]
|
client.clients_covering.return_value = [client]
|
||||||
return client
|
return client
|
||||||
|
|
||||||
|
|
@ -1691,22 +1681,18 @@ class TestMultipleCollectionsInstructions:
|
||||||
(create_rag, rag_text),
|
(create_rag, rag_text),
|
||||||
(create_analysis, analysis_text),
|
(create_analysis, analysis_text),
|
||||||
):
|
):
|
||||||
for config in (AppConfig(), self._config(alpha="/a.lancedb")):
|
one_at_a_path = factory(db_path=Path("/tmp/x.lancedb"), config=AppConfig())
|
||||||
capability = factory(db_path=Path("/tmp/x.lancedb"), config=config)
|
assert one_at_a_path.instruction_text == baseline()
|
||||||
assert capability.instruction_text == baseline()
|
one_configured = factory(config=self._config(alpha="/a.lancedb"))
|
||||||
|
assert one_configured.instruction_text == baseline()
|
||||||
|
|
||||||
def test_an_explicit_path_opens_one_database(self):
|
def test_a_path_beside_a_configured_set_is_refused(self):
|
||||||
"""A path names one database, whatever the configuration names."""
|
from haiku.rag.store.exceptions import AmbiguousDatabaseError
|
||||||
from haiku.rag.capabilities.analysis import instructions as analysis_text
|
|
||||||
from haiku.rag.capabilities.rag import instructions as rag_text
|
|
||||||
|
|
||||||
config = self._config(alpha="/a.lancedb", beta="/b.lancedb")
|
config = self._config(alpha="/a.lancedb", beta="/b.lancedb")
|
||||||
for factory, baseline in (
|
for factory in (create_rag, create_analysis):
|
||||||
(create_rag, rag_text),
|
with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"):
|
||||||
(create_analysis, analysis_text),
|
factory(db_path=Path("/tmp/one.lancedb"), config=config)
|
||||||
):
|
|
||||||
capability = factory(db_path=Path("/tmp/one.lancedb"), config=config)
|
|
||||||
assert capability.instruction_text == baseline()
|
|
||||||
|
|
||||||
def test_a_lent_client_covering_one_database_is_instructed_as_before(self):
|
def test_a_lent_client_covering_one_database_is_instructed_as_before(self):
|
||||||
from haiku.rag.capabilities.analysis import instructions as analysis_text
|
from haiku.rag.capabilities.analysis import instructions as analysis_text
|
||||||
|
|
|
||||||
|
|
@ -84,17 +84,18 @@ def test_chat_capabilities_read_the_named_database(tmp_path, monkeypatch):
|
||||||
from haiku.rag.chat import run_chat
|
from haiku.rag.chat import run_chat
|
||||||
|
|
||||||
run_chat(scope=DatabaseScope.resolve(config, database_name="b"))
|
run_chat(scope=DatabaseScope.resolve(config, database_name="b"))
|
||||||
|
named_scope = chat_app.call_args.kwargs["scope"]
|
||||||
[named] = chat_app.call_args.kwargs["capabilities"]
|
[named] = chat_app.call_args.kwargs["capabilities"]
|
||||||
|
|
||||||
run_chat(scope=DatabaseScope.resolve(config))
|
run_chat(scope=DatabaseScope.resolve(config))
|
||||||
|
covering_scope = chat_app.call_args.kwargs["scope"]
|
||||||
[covering] = chat_app.call_args.kwargs["capabilities"]
|
[covering] = chat_app.call_args.kwargs["capabilities"]
|
||||||
|
|
||||||
# The chat lends its own client, so this scope is the fallback: it places
|
# The app opens the scope it is handed and lends that client to the
|
||||||
# the named database alone.
|
# capabilities, which keep the configuration as the caller named it.
|
||||||
[placed] = named.scope.databases
|
assert named_scope.names == ("b",)
|
||||||
assert placed.db_path == tmp_path / "b.lancedb"
|
assert covering_scope.names == ("a", "b")
|
||||||
assert named.config.lancedb.databases == {}
|
assert set(named.config.lancedb.databases) == {"a", "b"}
|
||||||
assert covering.scope.names == ("a", "b")
|
|
||||||
assert set(covering.config.lancedb.databases) == {"a", "b"}
|
assert set(covering.config.lancedb.databases) == {"a", "b"}
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -158,8 +159,8 @@ def _make_mock_client():
|
||||||
# Covers one database; a bare AsyncMock answers `covers_multiple` with a
|
# Covers one database; a bare AsyncMock answers `covers_multiple` with a
|
||||||
# truthy Mock.
|
# truthy Mock.
|
||||||
mock_client.covers_multiple = False
|
mock_client.covers_multiple = False
|
||||||
mock_client.source_names = ()
|
mock_client.source_names = ("test",)
|
||||||
mock_client.source = None
|
mock_client.source = "test"
|
||||||
return mock_client
|
return mock_client
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -461,9 +462,9 @@ async def test_document_filter_updates_rag_state(temp_db_path: Path):
|
||||||
):
|
):
|
||||||
async with app.run_test():
|
async with app.run_test():
|
||||||
# The selection is document ids, so a repeated title cannot widen it.
|
# The selection is document ids, so a repeated title cannot widen it.
|
||||||
selected = [
|
selected: list[tuple[str | None, str]] = [
|
||||||
(None, "6f1c2d4e-0000-4000-8000-000000000001"),
|
("test", "6f1c2d4e-0000-4000-8000-000000000001"),
|
||||||
(None, "6f1c2d4e-0000-4000-8000-000000000002"),
|
("test", "6f1c2d4e-0000-4000-8000-000000000002"),
|
||||||
]
|
]
|
||||||
app.on_document_filter_modal_filter_changed(
|
app.on_document_filter_modal_filter_changed(
|
||||||
DocumentFilterModal.FilterChanged(selected)
|
DocumentFilterModal.FilterChanged(selected)
|
||||||
|
|
@ -477,7 +478,7 @@ async def test_document_filter_updates_rag_state(temp_db_path: Path):
|
||||||
assert rag_state.document_filter == expected_filter
|
assert rag_state.document_filter == expected_filter
|
||||||
assert rag_state.document_filter is not None
|
assert rag_state.document_filter is not None
|
||||||
assert "LIKE" not in rag_state.document_filter
|
assert "LIKE" not in rag_state.document_filter
|
||||||
# An unnamed database leaves the question unscoped by source.
|
# One database leaves the question unscoped by source.
|
||||||
assert rag_state.sources is None
|
assert rag_state.sources is None
|
||||||
|
|
||||||
# The state snapshot should also reflect the change
|
# The state snapshot should also reflect the change
|
||||||
|
|
@ -486,12 +487,15 @@ async def test_document_filter_updates_rag_state(temp_db_path: Path):
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_document_filter_narrows_sources_to_the_selection(temp_db_path: Path):
|
async def test_document_filter_narrows_sources_to_the_selection(temp_db_path: Path):
|
||||||
"""The filter carries ids, and `sources` restricts the question to the
|
"""Over a set, the filter carries ids and `sources` restricts the question
|
||||||
databases the selection names."""
|
to the databases the selection names."""
|
||||||
from haiku.rag.chat.app import RAG_STATE_NAMESPACE
|
from haiku.rag.chat.app import RAG_STATE_NAMESPACE
|
||||||
from haiku.rag.chat.widgets.document_filter_modal import DocumentFilterModal
|
from haiku.rag.chat.widgets.document_filter_modal import DocumentFilterModal
|
||||||
|
|
||||||
app, mock_client = _make_app_with_state(temp_db_path)
|
app, mock_client = _make_app_with_state(temp_db_path)
|
||||||
|
mock_client.covers_multiple = True
|
||||||
|
mock_client.source_names = ("alpha", "beta")
|
||||||
|
mock_client.source = None
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch("haiku.rag.chat.app.HaikuRAG") as _stub_rag,
|
patch("haiku.rag.chat.app.HaikuRAG") as _stub_rag,
|
||||||
|
|
@ -537,7 +541,7 @@ async def test_document_filter_cleared_when_empty(temp_db_path: Path):
|
||||||
async with app.run_test():
|
async with app.run_test():
|
||||||
# First set a filter
|
# First set a filter
|
||||||
app.on_document_filter_modal_filter_changed(
|
app.on_document_filter_modal_filter_changed(
|
||||||
DocumentFilterModal.FilterChanged([(None, "AI Overview")])
|
DocumentFilterModal.FilterChanged([("test", "AI Overview")])
|
||||||
)
|
)
|
||||||
rag_state = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE])
|
rag_state = RAGState.model_validate(app._state[RAG_STATE_NAMESPACE])
|
||||||
assert rag_state.document_filter is not None
|
assert rag_state.document_filter is not None
|
||||||
|
|
@ -684,6 +688,37 @@ class TestLendingTheClient:
|
||||||
assert borrowed == [client] * len(app._capabilities)
|
assert borrowed == [client] * len(app._capabilities)
|
||||||
assert borrowed
|
assert borrowed
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mounting_gives_every_capability_the_apps_scope(self, tmp_path):
|
||||||
|
"""A capability built over the configured set covers what the chat
|
||||||
|
selected once mounted: the analysis sandbox is built over that scope."""
|
||||||
|
from haiku.rag.chat.app import ChatApp
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
|
|
||||||
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(
|
||||||
|
databases={
|
||||||
|
"a": str(tmp_path / "a.lancedb"),
|
||||||
|
"b": str(tmp_path / "b.lancedb"),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
selected = DatabaseScope.resolve(config, database_name="b")
|
||||||
|
capability = create_capability(config=config)
|
||||||
|
assert capability.scope.covers_multiple
|
||||||
|
|
||||||
|
client = _make_mock_client()
|
||||||
|
app = ChatApp(scope=selected, capabilities=[capability], read_only=True)
|
||||||
|
with (
|
||||||
|
patch("haiku.rag.chat.app.HaikuRAG") as stub_rag,
|
||||||
|
_covering_returns(stub_rag, client),
|
||||||
|
):
|
||||||
|
async with app.run_test():
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert capability.scope == selected
|
||||||
|
|
||||||
|
|
||||||
class TestDocumentSelectionIdentity:
|
class TestDocumentSelectionIdentity:
|
||||||
"""Two documents can share a title, within a corpus and across databases, so
|
"""Two documents can share a title, within a corpus and across databases, so
|
||||||
|
|
@ -785,11 +820,23 @@ class TestDocumentSelectionIdentity:
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
((label, source, doc_id),) = _labelled(docs)
|
((label, source, doc_id),) = _labelled(docs, name_database=True)
|
||||||
box = DocumentCheckbox(label, source, doc_id, value=False)
|
box = DocumentCheckbox(label, source, doc_id, value=False)
|
||||||
|
|
||||||
assert str(box.label) == "Report [/red] (alpha [/x])"
|
assert str(box.label) == "Report [/red] (alpha [/x])"
|
||||||
|
|
||||||
|
def test_one_database_is_not_named_on_its_labels(self):
|
||||||
|
"""A single database names every document alike, so the label says
|
||||||
|
nothing a title does not."""
|
||||||
|
from haiku.rag.chat.widgets.document_filter_modal import _labelled
|
||||||
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
|
docs = [Document(id="id-one", content="", title="Report", source="test")]
|
||||||
|
|
||||||
|
((label, _, _),) = _labelled(docs)
|
||||||
|
|
||||||
|
assert label == "Report"
|
||||||
|
|
||||||
|
|
||||||
def test_a_citation_title_that_looks_like_markup_is_text():
|
def test_a_citation_title_that_looks_like_markup_is_text():
|
||||||
from rich.text import Text
|
from rich.text import Text
|
||||||
|
|
@ -1014,17 +1061,20 @@ class TestKeepingSelectionsReachable:
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
picked = [
|
picked = [
|
||||||
Document(id=f"sel-{i:04d}", content="", title=f"Selected {i:04d}")
|
Document(
|
||||||
|
id=f"sel-{i:04d}", content="", title=f"Selected {i:04d}", source="test"
|
||||||
|
)
|
||||||
for i in range(DOCUMENT_PAGE + 20)
|
for i in range(DOCUMENT_PAGE + 20)
|
||||||
]
|
]
|
||||||
by_id = {d.id: d for d in picked}
|
by_id = {d.id: d for d in picked}
|
||||||
matched = [
|
matched = [
|
||||||
Document(id=f"hit-{i}", content="", title=f"Hit {i}") for i in range(5)
|
Document(id=f"hit-{i}", content="", title=f"Hit {i}", source="test")
|
||||||
|
for i in range(5)
|
||||||
]
|
]
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.count_documents.return_value = 5
|
client.count_documents.return_value = 5
|
||||||
|
|
||||||
async def listing(limit=None, offset=0, filter=None):
|
async def listing(limit=None, offset=0, filter=None):
|
||||||
|
|
@ -1036,7 +1086,7 @@ class TestKeepingSelectionsReachable:
|
||||||
client.list_documents.side_effect = listing
|
client.list_documents.side_effect = listing
|
||||||
|
|
||||||
modal = DocumentFilterModal(
|
modal = DocumentFilterModal(
|
||||||
client=client, selected=[(None, d.id or "") for d in picked]
|
client=client, selected=[("test", d.id or "") for d in picked]
|
||||||
)
|
)
|
||||||
app, _ = _make_app(temp_db_path, client)
|
app, _ = _make_app(temp_db_path, client)
|
||||||
with (
|
with (
|
||||||
|
|
@ -1083,14 +1133,16 @@ class TestKeepingSelectionsReachable:
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
picked = [
|
picked = [
|
||||||
Document(id=f"sel-{i:04d}", content="", title=f"Selected {i:04d}")
|
Document(
|
||||||
|
id=f"sel-{i:04d}", content="", title=f"Selected {i:04d}", source="test"
|
||||||
|
)
|
||||||
for i in range(DOCUMENT_PAGE + 1)
|
for i in range(DOCUMENT_PAGE + 1)
|
||||||
]
|
]
|
||||||
by_id = {d.id: d for d in picked}
|
by_id = {d.id: d for d in picked}
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.count_documents.return_value = 0
|
client.count_documents.return_value = 0
|
||||||
|
|
||||||
async def listing(limit=None, offset=0, filter=None):
|
async def listing(limit=None, offset=0, filter=None):
|
||||||
|
|
@ -1102,7 +1154,7 @@ class TestKeepingSelectionsReachable:
|
||||||
client.list_documents.side_effect = listing
|
client.list_documents.side_effect = listing
|
||||||
|
|
||||||
modal = DocumentFilterModal(
|
modal = DocumentFilterModal(
|
||||||
client=client, selected=[(None, d.id or "") for d in picked]
|
client=client, selected=[("test", d.id or "") for d in picked]
|
||||||
)
|
)
|
||||||
app, _ = _make_app(temp_db_path, client)
|
app, _ = _make_app(temp_db_path, client)
|
||||||
with (
|
with (
|
||||||
|
|
@ -1134,7 +1186,9 @@ class TestKeepingSelectionsReachable:
|
||||||
# The row is gone from the listing, not merely unchecked.
|
# The row is gone from the listing, not merely unchecked.
|
||||||
assert "sel-0200" not in remaining
|
assert "sel-0200" not in remaining
|
||||||
assert len(remaining) == DOCUMENT_PAGE
|
assert len(remaining) == DOCUMENT_PAGE
|
||||||
assert modal._selected == {(None, d.id) for d in picked} - {(None, "sel-0200")}
|
assert modal._selected == {("test", d.id) for d in picked} - {
|
||||||
|
("test", "sel-0200")
|
||||||
|
}
|
||||||
# The page it was on no longer exists, so the modal does not report it.
|
# The page it was on no longer exists, so the modal does not report it.
|
||||||
assert modal._page == 0
|
assert modal._page == 0
|
||||||
assert "page" not in footer
|
assert "page" not in footer
|
||||||
|
|
@ -1154,10 +1208,10 @@ class TestKeepingSelectionsReachable:
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.count_documents.return_value = DOCUMENT_PAGE * 2
|
client.count_documents.return_value = DOCUMENT_PAGE * 2
|
||||||
client.list_documents.return_value = [
|
client.list_documents.return_value = [
|
||||||
Document(id="d1", content="", title="One")
|
Document(id="d1", content="", title="One", source="test")
|
||||||
]
|
]
|
||||||
|
|
||||||
modal = DocumentFilterModal(client=client)
|
modal = DocumentFilterModal(client=client)
|
||||||
|
|
@ -1187,7 +1241,7 @@ class TestKeepingSelectionsReachable:
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.list_documents.return_value = []
|
client.list_documents.return_value = []
|
||||||
client.count_documents.return_value = 0
|
client.count_documents.return_value = 0
|
||||||
|
|
||||||
|
|
@ -1222,10 +1276,10 @@ class TestKeepingSelectionsReachable:
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.list_documents.return_value = [
|
client.list_documents.return_value = [
|
||||||
Document(id="id-one", content="", title="Capital region"),
|
Document(id="id-one", content="", title="Capital region", source="test"),
|
||||||
Document(id="id-two", content="", title="Nobel laureates"),
|
Document(id="id-two", content="", title="Nobel laureates", source="test"),
|
||||||
]
|
]
|
||||||
client.count_documents.return_value = DOCUMENT_PAGE * 2
|
client.count_documents.return_value = DOCUMENT_PAGE * 2
|
||||||
|
|
||||||
|
|
@ -1328,10 +1382,10 @@ class TestDocumentSearchFilter:
|
||||||
|
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("test",)
|
||||||
client.list_documents.return_value = [
|
client.list_documents.return_value = [
|
||||||
Document(id="id-one", content="", title="Capital region"),
|
Document(id="id-one", content="", title="Capital region", source="test"),
|
||||||
Document(id="id-two", content="", title="Nobel laureates"),
|
Document(id="id-two", content="", title="Nobel laureates", source="test"),
|
||||||
]
|
]
|
||||||
client.count_documents.return_value = 2
|
client.count_documents.return_value = 2
|
||||||
|
|
||||||
|
|
@ -1347,7 +1401,9 @@ class TestDocumentSearchFilter:
|
||||||
assert len(list(modal.query(DocumentCheckbox))) == 2
|
assert len(list(modal.query(DocumentCheckbox))) == 2
|
||||||
|
|
||||||
client.list_documents.return_value = [
|
client.list_documents.return_value = [
|
||||||
Document(id="id-two", content="", title="Nobel laureates"),
|
Document(
|
||||||
|
id="id-two", content="", title="Nobel laureates", source="test"
|
||||||
|
),
|
||||||
]
|
]
|
||||||
client.count_documents.return_value = 1
|
client.count_documents.return_value = 1
|
||||||
await modal.on_input_submitted(Input.Submitted(Input(), "Nobel"))
|
await modal.on_input_submitted(Input.Submitted(Input(), "Nobel"))
|
||||||
|
|
|
||||||
|
|
@ -838,11 +838,11 @@ async def test_database_503_when_no_database_is_configured(state):
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_the_report_follows_a_configured_uri(tmp_path, jobs, sync):
|
async def test_the_report_follows_a_configured_uri(tmp_path, jobs, sync):
|
||||||
"""A configured `lancedb.uri` places the database, so the report opens that
|
"""A configured location places the database, so the report opens that
|
||||||
and not the local default."""
|
and not the local default."""
|
||||||
db_path = tmp_path / "configured.lancedb"
|
db_path = tmp_path / "configured.lancedb"
|
||||||
await _seed_lancedb(db_path)
|
await _seed_lancedb(db_path)
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri=str(db_path)))
|
config = AppConfig(lancedb=LanceDBConfig(databases={"configured": str(db_path)}))
|
||||||
state = APIState(
|
state = APIState(
|
||||||
config=config,
|
config=config,
|
||||||
job_repo=jobs,
|
job_repo=jobs,
|
||||||
|
|
|
||||||
|
|
@ -504,31 +504,58 @@ def test_cli_entry_point_exits_on_store_state_errors(monkeypatch, capsys, error)
|
||||||
|
|
||||||
class TestPlacingTheIngesterDatabase:
|
class TestPlacingTheIngesterDatabase:
|
||||||
"""The ingester writes wherever the configuration places the database, and
|
"""The ingester writes wherever the configuration places the database, and
|
||||||
resolves that once. A path is an explicit override of a configured
|
resolves that once. `--db PATH` is the operator's explicit override and
|
||||||
`lancedb.uri`, so no local default stands in for one."""
|
constructs the scope directly."""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _app(config: AppConfig, db_path=None):
|
def _app(config: AppConfig, db_path=None):
|
||||||
from haiku.rag.ingester.app import IngesterApp
|
from haiku.rag.ingester.app import IngesterApp
|
||||||
|
from haiku.rag.ingester.cli import _scope_for
|
||||||
|
|
||||||
return IngesterApp(config=config, db_path=db_path)
|
return IngesterApp(config=config, scope=_scope_for(db_path))
|
||||||
|
|
||||||
def test_a_configured_uri_becomes_the_scope(self, tmp_path):
|
def test_a_configured_location_becomes_the_scope(self, tmp_path):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/prod.lancedb"))
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(databases={"prod": "s3://bucket/prod.lancedb"})
|
||||||
|
)
|
||||||
|
|
||||||
[ref] = self._app(config)._scope.databases
|
[ref] = self._app(config)._scope.databases
|
||||||
|
|
||||||
assert ref.uri == "s3://bucket/prod.lancedb"
|
assert ref.name == "prod"
|
||||||
|
assert ref.location == "s3://bucket/prod.lancedb"
|
||||||
assert ref.db_path is None
|
assert ref.db_path is None
|
||||||
|
|
||||||
def test_an_override_names_the_database(self, tmp_path):
|
def test_an_override_names_the_database_by_its_stem(self, tmp_path):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/prod.lancedb"))
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(databases={"prod": "s3://bucket/prod.lancedb"})
|
||||||
|
)
|
||||||
override = tmp_path / "local.lancedb"
|
override = tmp_path / "local.lancedb"
|
||||||
|
|
||||||
[ref] = self._app(config, override)._scope.databases
|
[ref] = self._app(config, override)._scope.databases
|
||||||
|
|
||||||
assert ref.db_path == override
|
assert ref.name == "local"
|
||||||
assert ref.uri == ""
|
assert ref.location == override
|
||||||
|
|
||||||
|
def test_the_override_is_the_cli_s_alone(self, tmp_path):
|
||||||
|
"""`IngesterApp` takes a resolved scope, so a Python caller has no path
|
||||||
|
to slip past the configuration; only the CLI constructs one."""
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
from haiku.rag.ingester.cli import _scope_for
|
||||||
|
|
||||||
|
assert _scope_for(None) is None
|
||||||
|
assert _scope_for(tmp_path / "local.lancedb") == DatabaseScope.at(
|
||||||
|
tmp_path / "local.lancedb"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_a_path_without_a_stem_is_a_usage_error(self):
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import typer
|
||||||
|
|
||||||
|
from haiku.rag.ingester.cli import _scope_for
|
||||||
|
|
||||||
|
with pytest.raises(typer.BadParameter, match="no name"):
|
||||||
|
_scope_for(Path("/"))
|
||||||
|
|
||||||
def test_one_configured_database_is_accepted(self, tmp_path):
|
def test_one_configured_database_is_accepted(self, tmp_path):
|
||||||
"""A one-entry mapping names which database to write."""
|
"""A one-entry mapping names which database to write."""
|
||||||
|
|
@ -567,6 +594,10 @@ class TestPlacingTheIngesterDatabase:
|
||||||
f" a: {tmp_path / 'a.lancedb'}\n"
|
f" a: {tmp_path / 'a.lancedb'}\n"
|
||||||
f" b: {tmp_path / 'b.lancedb'}\n"
|
f" b: {tmp_path / 'b.lancedb'}\n"
|
||||||
)
|
)
|
||||||
|
import haiku.rag.config as config_module
|
||||||
|
|
||||||
|
# The CLI caches the loaded configuration process-wide.
|
||||||
|
monkeypatch.setattr(config_module, "_config", None)
|
||||||
monkeypatch.setattr(sys, "argv", ["haiku-ingester", "run-batch"])
|
monkeypatch.setattr(sys, "argv", ["haiku-ingester", "run-batch"])
|
||||||
monkeypatch.setenv("HAIKU_RAG_CONFIG_PATH", str(config_file))
|
monkeypatch.setenv("HAIKU_RAG_CONFIG_PATH", str(config_file))
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,7 @@ import pytest
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.client.exceptions import UnsupportedSourceError
|
from haiku.rag.client.exceptions import UnsupportedSourceError
|
||||||
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.config import (
|
from haiku.rag.config import (
|
||||||
APIConfig,
|
APIConfig,
|
||||||
AppConfig,
|
AppConfig,
|
||||||
|
|
@ -121,7 +122,7 @@ async def test_run_batch_drains_upserts(tmp_path, use_client):
|
||||||
use_client(client)
|
use_client(client)
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch()
|
).run_batch()
|
||||||
|
|
||||||
assert report.succeeded == 2
|
assert report.succeeded == 2
|
||||||
|
|
@ -146,7 +147,7 @@ async def test_run_batch_reports_progress(tmp_path, use_client):
|
||||||
progress = []
|
progress = []
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch(progress_callback=progress.append)
|
).run_batch(progress_callback=progress.append)
|
||||||
|
|
||||||
assert report.succeeded == 2
|
assert report.succeeded == 2
|
||||||
|
|
@ -171,7 +172,9 @@ async def test_run_batch_prunes_orphans(tmp_path, use_client):
|
||||||
db_path = tmp_path / "db.lancedb"
|
db_path = tmp_path / "db.lancedb"
|
||||||
|
|
||||||
# First batch ingests both files and records sync_state for each.
|
# First batch ingests both files and records sync_state for each.
|
||||||
first = await IngesterApp(config=config, db_path=db_path).run_batch()
|
first = await IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(db_path)
|
||||||
|
).run_batch()
|
||||||
assert first.succeeded == 2
|
assert first.succeeded == 2
|
||||||
client.delete_document.assert_not_awaited()
|
client.delete_document.assert_not_awaited()
|
||||||
|
|
||||||
|
|
@ -179,7 +182,9 @@ async def test_run_batch_prunes_orphans(tmp_path, use_client):
|
||||||
# sync_state but not on disk -> enqueues a DELETE for it.
|
# sync_state but not on disk -> enqueues a DELETE for it.
|
||||||
(tmp_path / "b.md").unlink()
|
(tmp_path / "b.md").unlink()
|
||||||
|
|
||||||
second = await IngesterApp(config=config, db_path=db_path).run_batch()
|
second = await IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(db_path)
|
||||||
|
).run_batch()
|
||||||
|
|
||||||
# a.md is unchanged (same mtime) so it's not re-ingested; only the orphan
|
# a.md is unchanged (same mtime) so it's not re-ingested; only the orphan
|
||||||
# delete runs.
|
# delete runs.
|
||||||
|
|
@ -198,7 +203,7 @@ async def test_run_batch_reports_dead_on_permanent_failure(tmp_path, use_client)
|
||||||
use_client(client)
|
use_client(client)
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch()
|
).run_batch()
|
||||||
|
|
||||||
assert report.succeeded == 0
|
assert report.succeeded == 0
|
||||||
|
|
@ -217,7 +222,9 @@ async def test_run_batch_recovered_doc_is_not_counted_as_dead(tmp_path, use_clie
|
||||||
failing = _mock_client()
|
failing = _mock_client()
|
||||||
failing.create_document_from_source.side_effect = UnsupportedSourceError("nope")
|
failing.create_document_from_source.side_effect = UnsupportedSourceError("nope")
|
||||||
use_client(failing)
|
use_client(failing)
|
||||||
first = await IngesterApp(config=config, db_path=db_path).run_batch()
|
first = await IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(db_path)
|
||||||
|
).run_batch()
|
||||||
assert first.dead == 1
|
assert first.dead == 1
|
||||||
|
|
||||||
healthy = _mock_client()
|
healthy = _mock_client()
|
||||||
|
|
@ -226,7 +233,9 @@ async def test_run_batch_recovered_doc_is_not_counted_as_dead(tmp_path, use_clie
|
||||||
# records the revision in sync_state, so a plain re-run no longer retries an
|
# records the revision in sync_state, so a plain re-run no longer retries an
|
||||||
# unchanged file — recovery needs the content (mtime) to change.
|
# unchanged file — recovery needs the content (mtime) to change.
|
||||||
(tmp_path / "a.md").write_text("hello again")
|
(tmp_path / "a.md").write_text("hello again")
|
||||||
second = await IngesterApp(config=config, db_path=db_path).run_batch()
|
second = await IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(db_path)
|
||||||
|
).run_batch()
|
||||||
assert second.dead == 0
|
assert second.dead == 0
|
||||||
assert second.succeeded == 1
|
assert second.succeeded == 1
|
||||||
|
|
||||||
|
|
@ -248,7 +257,7 @@ async def test_run_batch_reports_failed_sweep(
|
||||||
|
|
||||||
with caplog.at_level("ERROR", logger="haiku.rag.ingester.pollers.base"):
|
with caplog.at_level("ERROR", logger="haiku.rag.ingester.pollers.base"):
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch()
|
).run_batch()
|
||||||
|
|
||||||
assert report.failed_sweeps == ["local"]
|
assert report.failed_sweeps == ["local"]
|
||||||
|
|
@ -264,7 +273,7 @@ async def test_run_batch_empty_source_returns_immediately(tmp_path, use_client):
|
||||||
|
|
||||||
report = await asyncio.wait_for(
|
report = await asyncio.wait_for(
|
||||||
IngesterApp(
|
IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch(),
|
).run_batch(),
|
||||||
timeout=5.0,
|
timeout=5.0,
|
||||||
)
|
)
|
||||||
|
|
@ -287,7 +296,9 @@ async def test_run_batch_dry_run_reports_manifest_without_mutating_queue(tmp_pat
|
||||||
finally:
|
finally:
|
||||||
await engine.dispose()
|
await engine.dispose()
|
||||||
|
|
||||||
report = await IngesterApp(config=config, db_path=db_path).run_batch_dry_run()
|
report = await IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(db_path)
|
||||||
|
).run_batch_dry_run()
|
||||||
|
|
||||||
assert report.failed_sweeps == []
|
assert report.failed_sweeps == []
|
||||||
assert report.manifest.version == 1
|
assert report.manifest.version == 1
|
||||||
|
|
@ -322,7 +333,7 @@ async def test_run_batch_from_manifest_drains_changes_without_sweeping(
|
||||||
monkeypatch.setattr(PollerManager, "sweep_all", sweep_all)
|
monkeypatch.setattr(PollerManager, "sweep_all", sweep_all)
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(
|
).run_batch_from_manifest(
|
||||||
_manifest(
|
_manifest(
|
||||||
BatchChange(
|
BatchChange(
|
||||||
|
|
@ -350,7 +361,7 @@ async def test_run_batch_from_manifest_rejects_stale_upsert_revision(
|
||||||
use_client(client)
|
use_client(client)
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(
|
).run_batch_from_manifest(
|
||||||
_manifest(
|
_manifest(
|
||||||
BatchChange(
|
BatchChange(
|
||||||
|
|
@ -378,7 +389,7 @@ async def test_run_batch_from_manifest_delete_uses_manifest_even_if_file_reappea
|
||||||
use_client(client)
|
use_client(client)
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(
|
).run_batch_from_manifest(
|
||||||
_manifest(
|
_manifest(
|
||||||
BatchChange(
|
BatchChange(
|
||||||
|
|
@ -431,7 +442,7 @@ async def test_run_batch_from_manifest_resumes_same_manifest_work(tmp_path, use_
|
||||||
await engine.dispose()
|
await engine.dispose()
|
||||||
|
|
||||||
report = await IngesterApp(
|
report = await IngesterApp(
|
||||||
config=config, db_path=tmp_path / "db.lancedb"
|
config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(manifest)
|
).run_batch_from_manifest(manifest)
|
||||||
|
|
||||||
assert report.succeeded == 1
|
assert report.succeeded == 1
|
||||||
|
|
@ -456,7 +467,7 @@ async def test_run_batch_from_manifest_rejects_non_manifest_pending_work(
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="non-manifest pending work"):
|
with pytest.raises(ValueError, match="non-manifest pending work"):
|
||||||
await IngesterApp(
|
await IngesterApp(
|
||||||
config=config, db_path=tmp_path / "db.lancedb"
|
config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(
|
).run_batch_from_manifest(
|
||||||
_manifest(
|
_manifest(
|
||||||
BatchChange(
|
BatchChange(
|
||||||
|
|
@ -505,7 +516,7 @@ async def test_run_batch_from_manifest_rejects_different_manifest_pending_work(
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="non-manifest pending work"):
|
with pytest.raises(ValueError, match="non-manifest pending work"):
|
||||||
await IngesterApp(
|
await IngesterApp(
|
||||||
config=config, db_path=tmp_path / "db.lancedb"
|
config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(manifest)
|
).run_batch_from_manifest(manifest)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -526,7 +537,7 @@ async def test_run_batch_from_manifest_rejects_unrelated_pending_work(
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="non-manifest pending work"):
|
with pytest.raises(ValueError, match="non-manifest pending work"):
|
||||||
await IngesterApp(
|
await IngesterApp(
|
||||||
config=config, db_path=tmp_path / "db.lancedb"
|
config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(
|
).run_batch_from_manifest(
|
||||||
_manifest(
|
_manifest(
|
||||||
BatchChange(
|
BatchChange(
|
||||||
|
|
@ -554,7 +565,7 @@ async def test_run_batch_from_manifest_rejects_duplicate_changes(tmp_path, use_c
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="duplicate change"):
|
with pytest.raises(ValueError, match="duplicate change"):
|
||||||
await IngesterApp(
|
await IngesterApp(
|
||||||
config=_config(tmp_path), db_path=tmp_path / "db.lancedb"
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
).run_batch_from_manifest(_manifest(change, change))
|
).run_batch_from_manifest(_manifest(change, change))
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -583,7 +594,9 @@ async def test_run_batch_aborts_when_all_workers_die(
|
||||||
|
|
||||||
with caplog.at_level("ERROR", logger="haiku.rag.ingester.app"):
|
with caplog.at_level("ERROR", logger="haiku.rag.ingester.app"):
|
||||||
report = await asyncio.wait_for(
|
report = await asyncio.wait_for(
|
||||||
IngesterApp(config=config, db_path=tmp_path / "db.lancedb").run_batch(),
|
IngesterApp(
|
||||||
|
config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
|
).run_batch(),
|
||||||
timeout=10.0,
|
timeout=10.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -608,7 +621,7 @@ async def test_serve_starts_workers_pollers_and_shuts_down(tmp_path, use_client,
|
||||||
use_client(_mock_client())
|
use_client(_mock_client())
|
||||||
config = _config(tmp_path)
|
config = _config(tmp_path)
|
||||||
config.ingester.api = APIConfig(enabled=api, host="127.0.0.1", port=0)
|
config.ingester.api = APIConfig(enabled=api, host="127.0.0.1", port=0)
|
||||||
app = IngesterApp(config=config, db_path=tmp_path / "db.lancedb")
|
app = IngesterApp(config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb"))
|
||||||
|
|
||||||
task = asyncio.create_task(app.serve(api=api))
|
task = asyncio.create_task(app.serve(api=api))
|
||||||
try:
|
try:
|
||||||
|
|
@ -653,7 +666,7 @@ async def test_stop_pool_warns_when_shutdown_grace_elapses(tmp_path, caplog):
|
||||||
"""When a worker doesn't stop within the shutdown grace, _stop_pool logs a
|
"""When a worker doesn't stop within the shutdown grace, _stop_pool logs a
|
||||||
warning and still drains any pending cancel-cleanup releases."""
|
warning and still drains any pending cancel-cleanup releases."""
|
||||||
config = _config(tmp_path, shutdown_grace_s=0.01)
|
config = _config(tmp_path, shutdown_grace_s=0.01)
|
||||||
app = IngesterApp(config=config, db_path=tmp_path / "db.lancedb")
|
app = IngesterApp(config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb"))
|
||||||
pool = _SlowPool()
|
pool = _SlowPool()
|
||||||
app._pool = pool
|
app._pool = pool
|
||||||
|
|
||||||
|
|
@ -694,7 +707,9 @@ async def test_run_batch_closes_sources_after_pool_stops(
|
||||||
):
|
):
|
||||||
(tmp_path / "a.md").write_text("hello")
|
(tmp_path / "a.md").write_text("hello")
|
||||||
use_client(_mock_client())
|
use_client(_mock_client())
|
||||||
app = IngesterApp(config=_config(tmp_path), db_path=tmp_path / "db.lancedb")
|
app = IngesterApp(
|
||||||
|
config=_config(tmp_path), scope=DatabaseScope.at(tmp_path / "db.lancedb")
|
||||||
|
)
|
||||||
order = _record_close_order(monkeypatch)
|
order = _record_close_order(monkeypatch)
|
||||||
|
|
||||||
await app.run_batch()
|
await app.run_batch()
|
||||||
|
|
@ -707,7 +722,7 @@ async def test_serve_closes_sources_after_pool_stops(tmp_path, use_client, monke
|
||||||
use_client(_mock_client())
|
use_client(_mock_client())
|
||||||
config = _config(tmp_path)
|
config = _config(tmp_path)
|
||||||
config.ingester.api = APIConfig(enabled=False)
|
config.ingester.api = APIConfig(enabled=False)
|
||||||
app = IngesterApp(config=config, db_path=tmp_path / "db.lancedb")
|
app = IngesterApp(config=config, scope=DatabaseScope.at(tmp_path / "db.lancedb"))
|
||||||
order = _record_close_order(monkeypatch)
|
order = _record_close_order(monkeypatch)
|
||||||
|
|
||||||
task = asyncio.create_task(app.serve(api=False))
|
task = asyncio.create_task(app.serve(api=False))
|
||||||
|
|
|
||||||
|
|
@ -40,8 +40,8 @@ async def _restore_embedder(config, name, *, provider=None, model_name=None):
|
||||||
|
|
||||||
import lancedb
|
import lancedb
|
||||||
|
|
||||||
_, db_path = locate_database(config.lancedb.databases[name])
|
db_path = locate_database(config.lancedb.databases[name])
|
||||||
assert db_path is not None
|
assert not isinstance(db_path, str)
|
||||||
db = await lancedb.connect_async(str(db_path.resolve()))
|
db = await lancedb.connect_async(str(db_path.resolve()))
|
||||||
table = await db.open_table("settings")
|
table = await db.open_table("settings")
|
||||||
rows = (
|
rows = (
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,6 @@ import pytest
|
||||||
from haiku.rag.capabilities._tools import search_corpus
|
from haiku.rag.capabilities._tools import search_corpus
|
||||||
from haiku.rag.capabilities.rag import RAGState, create_capability
|
from haiku.rag.capabilities.rag import RAGState, create_capability
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.client.scope import DatabaseScope
|
|
||||||
from haiku.rag.client.session import FederatedSession
|
from haiku.rag.client.session import FederatedSession
|
||||||
from haiku.rag.sandbox import AnalysisContext, Sandbox
|
from haiku.rag.sandbox import AnalysisContext, Sandbox
|
||||||
from haiku.rag.store.exceptions import UnknownDatabaseError
|
from haiku.rag.store.exceptions import UnknownDatabaseError
|
||||||
|
|
@ -163,8 +162,9 @@ class TestCollectionIdentityForTheModel:
|
||||||
|
|
||||||
assert "Collection" not in result.format_for_agent()
|
assert "Collection" not in result.format_for_agent()
|
||||||
|
|
||||||
def test_an_unnamed_collection_is_never_mentioned(self):
|
def test_a_hand_built_result_without_a_source_is_never_labelled(self):
|
||||||
"""Nothing to name, whatever the caller asked for."""
|
"""A result built by hand carries no source to name, whatever the
|
||||||
|
caller asked for."""
|
||||||
result = SearchResult(content="body", score=0.9, chunk_id="c1")
|
result = SearchResult(content="body", score=0.9, chunk_id="c1")
|
||||||
|
|
||||||
assert "Collection" not in result.format_for_agent(include_collection=True)
|
assert "Collection" not in result.format_for_agent(include_collection=True)
|
||||||
|
|
@ -256,12 +256,9 @@ class TestLendingANamedClient:
|
||||||
await _seed(config, "alpha", ["alpha document about cats"])
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
await _seed(config, "beta", ["beta document about cats"])
|
await _seed(config, "beta", ["beta document about cats"])
|
||||||
|
|
||||||
# `run_chat` derives these for a single-database scope.
|
# What `run_chat` builds: the capability's own scope is the set, and
|
||||||
scope = DatabaseScope.resolve(config, database_name="alpha")
|
# the lent client is what narrows it.
|
||||||
one_config, one_path = scope.databases[0].connection(config)
|
capability = create_capability(config=config, defer_loading=False)
|
||||||
capability = create_capability(
|
|
||||||
db_path=one_path, config=one_config, defer_loading=False
|
|
||||||
)
|
|
||||||
|
|
||||||
async with HaikuRAG(config=config, sources=["alpha"]) as client:
|
async with HaikuRAG(config=config, sources=["alpha"]) as client:
|
||||||
# What `ChatApp.on_mount` does.
|
# What `ChatApp.on_mount` does.
|
||||||
|
|
|
||||||
|
|
@ -299,7 +299,7 @@ class TestCitationSource:
|
||||||
|
|
||||||
assert citation.chunk_id == "c1"
|
assert citation.chunk_id == "c1"
|
||||||
|
|
||||||
def test_a_single_database_citation_has_no_source(self):
|
def test_a_citation_from_a_hand_built_result_has_no_source(self):
|
||||||
result = SearchResult(
|
result = SearchResult(
|
||||||
content="body",
|
content="body",
|
||||||
score=0.9,
|
score=0.9,
|
||||||
|
|
|
||||||
|
|
@ -259,7 +259,8 @@ class TestLookupByIdentifier:
|
||||||
assert chunk is not None and chunk.content == "alpha one"
|
assert chunk is not None and chunk.content == "alpha one"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_unnamed_database_answers_to_no_name(self, temp_db_path):
|
async def test_a_database_at_a_path_answers_to_its_stem_alone(self, temp_db_path):
|
||||||
|
stem = temp_db_path.stem
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
docling = DoclingDocument(name="one")
|
docling = DoclingDocument(name="one")
|
||||||
docling.add_text(label=DocItemLabel.TEXT, text="body")
|
docling.add_text(label=DocItemLabel.TEXT, text="body")
|
||||||
|
|
@ -276,6 +277,8 @@ class TestLookupByIdentifier:
|
||||||
|
|
||||||
assert await rag.get_document_by_id(doc.id) is not None
|
assert await rag.get_document_by_id(doc.id) is not None
|
||||||
assert await rag.get_chunk_by_id(held.id) is not None
|
assert await rag.get_chunk_by_id(held.id) is not None
|
||||||
|
assert await rag.get_document_by_id(doc.id, stem) is not None
|
||||||
|
assert await rag.get_chunk_by_id(held.id, stem) is not None
|
||||||
with pytest.raises(UnknownDatabaseError):
|
with pytest.raises(UnknownDatabaseError):
|
||||||
await rag.get_document_by_id(doc.id, "alpha")
|
await rag.get_document_by_id(doc.id, "alpha")
|
||||||
with pytest.raises(UnknownDatabaseError):
|
with pytest.raises(UnknownDatabaseError):
|
||||||
|
|
@ -459,8 +462,9 @@ class TestDocumentsNameTheirDatabase:
|
||||||
assert by_uri is not None and by_uri.source == "alpha"
|
assert by_uri is not None and by_uri.source == "alpha"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_one_database_leaves_the_source_unset(self, tmp_path, temp_db_path):
|
async def test_one_database_at_a_path_is_named_by_its_stem(
|
||||||
"""Nothing names the database when there is only one to name."""
|
self, tmp_path, temp_db_path
|
||||||
|
):
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
dim = get_config().embeddings.model.vector_dim
|
dim = get_config().embeddings.model.vector_dim
|
||||||
doc = DoclingDocument(name="solo")
|
doc = DoclingDocument(name="solo")
|
||||||
|
|
@ -472,6 +476,7 @@ class TestDocumentsNameTheirDatabase:
|
||||||
)
|
)
|
||||||
|
|
||||||
[listed] = await rag.list_documents()
|
[listed] = await rag.list_documents()
|
||||||
assert listed.source is None
|
assert listed.source == temp_db_path.stem
|
||||||
assert listed.id is not None
|
assert listed.id is not None
|
||||||
assert (await rag.get_document_by_id(listed.id)).source is None
|
by_id = await rag.get_document_by_id(listed.id)
|
||||||
|
assert by_id is not None and by_id.source == temp_db_path.stem
|
||||||
|
|
|
||||||
|
|
@ -475,8 +475,27 @@ class TestFailureNaming:
|
||||||
assert caught.value.__cause__ is None
|
assert caught.value.__cause__ is None
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_unnamed_database_keeps_its_error(self, tmp_path):
|
async def test_a_missing_default_database_names_the_remedy(self, tmp_path):
|
||||||
"""Nothing named it, so there is no name to report."""
|
"""The location stays out of the message; the way to create the
|
||||||
|
database does not."""
|
||||||
|
from haiku.rag.config.models import AppConfig, StorageConfig
|
||||||
|
|
||||||
|
config = AppConfig(storage=StorageConfig(data_dir=tmp_path / "empty"))
|
||||||
|
|
||||||
|
with pytest.raises(SourceUnavailableError) as caught:
|
||||||
|
async with HaikuRAG(config=config):
|
||||||
|
pass
|
||||||
|
|
||||||
|
message = str(caught.value)
|
||||||
|
assert "database 'haiku.rag' does not exist" in message
|
||||||
|
assert "haiku-rag init" in message
|
||||||
|
assert "create=True" in message
|
||||||
|
assert str(tmp_path) not in message
|
||||||
|
assert caught.value.__cause__ is None
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_database_given_as_a_path_keeps_its_error(self, tmp_path):
|
||||||
|
"""The caller gave the path, so the error may name it."""
|
||||||
with pytest.raises(FileNotFoundError):
|
with pytest.raises(FileNotFoundError):
|
||||||
async with HaikuRAG(tmp_path / "nope.lancedb"):
|
async with HaikuRAG(tmp_path / "nope.lancedb"):
|
||||||
pass
|
pass
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@ from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
from haiku.rag.store.exceptions import (
|
from haiku.rag.store.exceptions import (
|
||||||
AmbiguousDatabaseError,
|
AmbiguousDatabaseError,
|
||||||
|
SourceUnavailableError,
|
||||||
UnknownDatabaseError,
|
UnknownDatabaseError,
|
||||||
)
|
)
|
||||||
from haiku.rag.utils import locate_database
|
from haiku.rag.utils import locate_database
|
||||||
|
|
@ -18,18 +19,13 @@ from tests.multi_db.helpers import (
|
||||||
|
|
||||||
|
|
||||||
class TestConfig:
|
class TestConfig:
|
||||||
def test_databases_and_uri_are_mutually_exclusive(self):
|
def test_databases_is_the_one_placement(self):
|
||||||
with pytest.raises(ValidationError, match="databases"):
|
|
||||||
LanceDBConfig(
|
|
||||||
uri="s3://b/one.lancedb", databases={"one": "s3://b/one.lancedb"}
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_databases_alone_is_fine(self):
|
|
||||||
config = LanceDBConfig(databases={"one": "s3://b/one.lancedb"})
|
config = LanceDBConfig(databases={"one": "s3://b/one.lancedb"})
|
||||||
assert config.databases == {"one": "s3://b/one.lancedb"}
|
assert config.databases == {"one": "s3://b/one.lancedb"}
|
||||||
|
|
||||||
def test_uri_alone_is_fine(self):
|
def test_uri_is_refused_naming_the_replacement(self):
|
||||||
assert LanceDBConfig(uri="s3://b/one.lancedb").databases == {}
|
with pytest.raises(ValidationError, match="lancedb.databases"):
|
||||||
|
LanceDBConfig.model_validate({"uri": "s3://b/one.lancedb"})
|
||||||
|
|
||||||
|
|
||||||
class TestNamingIsRequired:
|
class TestNamingIsRequired:
|
||||||
|
|
@ -51,18 +47,18 @@ class TestNamingIsRequired:
|
||||||
|
|
||||||
class TestNamingADatabaseDirectly:
|
class TestNamingADatabaseDirectly:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_explicit_db_path_wins_over_the_configured_set(
|
async def test_a_db_path_beside_the_configured_set_is_refused(
|
||||||
self, tmp_path, temp_db_path
|
self, tmp_path, temp_db_path
|
||||||
):
|
):
|
||||||
"""A caller that names a path means that database, not the configured
|
"""The configuration places databases; a path beside it is a second
|
||||||
set: the CLI resolves `--db` to one and must not fan out instead."""
|
placement, and the refusal names both."""
|
||||||
config = _config(tmp_path, ["alpha", "beta"])
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
await _seed(config, "alpha", ["alpha document about cats"])
|
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, config=config, create=True) as rag:
|
with pytest.raises(AmbiguousDatabaseError, match="alpha") as raised:
|
||||||
assert not rag.covers_multiple
|
async with HaikuRAG(temp_db_path, config=config, create=True):
|
||||||
assert rag.source is None
|
pass
|
||||||
assert rag.store.db_path == temp_db_path
|
assert str(temp_db_path) in str(raised.value)
|
||||||
|
assert not temp_db_path.exists()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_one_configured_database_is_opened_by_name(self, tmp_path):
|
async def test_one_configured_database_is_opened_by_name(self, tmp_path):
|
||||||
|
|
@ -79,42 +75,43 @@ class TestNamingADatabaseDirectly:
|
||||||
|
|
||||||
|
|
||||||
class TestOneConfiguredLocation:
|
class TestOneConfiguredLocation:
|
||||||
"""`lancedb.uri` places one unnamed database, at a URI or at a local path."""
|
"""One entry in `lancedb.databases` places one named database, at a URI or
|
||||||
|
at a local path."""
|
||||||
|
|
||||||
def _config(self, location) -> AppConfig:
|
def _config(self, location) -> AppConfig:
|
||||||
return AppConfig(lancedb=LanceDBConfig(uri=str(location)))
|
return AppConfig(lancedb=LanceDBConfig(databases={"notes": str(location)}))
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_a_local_uri_opens_the_configured_database(self, tmp_path):
|
async def test_a_local_location_opens_the_configured_database(self, tmp_path):
|
||||||
located = tmp_path / "notes.lancedb"
|
located = tmp_path / "notes.lancedb"
|
||||||
config = self._config(located)
|
config = self._config(located)
|
||||||
|
|
||||||
async with HaikuRAG(config=config, create=True) as rag:
|
async with HaikuRAG(config=config, create=True) as rag:
|
||||||
assert rag.store.db_path == located
|
assert rag.store.db_path == located
|
||||||
# It places a database without naming one: only `lancedb.databases`
|
assert rag.source == "notes"
|
||||||
# assigns the name results and citations carry.
|
|
||||||
assert rag.source is None
|
|
||||||
assert located.exists()
|
assert located.exists()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_explicit_path_overrides_a_local_uri(self, tmp_path):
|
async def test_a_path_beside_the_configured_database_is_refused(self, tmp_path):
|
||||||
"""`--db` overrides the configured location for one invocation."""
|
|
||||||
config = self._config(tmp_path / "configured.lancedb")
|
config = self._config(tmp_path / "configured.lancedb")
|
||||||
chosen = tmp_path / "chosen.lancedb"
|
chosen = tmp_path / "chosen.lancedb"
|
||||||
|
|
||||||
async with HaikuRAG(chosen, config=config, create=True) as rag:
|
with pytest.raises(AmbiguousDatabaseError, match="notes"):
|
||||||
assert rag.store.db_path == chosen
|
async with HaikuRAG(chosen, config=config, create=True):
|
||||||
assert chosen.exists()
|
pass
|
||||||
|
assert not chosen.exists()
|
||||||
assert not (tmp_path / "configured.lancedb").exists()
|
assert not (tmp_path / "configured.lancedb").exists()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_a_local_uri_that_does_not_exist_is_refused(self, tmp_path):
|
async def test_a_local_location_that_does_not_exist_is_refused(self, tmp_path):
|
||||||
"""A schemeless location is a local path and must exist."""
|
"""A schemeless location is a local path and must exist. The error names
|
||||||
|
the configured database, never its location."""
|
||||||
config = self._config(tmp_path / "typo.lancedb")
|
config = self._config(tmp_path / "typo.lancedb")
|
||||||
|
|
||||||
with pytest.raises(FileNotFoundError):
|
with pytest.raises(SourceUnavailableError, match="notes") as caught:
|
||||||
async with HaikuRAG(config=config):
|
async with HaikuRAG(config=config):
|
||||||
pass
|
pass
|
||||||
|
assert "typo.lancedb" not in str(caught.value)
|
||||||
assert not (tmp_path / "typo.lancedb").exists()
|
assert not (tmp_path / "typo.lancedb").exists()
|
||||||
|
|
||||||
def test_a_uri_with_a_scheme_stays_a_uri(self, tmp_path):
|
def test_a_uri_with_a_scheme_stays_a_uri(self, tmp_path):
|
||||||
|
|
@ -125,23 +122,54 @@ class TestOneConfiguredLocation:
|
||||||
config = self._config("s3://bucket/one.lancedb")
|
config = self._config("s3://bucket/one.lancedb")
|
||||||
|
|
||||||
[ref] = DatabaseScope.resolve(config).databases
|
[ref] = DatabaseScope.resolve(config).databases
|
||||||
one, db_path = ref.connection(config)
|
|
||||||
|
|
||||||
assert db_path is None
|
assert ref.location == "s3://bucket/one.lancedb"
|
||||||
assert ConnectionMode.from_config(one) == ConnectionMode.OBJECT_STORAGE
|
assert ConnectionMode.of(ref.location) == ConnectionMode.OBJECT_STORAGE
|
||||||
|
|
||||||
|
|
||||||
|
class TestSessionsOwnTheRef:
|
||||||
|
"""A session is built from the resolved reference and hands storage only
|
||||||
|
its location; the configuration it keeps is the one the caller named."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_session_opens_the_location_with_the_undivided_config(
|
||||||
|
self, tmp_path
|
||||||
|
):
|
||||||
|
from haiku.rag.client.session import SingleDatabaseSession
|
||||||
|
|
||||||
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
[ref] = DatabaseScope.resolve(config, database_name="alpha").databases
|
||||||
|
|
||||||
|
session = await SingleDatabaseSession(ref, config, read_only=True).open()
|
||||||
|
try:
|
||||||
|
assert session.source == "alpha"
|
||||||
|
assert session.location == ref.location
|
||||||
|
assert session.db_path == ref.location
|
||||||
|
assert session.store.location == ref.location
|
||||||
|
assert session.store._config is config
|
||||||
|
finally:
|
||||||
|
await session.aclose()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_client_keeps_the_configuration_it_was_given(self, tmp_path):
|
||||||
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
|
||||||
|
async with HaikuRAG(config=config, sources=["alpha"]) as rag:
|
||||||
|
assert rag._config is config
|
||||||
|
assert set(rag._config.lancedb.databases) == {"alpha", "beta"}
|
||||||
|
assert rag.store.location == tmp_path / "alpha.lancedb"
|
||||||
|
|
||||||
|
|
||||||
class TestLocate:
|
class TestLocate:
|
||||||
def test_a_scheme_is_a_uri(self):
|
def test_a_scheme_is_a_uri(self):
|
||||||
assert locate_database("s3://bucket/one.lancedb") == (
|
assert locate_database("s3://bucket/one.lancedb") == "s3://bucket/one.lancedb"
|
||||||
"s3://bucket/one.lancedb",
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_anything_else_is_a_local_path(self):
|
def test_anything_else_is_a_local_path(self):
|
||||||
uri, db_path = locate_database("/data/one.lancedb")
|
from pathlib import Path
|
||||||
assert uri == ""
|
|
||||||
assert db_path is not None and str(db_path) == "/data/one.lancedb"
|
assert locate_database("/data/one.lancedb") == Path("/data/one.lancedb")
|
||||||
|
|
||||||
|
|
||||||
class TestSelection:
|
class TestSelection:
|
||||||
|
|
@ -216,10 +244,39 @@ class TestPlacingADatabase:
|
||||||
assert {r.source for r in results} == {"alpha"}
|
assert {r.source for r in results} == {"alpha"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_unnamed_database_names_nothing(self, temp_db_path):
|
async def test_a_database_at_a_path_is_named_by_its_stem(self, temp_db_path):
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
assert rag.source_names == ()
|
assert rag.source_names == (temp_db_path.stem,)
|
||||||
assert rag.source is None
|
assert rag.source == temp_db_path.stem
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_the_default_database_is_selectable_by_name(self, tmp_path):
|
||||||
|
"""Nothing configured is the one entry `haiku.rag`, an ordinary
|
||||||
|
configured database that `sources` can name."""
|
||||||
|
from haiku.rag.config.models import StorageConfig
|
||||||
|
|
||||||
|
config = AppConfig(storage=StorageConfig(data_dir=tmp_path))
|
||||||
|
|
||||||
|
async with HaikuRAG(config=config, sources=["haiku.rag"], create=True) as rag:
|
||||||
|
assert rag.source == "haiku.rag"
|
||||||
|
assert rag.store.db_path == tmp_path / "haiku.rag.lancedb"
|
||||||
|
|
||||||
|
def test_coverage_is_known_before_the_client_opens(self, tmp_path):
|
||||||
|
"""Coverage is a fact of the resolved scope, readable before entering,
|
||||||
|
and `source_names` and `covers_multiple` agree on it."""
|
||||||
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
|
|
||||||
|
covering = HaikuRAG(config=config)
|
||||||
|
assert covering.source_names == ("alpha", "beta")
|
||||||
|
assert covering.covers_multiple
|
||||||
|
|
||||||
|
narrowed = HaikuRAG(config=config, sources=["beta"])
|
||||||
|
assert narrowed.source_names == ("beta",)
|
||||||
|
assert not narrowed.covers_multiple
|
||||||
|
|
||||||
|
at_path = HaikuRAG(tmp_path / "other.lancedb")
|
||||||
|
assert at_path.source_names == ("other",)
|
||||||
|
assert not at_path.covers_multiple
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_the_reader_for_a_database_is_the_client_holding_it(self, tmp_path):
|
async def test_the_reader_for_a_database_is_the_client_holding_it(self, tmp_path):
|
||||||
|
|
@ -280,10 +337,10 @@ class TestPlacingADatabase:
|
||||||
await alpha.reader_for("beta")
|
await alpha.reader_for("beta")
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_an_unnamed_database_refuses_any_name(self, temp_db_path):
|
async def test_a_database_at_a_path_answers_to_its_stem_alone(self, temp_db_path):
|
||||||
"""Nothing names it, so no name can be the one it covers."""
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
with pytest.raises(UnknownDatabaseError, match="single unnamed database"):
|
assert await rag.reader_for(temp_db_path.stem) is rag
|
||||||
|
with pytest.raises(UnknownDatabaseError, match=temp_db_path.stem):
|
||||||
await rag.reader_for("anything")
|
await rag.reader_for("anything")
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
|
||||||
|
|
@ -71,7 +71,7 @@ class TestFederatedSearch:
|
||||||
|
|
||||||
class TestSingleDatabaseUnchanged:
|
class TestSingleDatabaseUnchanged:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_source_is_unset_without_configured_databases(self, temp_db_path):
|
async def test_source_is_the_stem_without_configured_databases(self, temp_db_path):
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
doc = DoclingDocument(name="one")
|
doc = DoclingDocument(name="one")
|
||||||
doc.add_text(label=DocItemLabel.TEXT, text="a document about cats")
|
doc.add_text(label=DocItemLabel.TEXT, text="a document about cats")
|
||||||
|
|
@ -89,7 +89,7 @@ class TestSingleDatabaseUnchanged:
|
||||||
results = await rag.search("cats", search_type="fts")
|
results = await rag.search("cats", search_type="fts")
|
||||||
|
|
||||||
assert results
|
assert results
|
||||||
assert all(r.source is None for r in results)
|
assert all(r.source == temp_db_path.stem for r in results)
|
||||||
|
|
||||||
|
|
||||||
class TestOneQueryVector:
|
class TestOneQueryVector:
|
||||||
|
|
|
||||||
|
|
@ -110,8 +110,7 @@ class TestSandboxListDocuments:
|
||||||
assert result.success
|
assert result.success
|
||||||
assert "1" in result.stdout
|
assert "1" in result.stdout
|
||||||
assert "Test Document" in result.stdout
|
assert "Test Document" in result.stdout
|
||||||
# Nothing names this database, so there is no name to report.
|
assert temp_db_path.stem in result.stdout
|
||||||
assert "None" in result.stdout
|
|
||||||
|
|
||||||
|
|
||||||
class TestSandboxSearch:
|
class TestSandboxSearch:
|
||||||
|
|
|
||||||
|
|
@ -159,12 +159,14 @@ class TestTheSandboxConstructors:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_the_public_constructor_resolves_the_path_it_is_given(self, tmp_path):
|
async def test_the_public_constructor_resolves_the_path_it_is_given(self, tmp_path):
|
||||||
|
from haiku.rag.config.models import AppConfig
|
||||||
|
|
||||||
config = _config(tmp_path, ["alpha", "beta"])
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
await _seed(config, "alpha", ["alpha document about cats"])
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
|
||||||
sandbox = Sandbox(
|
sandbox = Sandbox(
|
||||||
db_path=tmp_path / "alpha.lancedb",
|
db_path=tmp_path / "alpha.lancedb",
|
||||||
config=config,
|
config=AppConfig(),
|
||||||
context=AnalysisContext(),
|
context=AnalysisContext(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -202,19 +204,20 @@ class TestTheSandboxCoversWhatTheCapabilityCovers:
|
||||||
"""The sandbox covers the scope the capability resolved, as handed
|
"""The sandbox covers the scope the capability resolved, as handed
|
||||||
over."""
|
over."""
|
||||||
from haiku.rag.capabilities.analysis import AnalysisState, create_capability
|
from haiku.rag.capabilities.analysis import AnalysisState, create_capability
|
||||||
|
from haiku.rag.config.models import AppConfig
|
||||||
|
|
||||||
config = _config(tmp_path, ["alpha", "beta"])
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
await _seed(config, "alpha", ["alpha document about cats"])
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
|
||||||
capability = create_capability(
|
capability = create_capability(
|
||||||
db_path=tmp_path / "alpha.lancedb", config=config, defer_loading=False
|
db_path=tmp_path / "alpha.lancedb", config=AppConfig(), defer_loading=False
|
||||||
)
|
)
|
||||||
capability.state = AnalysisState()
|
capability.state = AnalysisState()
|
||||||
|
|
||||||
sandbox = await capability._ensure_sandbox()
|
sandbox = await capability._ensure_sandbox()
|
||||||
try:
|
try:
|
||||||
assert sandbox._scope is capability.scope
|
assert sandbox._scope is capability.scope
|
||||||
assert capability.scope.names == ()
|
assert capability.scope.names == ("alpha",)
|
||||||
finally:
|
finally:
|
||||||
await capability._close()
|
await capability._close()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -69,7 +69,7 @@ async def _seed(temp_db_path, *, version: str, with_items: bool = True):
|
||||||
async def test_gather_database_info_reports_tables_and_settings(temp_db_path):
|
async def test_gather_database_info_reports_tables_and_settings(temp_db_path):
|
||||||
await _seed(temp_db_path, version="1.2.3")
|
await _seed(temp_db_path, version="1.2.3")
|
||||||
|
|
||||||
info = await gather_database_info(AppConfig(), temp_db_path)
|
info = await gather_database_info(temp_db_path, AppConfig())
|
||||||
|
|
||||||
assert info.exists is True
|
assert info.exists is True
|
||||||
assert info.path == str(temp_db_path)
|
assert info.path == str(temp_db_path)
|
||||||
|
|
@ -98,7 +98,7 @@ async def test_gather_database_info_flags_missing_table_and_pending_migrations(
|
||||||
):
|
):
|
||||||
await _seed(temp_db_path, version="0.39.0", with_items=False)
|
await _seed(temp_db_path, version="0.39.0", with_items=False)
|
||||||
|
|
||||||
info = await gather_database_info(AppConfig(), temp_db_path)
|
info = await gather_database_info(temp_db_path, AppConfig())
|
||||||
|
|
||||||
tables = {t.name: t for t in info.tables}
|
tables = {t.name: t for t in info.tables}
|
||||||
assert tables["document_items"].exists is False
|
assert tables["document_items"].exists is False
|
||||||
|
|
@ -112,7 +112,32 @@ async def test_gather_database_info_empty_database(temp_db_path):
|
||||||
|
|
||||||
await lancedb.connect_async(temp_db_path) # creates the dir, no tables
|
await lancedb.connect_async(temp_db_path) # creates the dir, no tables
|
||||||
|
|
||||||
info = await gather_database_info(AppConfig(), temp_db_path)
|
info = await gather_database_info(temp_db_path, AppConfig())
|
||||||
|
|
||||||
assert info.exists is False
|
assert info.exists is False
|
||||||
assert info.path == str(temp_db_path)
|
assert info.path == str(temp_db_path)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gather_database_info_connects_to_the_location_it_is_given():
|
||||||
|
"""A remote location is passed to the connection as is and reported back
|
||||||
|
as the path; the configuration's own `uri` plays no part."""
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
from haiku.rag.config.models import LanceDBConfig
|
||||||
|
|
||||||
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(databases={"other": "s3://elsewhere/other.lancedb"})
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"haiku.rag.store.info.connect_lancedb", new_callable=AsyncMock
|
||||||
|
) as mock_connect:
|
||||||
|
listing = MagicMock()
|
||||||
|
listing.tables = []
|
||||||
|
mock_connect.return_value.list_tables = AsyncMock(return_value=listing)
|
||||||
|
|
||||||
|
info = await gather_database_info("s3://bucket/papers.lancedb", config)
|
||||||
|
|
||||||
|
assert mock_connect.call_args.args[0] == "s3://bucket/papers.lancedb"
|
||||||
|
assert info.path == "s3://bucket/papers.lancedb"
|
||||||
|
assert info.exists is False
|
||||||
|
|
|
||||||
|
|
@ -126,6 +126,43 @@ class TestV0_58_0MigrationEdgeCases:
|
||||||
assert set(by_id) == {"a", "b"} # exactly one row each, no duplicates
|
assert set(by_id) == {"a", "b"} # exactly one row each, no duplicates
|
||||||
assert len(rows) == 2
|
assert len(rows) == 2
|
||||||
|
|
||||||
|
async def test_a_remote_store_has_no_disk_to_check(self, temp_db_path, monkeypatch):
|
||||||
|
"""A store behind a URI has no local path: the reclaim vacuum runs
|
||||||
|
without a free-disk check."""
|
||||||
|
from haiku.rag.store.upgrades import v0_58_0
|
||||||
|
|
||||||
|
async with Store(temp_db_path, create=True, skip_migration_check=True) as store:
|
||||||
|
await seed_legacy_documents(
|
||||||
|
store,
|
||||||
|
[LegacyDocumentRecord(id="a", content="x", uri="u", metadata="{}")],
|
||||||
|
)
|
||||||
|
await store.set_haiku_version("0.57.0")
|
||||||
|
|
||||||
|
def no_disk(_path):
|
||||||
|
raise AssertionError("disk_usage consulted for a remote store")
|
||||||
|
|
||||||
|
monkeypatch.setattr(v0_58_0.shutil, "disk_usage", no_disk)
|
||||||
|
vacuum_calls: list[int] = []
|
||||||
|
|
||||||
|
async with Store(temp_db_path, skip_migration_check=True) as store:
|
||||||
|
store.db_path = None
|
||||||
|
|
||||||
|
async def fake_stats():
|
||||||
|
return {"total_bytes": 10_000_000}
|
||||||
|
|
||||||
|
monkeypatch.setattr(store.documents_table, "stats", fake_stats)
|
||||||
|
orig_vacuum = store.vacuum
|
||||||
|
|
||||||
|
async def tracking_vacuum(*args, **kwargs):
|
||||||
|
vacuum_calls.append(1)
|
||||||
|
return await orig_vacuum(*args, **kwargs)
|
||||||
|
|
||||||
|
monkeypatch.setattr(store, "vacuum", tracking_vacuum)
|
||||||
|
|
||||||
|
await store.migrate()
|
||||||
|
|
||||||
|
assert vacuum_calls == [1]
|
||||||
|
|
||||||
async def test_skips_vacuum_when_disk_is_tight(self, temp_db_path, monkeypatch):
|
async def test_skips_vacuum_when_disk_is_tight(self, temp_db_path, monkeypatch):
|
||||||
"""When free disk can't cover one compacted copy, the split still
|
"""When free disk can't cover one compacted copy, the split still
|
||||||
completes but the reclaim vacuum is skipped."""
|
completes but the reclaim vacuum is skipped."""
|
||||||
|
|
|
||||||
|
|
@ -456,8 +456,10 @@ async def test_create_index_rebuilds_an_existing_one(app, client):
|
||||||
|
|
||||||
|
|
||||||
def test_show_settings_hides_secrets(tmp_path):
|
def test_show_settings_hides_secrets(tmp_path):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="db://x", api_key="secret-value"))
|
config = AppConfig(
|
||||||
app = HaikuRAGApp(scope=for_path(tmp_path / "db", config), config=config)
|
lancedb=LanceDBConfig(databases={"x": "db://x"}, api_key="secret-value")
|
||||||
|
)
|
||||||
|
app = HaikuRAGApp(scope=DatabaseScope.resolve(config), config=config)
|
||||||
app.console = Console(record=True, width=200)
|
app.console = Console(record=True, width=200)
|
||||||
|
|
||||||
app.show_settings()
|
app.show_settings()
|
||||||
|
|
@ -476,7 +478,7 @@ def test_show_settings_renders_the_shape_a_config_file_has(tmp_path):
|
||||||
lancedb=LanceDBConfig(databases={"alpha": "/tmp/a.lancedb"}),
|
lancedb=LanceDBConfig(databases={"alpha": "/tmp/a.lancedb"}),
|
||||||
storage=StorageConfig(data_dir=tmp_path),
|
storage=StorageConfig(data_dir=tmp_path),
|
||||||
)
|
)
|
||||||
app = HaikuRAGApp(scope=for_path(tmp_path / "db", config), config=config)
|
app = HaikuRAGApp(scope=DatabaseScope.resolve(config), config=config)
|
||||||
app.console = Console(record=True, width=200)
|
app.console = Console(record=True, width=200)
|
||||||
|
|
||||||
app.show_settings()
|
app.show_settings()
|
||||||
|
|
@ -515,7 +517,7 @@ def test_show_settings_survives_a_narrow_console_and_bracketed_values(tmp_path):
|
||||||
|
|
||||||
|
|
||||||
def test_remote_uri_is_the_display_path(tmp_path):
|
def test_remote_uri_is_the_display_path(tmp_path):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
|
config = AppConfig(lancedb=LanceDBConfig(databases={"path": "s3://bucket/path"}))
|
||||||
app = HaikuRAGApp(scope=for_path(None, config), config=config)
|
app = HaikuRAGApp(scope=for_path(None, config), config=config)
|
||||||
|
|
||||||
assert app.display_path == "s3://bucket/path"
|
assert app.display_path == "s3://bucket/path"
|
||||||
|
|
|
||||||
|
|
@ -158,6 +158,27 @@ class TestOneDatabaseCommands:
|
||||||
[ref] = resolve_scope(Path("/db/other.lancedb")).databases
|
[ref] = resolve_scope(Path("/db/other.lancedb")).databases
|
||||||
assert ref.db_path == Path("/db/other.lancedb")
|
assert ref.db_path == Path("/db/other.lancedb")
|
||||||
|
|
||||||
|
def test_naming_a_path_overrides_a_configured_set(self, monkeypatch):
|
||||||
|
"""`--db` is the operator's explicit override: it constructs the scope
|
||||||
|
directly, where a Python caller passing a path beside `databases` is
|
||||||
|
refused."""
|
||||||
|
self._install(monkeypatch, alpha="/db/a.lancedb", beta="/db/b.lancedb")
|
||||||
|
|
||||||
|
scope = resolve_scope(Path("/db/other.lancedb"))
|
||||||
|
|
||||||
|
assert scope.names == ("other",)
|
||||||
|
assert not scope.covers_multiple
|
||||||
|
|
||||||
|
def test_a_path_without_a_stem_is_a_usage_error(self, monkeypatch):
|
||||||
|
"""A path that names no database is the operator's mistake, reported as
|
||||||
|
one."""
|
||||||
|
import typer
|
||||||
|
|
||||||
|
self._install(monkeypatch)
|
||||||
|
|
||||||
|
with pytest.raises(typer.BadParameter, match="no name"):
|
||||||
|
resolve_scope(Path("/"))
|
||||||
|
|
||||||
def test_no_configured_databases_is_allowed(self, monkeypatch, tmp_path):
|
def test_no_configured_databases_is_allowed(self, monkeypatch, tmp_path):
|
||||||
import haiku.rag.config as config_module
|
import haiku.rag.config as config_module
|
||||||
|
|
||||||
|
|
@ -207,7 +228,6 @@ class TestSelectingADatabaseByName:
|
||||||
"notes": "/data/notes.lancedb",
|
"notes": "/data/notes.lancedb",
|
||||||
"other": "/data/o.lancedb",
|
"other": "/data/o.lancedb",
|
||||||
}
|
}
|
||||||
assert config.lancedb.uri == ""
|
|
||||||
|
|
||||||
def test_an_unknown_name_names_the_configured_ones(self, monkeypatch):
|
def test_an_unknown_name_names_the_configured_ones(self, monkeypatch):
|
||||||
self._install(monkeypatch, alpha="/data/a.lancedb", beta="/data/b.lancedb")
|
self._install(monkeypatch, alpha="/data/a.lancedb", beta="/data/b.lancedb")
|
||||||
|
|
@ -225,11 +245,15 @@ class TestSelectingADatabaseByName:
|
||||||
|
|
||||||
assert "bucket" not in str(raised.value)
|
assert "bucket" not in str(raised.value)
|
||||||
|
|
||||||
def test_selecting_nothing_reports_an_empty_mapping(self, monkeypatch):
|
def test_an_unknown_name_with_nothing_configured_names_the_default(
|
||||||
|
self, monkeypatch
|
||||||
|
):
|
||||||
|
"""Nothing configured is the one entry `haiku.rag`, which the message
|
||||||
|
offers."""
|
||||||
self._install(monkeypatch)
|
self._install(monkeypatch)
|
||||||
monkeypatch.setattr("haiku.rag.cli._db_name", "papers")
|
monkeypatch.setattr("haiku.rag.cli._db_name", "papers")
|
||||||
|
|
||||||
with pytest.raises(UnknownDatabaseError, match="nothing"):
|
with pytest.raises(UnknownDatabaseError, match="haiku.rag"):
|
||||||
resolve_scope(None)
|
resolve_scope(None)
|
||||||
|
|
||||||
def test_the_callback_selects_before_a_command_runs(self, tmp_path, monkeypatch):
|
def test_the_callback_selects_before_a_command_runs(self, tmp_path, monkeypatch):
|
||||||
|
|
@ -288,7 +312,6 @@ class TestSelectingADatabaseByName:
|
||||||
cli, ["--config", str(config_file), "--db-name", "alpha", "settings"]
|
cli, ["--config", str(config_file), "--db-name", "alpha", "settings"]
|
||||||
)
|
)
|
||||||
# Naming one leaves the configuration naming both.
|
# Naming one leaves the configuration naming both.
|
||||||
assert get_config().lancedb.uri == ""
|
|
||||||
assert set(get_config().lancedb.databases) == {"alpha", "beta"}
|
assert set(get_config().lancedb.databases) == {"alpha", "beta"}
|
||||||
|
|
||||||
runner.invoke(cli, ["--config", str(config_file), "settings"])
|
runner.invoke(cli, ["--config", str(config_file), "settings"])
|
||||||
|
|
@ -384,12 +407,12 @@ class TestResolvingTheDatabasePath:
|
||||||
resolve_scope(Path("/data/other.lancedb"))
|
resolve_scope(Path("/data/other.lancedb"))
|
||||||
|
|
||||||
|
|
||||||
class TestConfiguredLocalUri:
|
class TestConfiguredLocalDatabase:
|
||||||
"""`lancedb.uri` with a local path."""
|
"""One entry in `lancedb.databases` with a local path."""
|
||||||
|
|
||||||
def _config_file(self, tmp_path, located: Path) -> Path:
|
def _config_file(self, tmp_path, located: Path) -> Path:
|
||||||
config_file = tmp_path / "haiku.rag.yaml"
|
config_file = tmp_path / "haiku.rag.yaml"
|
||||||
config_file.write_text(f"lancedb:\n uri: {located}\n")
|
config_file.write_text(f"lancedb:\n databases:\n notes: {located}\n")
|
||||||
return config_file
|
return config_file
|
||||||
|
|
||||||
def _fresh(self, monkeypatch) -> None:
|
def _fresh(self, monkeypatch) -> None:
|
||||||
|
|
@ -429,7 +452,7 @@ class TestConfiguredLocalUri:
|
||||||
assert "does not exist" in result.output
|
assert "does not exist" in result.output
|
||||||
assert not located.exists()
|
assert not located.exists()
|
||||||
|
|
||||||
def test_db_overrides_the_configured_uri(self, tmp_path, monkeypatch):
|
def test_db_overrides_the_configured_database(self, tmp_path, monkeypatch):
|
||||||
self._fresh(monkeypatch)
|
self._fresh(monkeypatch)
|
||||||
configured = tmp_path / "configured.lancedb"
|
configured = tmp_path / "configured.lancedb"
|
||||||
chosen = tmp_path / "chosen.lancedb"
|
chosen = tmp_path / "chosen.lancedb"
|
||||||
|
|
|
||||||
|
|
@ -800,3 +800,31 @@ def test_complete_example_matches_the_defaults():
|
||||||
}
|
}
|
||||||
|
|
||||||
assert not drifted, f"documented value != default: {drifted}"
|
assert not drifted, f"documented value != default: {drifted}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_lancedb_uri_is_refused_with_the_replacement_named():
|
||||||
|
"""`lancedb.uri` no longer places a database; a config carrying it fails to
|
||||||
|
load with the `databases` spelling to use instead."""
|
||||||
|
from haiku.rag.config.models import LanceDBConfig
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError) as raised:
|
||||||
|
AppConfig.model_validate({"lancedb": {"uri": "s3://bucket/notes.lancedb"}})
|
||||||
|
|
||||||
|
message = str(raised.value)
|
||||||
|
assert "lancedb.uri" in message
|
||||||
|
assert "lancedb.databases: {NAME: 's3://bucket/notes.lancedb'}" in message
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError, match="lancedb.databases"):
|
||||||
|
LanceDBConfig.model_validate({"uri": "/data/notes.lancedb"})
|
||||||
|
|
||||||
|
|
||||||
|
def test_an_empty_lancedb_uri_is_refused_with_removal_as_the_remedy():
|
||||||
|
"""Generated configs carried `uri: ""` for the local default. The remedy is
|
||||||
|
to delete the key, never a mapping with an empty location."""
|
||||||
|
with pytest.raises(ValidationError) as raised:
|
||||||
|
AppConfig.model_validate({"lancedb": {"uri": ""}})
|
||||||
|
|
||||||
|
message = str(raised.value)
|
||||||
|
assert "lancedb.uri" in message
|
||||||
|
assert "remove" in message
|
||||||
|
assert "{NAME" not in message
|
||||||
|
|
|
||||||
|
|
@ -55,7 +55,8 @@ def test_default_db_path_comes_from_storage_data_dir(tmp_path):
|
||||||
|
|
||||||
[ref] = HaikuRAG(config=config)._resolve_scope().databases
|
[ref] = HaikuRAG(config=config)._resolve_scope().databases
|
||||||
|
|
||||||
assert ref.db_path == tmp_path / "haiku.rag.lancedb"
|
assert ref.location == tmp_path / "haiku.rag.lancedb"
|
||||||
|
assert ref.name == "haiku.rag"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,7 @@ from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.client.scope import DatabaseRef, DatabaseScope
|
from haiku.rag.client.scope import DatabaseRef, DatabaseScope, database_name
|
||||||
from haiku.rag.client.session import default_db_path
|
|
||||||
from haiku.rag.config.models import AppConfig, LanceDBConfig, StorageConfig
|
from haiku.rag.config.models import AppConfig, LanceDBConfig, StorageConfig
|
||||||
from haiku.rag.store.exceptions import (
|
from haiku.rag.store.exceptions import (
|
||||||
AmbiguousDatabaseError,
|
AmbiguousDatabaseError,
|
||||||
|
|
@ -26,31 +25,50 @@ class TestResolution:
|
||||||
config, database_name="alpha", database_path=Path("/data/other.lancedb")
|
config, database_name="alpha", database_path=Path("/data/other.lancedb")
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_a_path_names_one_unnamed_database(self):
|
def test_a_path_places_the_database_where_the_configuration_places_none(self):
|
||||||
"""A path says which database, not what it is called, even where the
|
scope = DatabaseScope.resolve(_config(), database_path="/data/other.lancedb")
|
||||||
configuration names one."""
|
|
||||||
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
|
||||||
|
|
||||||
scope = DatabaseScope.resolve(config, database_path=Path("/data/other.lancedb"))
|
|
||||||
|
|
||||||
assert scope.databases == (DatabaseRef.at("/data/other.lancedb"),)
|
assert scope.databases == (DatabaseRef.at("/data/other.lancedb"),)
|
||||||
assert scope.names == ()
|
assert scope.names == ("other",)
|
||||||
assert not scope.covers_multiple
|
assert not scope.covers_multiple
|
||||||
|
|
||||||
|
def test_a_path_beside_a_configured_placement_is_refused(self):
|
||||||
|
"""The configuration places databases; a path beside it is a second
|
||||||
|
placement, and the refusal names both."""
|
||||||
|
config = _config(databases={"alpha": "/data/alpha.lancedb", "beta": "b://b"})
|
||||||
|
|
||||||
|
with pytest.raises(AmbiguousDatabaseError) as raised:
|
||||||
|
DatabaseScope.resolve(config, database_path=Path("/data/other.lancedb"))
|
||||||
|
|
||||||
|
message = str(raised.value)
|
||||||
|
assert "/data/other.lancedb" in message
|
||||||
|
assert "alpha" in message and "beta" in message
|
||||||
|
assert "lancedb.databases" in message
|
||||||
|
|
||||||
def test_a_named_database_keeps_its_name(self):
|
def test_a_named_database_keeps_its_name(self):
|
||||||
config = _config(databases={"alpha": "/data/alpha.lancedb", "beta": "b://b"})
|
config = _config(databases={"alpha": "/data/alpha.lancedb", "beta": "b://b"})
|
||||||
|
|
||||||
scope = DatabaseScope.resolve(config, database_name="beta")
|
scope = DatabaseScope.resolve(config, database_name="beta")
|
||||||
|
|
||||||
assert scope.databases == (DatabaseRef("beta", "b://b", None),)
|
assert scope.databases == (DatabaseRef("beta", "b://b"),)
|
||||||
assert scope.names == ("beta",)
|
assert scope.names == ("beta",)
|
||||||
|
|
||||||
def test_an_unknown_name_is_refused(self):
|
def test_an_unknown_name_is_refused(self):
|
||||||
|
"""The message lists the databases there are, configured or default."""
|
||||||
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
||||||
|
|
||||||
with pytest.raises(UnknownDatabaseError, match="unknown database 'nope'"):
|
with pytest.raises(
|
||||||
|
UnknownDatabaseError,
|
||||||
|
match="unknown database 'nope'.*the databases are alpha",
|
||||||
|
):
|
||||||
DatabaseScope.resolve(config, database_name="nope")
|
DatabaseScope.resolve(config, database_name="nope")
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
UnknownDatabaseError,
|
||||||
|
match="unknown database 'nope'.*the databases are haiku.rag",
|
||||||
|
):
|
||||||
|
DatabaseScope.resolve(_config(), database_name="nope")
|
||||||
|
|
||||||
def test_no_selector_covers_the_configured_set_in_order(self):
|
def test_no_selector_covers_the_configured_set_in_order(self):
|
||||||
config = _config(
|
config = _config(
|
||||||
databases={"beta": "/data/b.lancedb", "alpha": "/data/a.lancedb"}
|
databases={"beta": "/data/b.lancedb", "alpha": "/data/a.lancedb"}
|
||||||
|
|
@ -67,49 +85,41 @@ class TestResolution:
|
||||||
|
|
||||||
scope = DatabaseScope.resolve(config)
|
scope = DatabaseScope.resolve(config)
|
||||||
|
|
||||||
assert scope.databases == (
|
assert scope.databases == (DatabaseRef("alpha", Path("/data/alpha.lancedb")),)
|
||||||
DatabaseRef.configured("alpha", "/data/alpha.lancedb"),
|
|
||||||
)
|
|
||||||
assert not scope.covers_multiple
|
assert not scope.covers_multiple
|
||||||
|
|
||||||
def test_a_bare_uri_is_one_unnamed_database(self):
|
def test_nothing_configured_is_the_default_entry(self, tmp_path):
|
||||||
scope = DatabaseScope.resolve(_config(uri="s3://bucket/one.lancedb"))
|
"""No `databases` reads as one entry, `haiku.rag`, under the data
|
||||||
|
directory: an ordinary configured database in every respect."""
|
||||||
assert scope.databases == (DatabaseRef(None, "s3://bucket/one.lancedb", None),)
|
|
||||||
|
|
||||||
def test_a_bare_uri_without_a_scheme_is_a_local_path(self):
|
|
||||||
"""`lancedb.uri` places one database the same way an entry in
|
|
||||||
`lancedb.databases` does, so a schemeless value is a path and gets the
|
|
||||||
existence check a local database gets."""
|
|
||||||
scope = DatabaseScope.resolve(_config(uri="/data/notes.lancedb"))
|
|
||||||
|
|
||||||
[ref] = scope.databases
|
|
||||||
assert ref.name is None
|
|
||||||
assert ref.db_path == Path("/data/notes.lancedb")
|
|
||||||
assert ref.uri == ""
|
|
||||||
|
|
||||||
def test_a_path_selects_the_database_over_a_configured_uri(self):
|
|
||||||
"""`--db` exists to override what is configured, and the configuration
|
|
||||||
derived from the ref is what makes the connection follow it."""
|
|
||||||
config = _config(uri="s3://bucket/one.lancedb")
|
|
||||||
|
|
||||||
scope = DatabaseScope.resolve(config, database_path=Path("/data/local"))
|
|
||||||
|
|
||||||
[ref] = scope.databases
|
|
||||||
assert ref.db_path == Path("/data/local")
|
|
||||||
one, _ = ref.connection(config)
|
|
||||||
assert one.lancedb.uri == ""
|
|
||||||
|
|
||||||
def test_nothing_configured_falls_back_to_the_data_directory(self, tmp_path):
|
|
||||||
config = AppConfig(storage=StorageConfig(data_dir=tmp_path))
|
config = AppConfig(storage=StorageConfig(data_dir=tmp_path))
|
||||||
|
|
||||||
scope = DatabaseScope.resolve(config)
|
scope = DatabaseScope.resolve(config)
|
||||||
|
|
||||||
assert scope.databases == (DatabaseRef.at(tmp_path / "haiku.rag.lancedb"),)
|
assert scope.databases == (
|
||||||
|
DatabaseRef("haiku.rag", tmp_path / "haiku.rag.lancedb"),
|
||||||
|
)
|
||||||
|
assert scope.names == ("haiku.rag",)
|
||||||
|
|
||||||
def test_the_environment_is_not_consulted(self, monkeypatch, tmp_path):
|
def test_the_default_entry_is_selectable_by_name(self, tmp_path):
|
||||||
"""HAIKU_RAG_DB is honoured by the capability entry point alone;
|
config = AppConfig(storage=StorageConfig(data_dir=tmp_path))
|
||||||
resolution never reads the environment."""
|
|
||||||
|
by_name = DatabaseScope.resolve(config, database_name="haiku.rag")
|
||||||
|
selected = DatabaseScope.resolve(config).select(["haiku.rag"])
|
||||||
|
|
||||||
|
assert by_name == selected == DatabaseScope.resolve(config)
|
||||||
|
|
||||||
|
def test_the_default_entry_does_not_hide_a_configured_set(self, tmp_path):
|
||||||
|
"""The default stands in only where nothing is configured."""
|
||||||
|
config = AppConfig(
|
||||||
|
storage=StorageConfig(data_dir=tmp_path),
|
||||||
|
lancedb=LanceDBConfig(databases={"alpha": "/data/alpha.lancedb"}),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(UnknownDatabaseError, match="haiku.rag"):
|
||||||
|
DatabaseScope.resolve(config, database_name="haiku.rag")
|
||||||
|
|
||||||
|
def test_the_environment_is_not_consulted(self, monkeypatch):
|
||||||
|
"""Resolution reads the configuration alone."""
|
||||||
monkeypatch.setenv("HAIKU_RAG_DB", "/data/from-the-environment.lancedb")
|
monkeypatch.setenv("HAIKU_RAG_DB", "/data/from-the-environment.lancedb")
|
||||||
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
||||||
|
|
||||||
|
|
@ -124,8 +134,7 @@ class TestResolution:
|
||||||
)
|
)
|
||||||
|
|
||||||
[ref] = scope.databases
|
[ref] = scope.databases
|
||||||
assert ref.db_path == Path("s3://bucket/looks-like-a-uri.lancedb")
|
assert ref.location == Path("s3://bucket/looks-like-a-uri.lancedb")
|
||||||
assert ref.uri == ""
|
|
||||||
|
|
||||||
def test_a_configured_location_with_a_scheme_is_a_uri(self):
|
def test_a_configured_location_with_a_scheme_is_a_uri(self):
|
||||||
"""A configured value is a URI or a path depending on its scheme, which is
|
"""A configured value is a URI or a path depending on its scheme, which is
|
||||||
|
|
@ -134,22 +143,16 @@ class TestResolution:
|
||||||
|
|
||||||
[ref] = DatabaseScope.resolve(config).databases
|
[ref] = DatabaseScope.resolve(config).databases
|
||||||
|
|
||||||
assert ref.uri == "s3://bucket/alpha.lancedb"
|
assert ref.location == "s3://bucket/alpha.lancedb"
|
||||||
assert ref.db_path is None
|
assert ref.db_path is None
|
||||||
|
|
||||||
def test_a_database_is_a_uri_or_a_path(self):
|
def test_a_configured_location_without_a_scheme_is_a_path(self):
|
||||||
"""A ref holding both, or neither, is refused at construction.
|
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
||||||
|
|
||||||
The message names what it was given: this is a programming error raised
|
[ref] = DatabaseScope.resolve(config).databases
|
||||||
in the caller's own process, not one an operator or a model ever sees.
|
|
||||||
"""
|
|
||||||
with pytest.raises(ValueError, match="either a URI or a local path") as both:
|
|
||||||
DatabaseRef(None, "s3://bucket/a.lancedb", Path("/data/a.lancedb"))
|
|
||||||
assert "s3://bucket/a.lancedb" in str(both.value)
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="either a URI or a local path") as neither:
|
assert ref.location == Path("/data/alpha.lancedb")
|
||||||
DatabaseRef(None, "", None)
|
assert ref.db_path == Path("/data/alpha.lancedb")
|
||||||
assert "db_path=None" in str(neither.value)
|
|
||||||
|
|
||||||
def test_a_scope_covers_at_least_one_database(self):
|
def test_a_scope_covers_at_least_one_database(self):
|
||||||
"""Every resolution reaches a database, and the sessions built from a
|
"""Every resolution reaches a database, and the sessions built from a
|
||||||
|
|
@ -158,61 +161,69 @@ class TestResolution:
|
||||||
DatabaseScope(())
|
DatabaseScope(())
|
||||||
|
|
||||||
|
|
||||||
class TestConnectionDerivation:
|
class TestTheReference:
|
||||||
"""Opening one of a set must not disturb the configuration it came from."""
|
"""Constructed directly, a reference still holds what it advertises."""
|
||||||
|
|
||||||
def test_a_local_location_becomes_a_path(self):
|
def test_a_schemeless_string_location_is_a_path(self):
|
||||||
config = _config(databases={"alpha": "/data/alpha.lancedb"})
|
ref = DatabaseRef("x", "local.lancedb")
|
||||||
[ref] = DatabaseScope.resolve(config).databases
|
|
||||||
|
|
||||||
one, db_path = ref.connection(config)
|
assert ref.location == Path("local.lancedb")
|
||||||
|
assert ref.db_path == Path("local.lancedb")
|
||||||
|
|
||||||
assert db_path == Path("/data/alpha.lancedb")
|
def test_a_location_with_a_scheme_stays_a_uri(self):
|
||||||
assert one.lancedb.uri == ""
|
ref = DatabaseRef("x", "s3://bucket/x.lancedb")
|
||||||
assert one.lancedb.databases == {}
|
|
||||||
|
|
||||||
def test_a_uri_location_stays_a_uri(self):
|
assert ref.location == "s3://bucket/x.lancedb"
|
||||||
config = _config(databases={"alpha": "s3://bucket/alpha.lancedb"})
|
assert ref.db_path is None
|
||||||
[ref] = DatabaseScope.resolve(config).databases
|
|
||||||
|
|
||||||
one, db_path = ref.connection(config)
|
def test_a_blank_name_is_refused(self):
|
||||||
|
with pytest.raises(ValueError, match="no name"):
|
||||||
|
DatabaseRef("", "/data/x.lancedb")
|
||||||
|
with pytest.raises(ValueError, match="no name"):
|
||||||
|
DatabaseRef(" ", "s3://bucket/x.lancedb")
|
||||||
|
|
||||||
assert db_path is None
|
def test_a_blank_location_is_refused(self):
|
||||||
assert one.lancedb.uri == "s3://bucket/alpha.lancedb"
|
"""A blank string would resolve to the working directory."""
|
||||||
|
with pytest.raises(ValueError, match="no location"):
|
||||||
|
DatabaseRef("x", "")
|
||||||
|
with pytest.raises(ValueError, match="no location"):
|
||||||
|
DatabaseRef("x", " ")
|
||||||
|
|
||||||
def test_the_original_configuration_is_untouched(self):
|
def test_a_given_database_is_a_local_path(self):
|
||||||
"""Rewriting it in place is what left downstream code unable to tell a set
|
"""Only a path can be given: a given database's errors name its
|
||||||
had been named."""
|
location, and a URI must never travel that way."""
|
||||||
config = _config(databases={"alpha": "/a.lancedb", "beta": "/b.lancedb"})
|
assert DatabaseRef("x", "local.lancedb", given=True).location == Path(
|
||||||
|
"local.lancedb"
|
||||||
for ref in DatabaseScope.resolve(config).databases:
|
)
|
||||||
ref.connection(config)
|
with pytest.raises(ValueError, match="is a URI"):
|
||||||
|
DatabaseRef("x", "s3://bucket/x.lancedb", given=True)
|
||||||
assert config.lancedb.databases == {"alpha": "/a.lancedb", "beta": "/b.lancedb"}
|
|
||||||
assert config.lancedb.uri == ""
|
|
||||||
|
|
||||||
def test_each_derived_configuration_is_its_own_copy(self):
|
|
||||||
config = _config(databases={"alpha": "/a.lancedb", "beta": "s3://b/b.lancedb"})
|
|
||||||
alpha, beta = DatabaseScope.resolve(config).databases
|
|
||||||
|
|
||||||
one, _ = alpha.connection(config)
|
|
||||||
other, _ = beta.connection(config)
|
|
||||||
|
|
||||||
assert one is not other
|
|
||||||
assert one.lancedb.uri == ""
|
|
||||||
assert other.lancedb.uri == "s3://b/b.lancedb"
|
|
||||||
|
|
||||||
|
|
||||||
def test_a_database_behind_a_uri_has_no_path_of_its_own(tmp_path):
|
class TestNamingAPath:
|
||||||
"""`connection` hands back no path for a URI, and the store still needs one:
|
"""A path the caller gave is named by its stem, the one rule for the
|
||||||
the default stands in, and the URI is what decides where it connects."""
|
default database and for `--db`."""
|
||||||
config = AppConfig(
|
|
||||||
storage=StorageConfig(data_dir=tmp_path),
|
|
||||||
lancedb=LanceDBConfig(databases={"alpha": "s3://bucket/alpha.lancedb"}),
|
|
||||||
)
|
|
||||||
[ref] = DatabaseScope.resolve(config).databases
|
|
||||||
|
|
||||||
one, db_path = ref.connection(config)
|
def test_the_stem_names_the_database(self):
|
||||||
|
assert database_name(Path("/data/foo.lancedb")) == "foo"
|
||||||
|
assert database_name(Path("relative.lancedb")) == "relative"
|
||||||
|
assert database_name(Path("/data/haiku.rag.lancedb")) == "haiku.rag"
|
||||||
|
|
||||||
assert db_path is None
|
def test_a_path_with_no_stem_is_refused(self):
|
||||||
assert default_db_path(one) == tmp_path / "haiku.rag.lancedb"
|
with pytest.raises(ValueError, match="no name"):
|
||||||
|
database_name(Path("/"))
|
||||||
|
|
||||||
|
def test_at_names_the_path_it_is_given(self):
|
||||||
|
"""A given path is marked as such: its errors may name it, since the
|
||||||
|
caller already knows where it is."""
|
||||||
|
assert DatabaseRef.at("/data/other.lancedb") == DatabaseRef(
|
||||||
|
"other", Path("/data/other.lancedb"), given=True
|
||||||
|
)
|
||||||
|
assert not DatabaseRef.configured("other", "/data/other.lancedb").given
|
||||||
|
|
||||||
|
def test_a_scope_at_a_path_ignores_the_configuration(self):
|
||||||
|
"""The CLI's `--db` is a human's explicit override: it constructs the
|
||||||
|
scope directly and consults no configuration."""
|
||||||
|
scope = DatabaseScope.at(Path("/data/other.lancedb"))
|
||||||
|
|
||||||
|
assert scope.databases == (DatabaseRef.at("/data/other.lancedb"),)
|
||||||
|
assert scope.names == ("other",)
|
||||||
|
|
|
||||||
|
|
@ -177,8 +177,7 @@ async def test_app_info_opens_a_named_remote_database(tmp_path):
|
||||||
app = HaikuRAGApp(scope=scope, config=config)
|
app = HaikuRAGApp(scope=scope, config=config)
|
||||||
|
|
||||||
assert app._is_local is False
|
assert app._is_local is False
|
||||||
assert app._store_config.lancedb.uri == "s3://bucket/papers.lancedb"
|
assert app._location == "s3://bucket/papers.lancedb"
|
||||||
assert app._store_config.lancedb.databases == {}
|
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.info.connect_lancedb", new_callable=AsyncMock
|
"haiku.rag.store.info.connect_lancedb", new_callable=AsyncMock
|
||||||
|
|
@ -189,13 +188,12 @@ async def test_app_info_opens_a_named_remote_database(tmp_path):
|
||||||
mock_db.list_tables = AsyncMock(return_value=mock_list_result)
|
mock_db.list_tables = AsyncMock(return_value=mock_list_result)
|
||||||
await app.info()
|
await app.info()
|
||||||
|
|
||||||
opened = mock_connect.call_args.args[0]
|
assert mock_connect.call_args.args[0] == "s3://bucket/papers.lancedb"
|
||||||
assert opened.lancedb.uri == "s3://bucket/papers.lancedb"
|
|
||||||
|
|
||||||
|
|
||||||
async def test_app_doctor_opens_a_named_remote_database():
|
async def test_app_doctor_opens_a_named_remote_database():
|
||||||
"""`run_doctor` connects with the configuration it is handed: the one
|
"""`run_doctor` is handed the database's location, not the configuration
|
||||||
derived for the database, not the one naming the set."""
|
naming the set."""
|
||||||
from haiku.rag.client.scope import DatabaseScope
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
|
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
|
|
@ -207,14 +205,14 @@ async def test_app_doctor_opens_a_named_remote_database():
|
||||||
run.return_value = MagicMock(checks=[], ok=True, duplicates=None)
|
run.return_value = MagicMock(checks=[], ok=True, duplicates=None)
|
||||||
await app.doctor()
|
await app.doctor()
|
||||||
|
|
||||||
assert run.call_args.args[0].lancedb.uri == "s3://bucket/papers.lancedb"
|
assert run.call_args.args[1] == "s3://bucket/papers.lancedb"
|
||||||
|
|
||||||
|
|
||||||
async def test_app_info_uses_connect_lancedb_for_remote(tmp_path):
|
async def test_app_info_uses_connect_lancedb_for_remote(tmp_path):
|
||||||
"""info() should use connect_lancedb() instead of direct lancedb.connect() for remote URIs."""
|
"""info() should use connect_lancedb() instead of direct lancedb.connect() for remote URIs."""
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
databases={"path": "s3://bucket/path"},
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
storage_options={"endpoint": "http://localhost:9000"},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -230,9 +228,8 @@ async def test_app_info_uses_connect_lancedb_for_remote(tmp_path):
|
||||||
mock_db.list_tables = AsyncMock(return_value=mock_list_result)
|
mock_db.list_tables = AsyncMock(return_value=mock_list_result)
|
||||||
await app.info()
|
await app.info()
|
||||||
|
|
||||||
# The uri decides where it connects; the path argument is not read.
|
|
||||||
mock_connect.assert_called_once()
|
mock_connect.assert_called_once()
|
||||||
assert mock_connect.call_args.args[0].lancedb.uri == "s3://bucket/path"
|
assert mock_connect.call_args.args[0] == "s3://bucket/path"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -371,7 +368,7 @@ async def test_app_init_skips_exists_check_for_remote(tmp_path):
|
||||||
"""init() should not check db_path.exists() for remote URIs."""
|
"""init() should not check db_path.exists() for remote URIs."""
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
databases={"path": "s3://bucket/path"},
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
storage_options={"endpoint": "http://localhost:9000"},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -392,7 +389,7 @@ async def test_app_history_skips_exists_check_for_remote(tmp_path):
|
||||||
"""history() should not check db_path.exists() for remote URIs."""
|
"""history() should not check db_path.exists() for remote URIs."""
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
databases={"path": "s3://bucket/path"},
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
storage_options={"endpoint": "http://localhost:9000"},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -419,7 +416,7 @@ async def test_app_tag_rendering_escapes_markup(tmp_path):
|
||||||
|
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
databases={"path": "s3://bucket/path"},
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
storage_options={"endpoint": "http://localhost:9000"},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -453,7 +450,7 @@ async def test_app_history_survives_tag_annotation_failure(tmp_path):
|
||||||
|
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
databases={"path": "s3://bucket/path"},
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
storage_options={"endpoint": "http://localhost:9000"},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -340,23 +340,19 @@ class TestReportedLocation:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _session(location: str):
|
def _session(location: str):
|
||||||
from haiku.rag.client.scope import DatabaseScope
|
from haiku.rag.client.scope import DatabaseScope
|
||||||
from haiku.rag.client.session import SingleDatabaseSession, default_db_path
|
from haiku.rag.client.session import SingleDatabaseSession
|
||||||
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
|
|
||||||
config = AppConfig(lancedb=LanceDBConfig(databases={"alpha": location}))
|
config = AppConfig(lancedb=LanceDBConfig(databases={"alpha": location}))
|
||||||
[ref] = DatabaseScope.resolve(config, database_name="alpha").databases
|
[ref] = DatabaseScope.resolve(config, database_name="alpha").databases
|
||||||
one, db_path = ref.connection(config)
|
return SingleDatabaseSession(ref, config)
|
||||||
return SingleDatabaseSession(
|
|
||||||
db_path if db_path is not None else default_db_path(one),
|
|
||||||
one,
|
|
||||||
source="alpha",
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_a_named_remote_database_reports_its_uri(self):
|
def test_a_named_remote_database_reports_its_uri(self):
|
||||||
session = self._session("s3://bucket/alpha.lancedb")
|
session = self._session("s3://bucket/alpha.lancedb")
|
||||||
|
|
||||||
assert isinstance(session.db_path, Path)
|
assert session.db_path is None
|
||||||
assert session.location == "s3://bucket/alpha.lancedb"
|
assert session.location == "s3://bucket/alpha.lancedb"
|
||||||
|
assert session.source == "alpha"
|
||||||
|
|
||||||
def test_a_named_local_database_reports_its_path(self):
|
def test_a_named_local_database_reports_its_path(self):
|
||||||
session = self._session("/data/alpha.lancedb")
|
session = self._session("/data/alpha.lancedb")
|
||||||
|
|
|
||||||
|
|
@ -4,53 +4,44 @@ from unittest.mock import AsyncMock, patch
|
||||||
import pytest
|
import pytest
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
|
|
||||||
from haiku.rag.config import get_config
|
|
||||||
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
from haiku.rag.store.engine import ConnectionMode, Store, connect_lancedb
|
from haiku.rag.store.engine import ConnectionMode, Store, connect_lancedb
|
||||||
|
|
||||||
|
|
||||||
class TestConnectionMode:
|
class TestConnectionMode:
|
||||||
def test_local_when_uri_empty(self):
|
"""A location is classified by itself: a path is local, `db://` is LanceDB
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
Cloud, any other scheme is object storage."""
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.LOCAL
|
|
||||||
|
def test_a_path_is_local(self, tmp_path):
|
||||||
|
assert ConnectionMode.of(tmp_path / "db.lancedb") == ConnectionMode.LOCAL
|
||||||
|
|
||||||
|
def test_a_schemeless_string_is_local(self):
|
||||||
|
assert ConnectionMode.of("/data/db.lancedb") == ConnectionMode.LOCAL
|
||||||
|
|
||||||
def test_cloud_when_db_uri(self):
|
def test_cloud_when_db_uri(self):
|
||||||
config = AppConfig(
|
assert ConnectionMode.of("db://my-database") == ConnectionMode.CLOUD
|
||||||
lancedb=LanceDBConfig(
|
|
||||||
uri="db://my-database", api_key="key", region="us-east-1"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.CLOUD
|
|
||||||
|
|
||||||
def test_object_storage_s3(self):
|
@pytest.mark.parametrize(
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
|
"uri",
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
|
[
|
||||||
|
"s3://bucket/path",
|
||||||
def test_object_storage_gs(self):
|
"gs://bucket/path",
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="gs://bucket/path"))
|
"az://container/path",
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
|
"hdfs://namenode/path",
|
||||||
|
"custom://something",
|
||||||
def test_object_storage_az(self):
|
],
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="az://container/path"))
|
)
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
|
def test_any_other_scheme_is_object_storage(self, uri):
|
||||||
|
assert ConnectionMode.of(uri) == ConnectionMode.OBJECT_STORAGE
|
||||||
def test_object_storage_hdfs(self):
|
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="hdfs://namenode/path"))
|
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
|
|
||||||
|
|
||||||
def test_unknown_uri_treated_as_object_storage(self):
|
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="custom://something"))
|
|
||||||
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
|
|
||||||
|
|
||||||
|
|
||||||
class TestConnectLancedb:
|
class TestConnectLancedb:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_local_passes_absolute_db_path(self, temp_db_path):
|
async def test_local_passes_absolute_db_path(self, temp_db_path):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config, db_path=temp_db_path)
|
await connect_lancedb(temp_db_path, AppConfig())
|
||||||
mock_connect.assert_awaited_once()
|
mock_connect.assert_awaited_once()
|
||||||
assert mock_connect.call_args.args == (temp_db_path.absolute(),)
|
assert mock_connect.call_args.args == (temp_db_path.absolute(),)
|
||||||
|
|
||||||
|
|
@ -60,25 +51,36 @@ class TestConnectLancedb:
|
||||||
|
|
||||||
monkeypatch.chdir(tmp_path)
|
monkeypatch.chdir(tmp_path)
|
||||||
relative = Path("db/rag.lancedb")
|
relative = Path("db/rag.lancedb")
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config, db_path=relative)
|
await connect_lancedb(relative, AppConfig())
|
||||||
mock_connect.assert_awaited_once()
|
mock_connect.assert_awaited_once()
|
||||||
assert mock_connect.call_args.args == (relative.absolute(),)
|
assert mock_connect.call_args.args == (relative.absolute(),)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cloud_passes_uri_api_key_region(self):
|
async def test_the_configured_uri_is_not_consulted(self, temp_db_path):
|
||||||
|
"""Storage connects to the location it is handed; placement is the
|
||||||
|
caller's, and the configuration's own `uri` never redirects it."""
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(databases={"elsewhere": "s3://elsewhere/db.lancedb"})
|
||||||
uri="db://my-database", api_key="test-key", region="us-west-2"
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb(temp_db_path, config)
|
||||||
|
assert mock_connect.call_args.args == (temp_db_path.absolute(),)
|
||||||
|
assert "uri" not in mock_connect.call_args.kwargs
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cloud_passes_uri_api_key_region(self):
|
||||||
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(api_key="test-key", region="us-west-2")
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
|
) as mock_connect:
|
||||||
|
await connect_lancedb("db://my-database", config)
|
||||||
mock_connect.assert_awaited_once()
|
mock_connect.assert_awaited_once()
|
||||||
kwargs = mock_connect.call_args.kwargs
|
kwargs = mock_connect.call_args.kwargs
|
||||||
assert kwargs["uri"] == "db://my-database"
|
assert kwargs["uri"] == "db://my-database"
|
||||||
|
|
@ -89,7 +91,6 @@ class TestConnectLancedb:
|
||||||
async def test_object_storage_passes_uri_and_storage_options(self):
|
async def test_object_storage_passes_uri_and_storage_options(self):
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
|
||||||
storage_options={
|
storage_options={
|
||||||
"endpoint": "http://minio:9000",
|
"endpoint": "http://minio:9000",
|
||||||
"region": "us-east-1",
|
"region": "us-east-1",
|
||||||
|
|
@ -99,7 +100,7 @@ class TestConnectLancedb:
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
mock_connect.assert_awaited_once()
|
mock_connect.assert_awaited_once()
|
||||||
kwargs = mock_connect.call_args.kwargs
|
kwargs = mock_connect.call_args.kwargs
|
||||||
assert kwargs["uri"] == "s3://bucket/path"
|
assert kwargs["uri"] == "s3://bucket/path"
|
||||||
|
|
@ -110,23 +111,25 @@ class TestConnectLancedb:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_object_storage_without_storage_options(self):
|
async def test_object_storage_without_storage_options(self):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", AppConfig())
|
||||||
mock_connect.assert_awaited_once()
|
mock_connect.assert_awaited_once()
|
||||||
kwargs = mock_connect.call_args.kwargs
|
kwargs = mock_connect.call_args.kwargs
|
||||||
assert kwargs["uri"] == "s3://bucket/path"
|
assert kwargs["uri"] == "s3://bucket/path"
|
||||||
assert "storage_options" not in kwargs
|
assert "storage_options" not in kwargs
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_local_without_db_path_raises(self):
|
def _remote_store(location: str, config: AppConfig | None = None) -> Store:
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
"""A store over a remote location, opened against a mocked connection."""
|
||||||
with pytest.raises(
|
return Store(
|
||||||
ValueError, match="No lancedb.uri configured and no db_path provided"
|
location,
|
||||||
):
|
config=config,
|
||||||
await connect_lancedb(config)
|
create=True,
|
||||||
|
skip_validation=True,
|
||||||
|
skip_migration_check=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestStoreConnectionMode:
|
class TestStoreConnectionMode:
|
||||||
|
|
@ -134,132 +137,129 @@ class TestStoreConnectionMode:
|
||||||
async def test_store_connection_mode_local(self, temp_db_path):
|
async def test_store_connection_mode_local(self, temp_db_path):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
async with Store(temp_db_path, create=True) as store:
|
||||||
assert store._connection_mode == ConnectionMode.LOCAL
|
assert store._connection_mode == ConnectionMode.LOCAL
|
||||||
|
assert store.location == temp_db_path
|
||||||
|
assert store.db_path == temp_db_path
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_store_connection_mode_cloud(self, temp_db_path):
|
async def test_a_local_store_ignores_the_configured_uri(self, temp_db_path):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
config = AppConfig(
|
||||||
with (
|
lancedb=LanceDBConfig(databases={"elsewhere": "s3://elsewhere/db.lancedb"})
|
||||||
patch.object(get_config().lancedb, "uri", "db://test-database"),
|
)
|
||||||
patch.object(get_config().lancedb, "api_key", "test-api-key"),
|
async with Store(temp_db_path, config=config, create=True) as store:
|
||||||
patch.object(get_config().lancedb, "region", "us-east-1"),
|
assert store._connection_mode == ConnectionMode.LOCAL
|
||||||
):
|
assert store.db_path == temp_db_path
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_store_connection_mode_cloud(self):
|
||||||
|
config = AppConfig(lancedb=LanceDBConfig(api_key="key", region="us-east-1"))
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
|
),
|
||||||
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
|
):
|
||||||
|
async with _remote_store("db://test-database", config) as store:
|
||||||
assert store._connection_mode == ConnectionMode.CLOUD
|
assert store._connection_mode == ConnectionMode.CLOUD
|
||||||
|
assert store.location == "db://test-database"
|
||||||
|
assert store.db_path is None
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_store_connection_mode_object_storage(self, temp_db_path):
|
async def test_store_connection_mode_object_storage(self):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
with (
|
||||||
with patch.object(get_config().lancedb, "uri", "s3://bucket/path"):
|
patch(
|
||||||
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
|
),
|
||||||
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
|
):
|
||||||
|
async with _remote_store("s3://bucket/path") as store:
|
||||||
assert store._connection_mode == ConnectionMode.OBJECT_STORAGE
|
assert store._connection_mode == ConnectionMode.OBJECT_STORAGE
|
||||||
|
assert store.db_path is None
|
||||||
|
|
||||||
|
|
||||||
|
def _remote_store_with_mock_tables(location: str) -> Store:
|
||||||
|
"""A remote store whose tables are mocks: the mode decision is under test,
|
||||||
|
not the tables."""
|
||||||
|
store = _remote_store(location)
|
||||||
|
store.chunks_table = AsyncMock()
|
||||||
|
return store
|
||||||
|
|
||||||
|
|
||||||
class TestVacuumByConnectionMode:
|
class TestVacuumByConnectionMode:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cloud_skips_vacuum(self, temp_db_path):
|
async def test_cloud_skips_vacuum(self):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
with (
|
||||||
with (
|
patch(
|
||||||
patch.object(get_config().lancedb, "uri", "db://test-database"),
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
patch.object(get_config().lancedb, "api_key", "test-api-key"),
|
),
|
||||||
patch.object(get_config().lancedb, "region", "us-east-1"),
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
):
|
):
|
||||||
with patch.object(
|
async with _remote_store_with_mock_tables("db://test-database") as store:
|
||||||
store.chunks_table, "optimize", new_callable=AsyncMock
|
await store.vacuum()
|
||||||
) as mock_optimize:
|
store.chunks_table.optimize.assert_not_awaited()
|
||||||
await store.vacuum()
|
|
||||||
mock_optimize.assert_not_called()
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_object_storage_runs_vacuum(self, temp_db_path):
|
async def test_object_storage_runs_vacuum(self):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
with (
|
||||||
with patch.object(get_config().lancedb, "uri", "s3://bucket/path"):
|
patch(
|
||||||
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
|
),
|
||||||
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
|
):
|
||||||
|
async with _remote_store_with_mock_tables("s3://bucket/path") as store:
|
||||||
|
store.chunks_table.tags.list = AsyncMock(return_value={})
|
||||||
with patch.object(
|
with patch.object(
|
||||||
store.chunks_table, "optimize", new_callable=AsyncMock
|
store, "_tables", return_value={"chunks": store.chunks_table}
|
||||||
) as mock_optimize:
|
):
|
||||||
await store.vacuum()
|
await store.vacuum()
|
||||||
mock_optimize.assert_called()
|
store.chunks_table.optimize.assert_awaited_once()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_local_runs_vacuum(self, temp_db_path):
|
async def test_local_runs_vacuum(self, temp_db_path):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
async with Store(temp_db_path, create=True) as store:
|
||||||
with patch.object(get_config().lancedb, "uri", ""):
|
with patch.object(
|
||||||
with patch.object(
|
store.chunks_table, "optimize", new_callable=AsyncMock
|
||||||
store.chunks_table, "optimize", new_callable=AsyncMock
|
) as mock_optimize:
|
||||||
) as mock_optimize:
|
await store.vacuum()
|
||||||
await store.vacuum()
|
mock_optimize.assert_called()
|
||||||
mock_optimize.assert_called()
|
|
||||||
|
|
||||||
|
|
||||||
class TestVectorIndexByConnectionMode:
|
class TestVectorIndexByConnectionMode:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cloud_skips_index_creation(self, temp_db_path):
|
async def test_cloud_skips_index_creation(self):
|
||||||
async with Store(temp_db_path, create=True) as store:
|
with (
|
||||||
with (
|
patch(
|
||||||
patch.object(get_config().lancedb, "uri", "db://test-database"),
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
patch.object(get_config().lancedb, "api_key", "test-api-key"),
|
),
|
||||||
patch.object(get_config().lancedb, "region", "us-east-1"),
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
):
|
|
||||||
with patch.object(
|
|
||||||
store.chunks_table, "count_rows", new_callable=AsyncMock
|
|
||||||
) as mock_count:
|
|
||||||
await store._ensure_vector_index()
|
|
||||||
mock_count.assert_not_called()
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_object_storage_runs_index_creation(self, temp_db_path):
|
|
||||||
async with Store(temp_db_path, create=True) as store:
|
|
||||||
with patch.object(get_config().lancedb, "uri", "s3://bucket/path"):
|
|
||||||
with patch.object(
|
|
||||||
store.chunks_table,
|
|
||||||
"count_rows",
|
|
||||||
new_callable=AsyncMock,
|
|
||||||
return_value=0,
|
|
||||||
) as mock_count:
|
|
||||||
await store._ensure_vector_index()
|
|
||||||
mock_count.assert_called()
|
|
||||||
|
|
||||||
|
|
||||||
class TestStoreSkipsPathValidationForRemote:
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_skips_path_check_for_cloud(self, tmp_path):
|
|
||||||
nonexistent = tmp_path / "does_not_exist" / "db.lancedb"
|
|
||||||
config = AppConfig(
|
|
||||||
lancedb=LanceDBConfig(
|
|
||||||
uri="db://test-database", api_key="key", region="us-east-1"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
with patch(
|
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
|
||||||
):
|
):
|
||||||
with patch.object(Store, "_init_tables", new_callable=AsyncMock):
|
async with _remote_store_with_mock_tables("db://test-database") as store:
|
||||||
async with Store(
|
await store._ensure_vector_index()
|
||||||
nonexistent,
|
store.chunks_table.count_rows.assert_not_awaited()
|
||||||
config=config,
|
|
||||||
create=True,
|
|
||||||
skip_validation=True,
|
|
||||||
skip_migration_check=True,
|
|
||||||
) as store:
|
|
||||||
assert store is not None
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_skips_path_check_for_object_storage(self, tmp_path):
|
async def test_object_storage_runs_index_creation(self):
|
||||||
nonexistent = tmp_path / "does_not_exist" / "db.lancedb"
|
with (
|
||||||
config = AppConfig(
|
patch(
|
||||||
lancedb=LanceDBConfig(
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
uri="s3://bucket/path",
|
),
|
||||||
storage_options={"endpoint": "http://localhost:9000"},
|
patch.object(Store, "_init_tables", new_callable=AsyncMock),
|
||||||
)
|
|
||||||
)
|
|
||||||
with patch(
|
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
|
||||||
):
|
):
|
||||||
with patch.object(Store, "_init_tables", new_callable=AsyncMock):
|
async with _remote_store_with_mock_tables("s3://bucket/path") as store:
|
||||||
async with Store(
|
store.chunks_table.count_rows = AsyncMock(return_value=0)
|
||||||
nonexistent,
|
await store._ensure_vector_index()
|
||||||
config=config,
|
store.chunks_table.count_rows.assert_awaited_once()
|
||||||
create=True,
|
|
||||||
skip_validation=True,
|
|
||||||
skip_migration_check=True,
|
class TestLocationIsFixed:
|
||||||
) as store:
|
@pytest.mark.asyncio
|
||||||
assert store is not None
|
async def test_a_store_keeps_the_location_it_opened(self, temp_db_path):
|
||||||
|
"""`db_path` and the connection mode derive from the location once; a
|
||||||
|
store cannot be pointed elsewhere after it is built."""
|
||||||
|
async with Store(temp_db_path, create=True) as store:
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
store.location = "s3://bucket/path" # type: ignore[misc]
|
||||||
|
assert store.location == temp_db_path
|
||||||
|
assert store._connection_mode == ConnectionMode.LOCAL
|
||||||
|
|
||||||
|
|
||||||
class TestInitFailureCleanup:
|
class TestInitFailureCleanup:
|
||||||
|
|
@ -412,33 +412,25 @@ class TestStoreMiscellany:
|
||||||
class TestSessionAndConsistency:
|
class TestSessionAndConsistency:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_session_is_shared_across_connections(self):
|
async def test_session_is_shared_across_connections(self):
|
||||||
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
|
config = AppConfig()
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
|
|
||||||
sessions = [c.kwargs["session"] for c in mock_connect.call_args_list]
|
sessions = [c.kwargs["session"] for c in mock_connect.call_args_list]
|
||||||
assert sessions[0] is sessions[1]
|
assert sessions[0] is sessions[1]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cache_sizes_select_distinct_sessions(self):
|
async def test_cache_sizes_select_distinct_sessions(self):
|
||||||
small = AppConfig(
|
small = AppConfig(lancedb=LanceDBConfig(index_cache_size_bytes=1 << 20))
|
||||||
lancedb=LanceDBConfig(
|
large = AppConfig(lancedb=LanceDBConfig(index_cache_size_bytes=1 << 30))
|
||||||
uri="s3://bucket/path", index_cache_size_bytes=1 << 20
|
|
||||||
)
|
|
||||||
)
|
|
||||||
large = AppConfig(
|
|
||||||
lancedb=LanceDBConfig(
|
|
||||||
uri="s3://bucket/path", index_cache_size_bytes=1 << 30
|
|
||||||
)
|
|
||||||
)
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(small)
|
await connect_lancedb("s3://bucket/path", small)
|
||||||
await connect_lancedb(large)
|
await connect_lancedb("s3://bucket/path", large)
|
||||||
|
|
||||||
sessions = [c.kwargs["session"] for c in mock_connect.call_args_list]
|
sessions = [c.kwargs["session"] for c in mock_connect.call_args_list]
|
||||||
assert sessions[0] is not sessions[1]
|
assert sessions[0] is not sessions[1]
|
||||||
|
|
@ -447,7 +439,6 @@ class TestSessionAndConsistency:
|
||||||
async def test_both_cache_sizes_are_applied(self):
|
async def test_both_cache_sizes_are_applied(self):
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri="s3://bucket/path",
|
|
||||||
index_cache_size_bytes=2 << 20,
|
index_cache_size_bytes=2 << 20,
|
||||||
metadata_cache_size_bytes=4 << 20,
|
metadata_cache_size_bytes=4 << 20,
|
||||||
)
|
)
|
||||||
|
|
@ -458,7 +449,7 @@ class TestSessionAndConsistency:
|
||||||
),
|
),
|
||||||
patch("haiku.rag.store.engine.lancedb.Session") as mock_session,
|
patch("haiku.rag.store.engine.lancedb.Session") as mock_session,
|
||||||
):
|
):
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
|
|
||||||
mock_session.assert_called_once_with(
|
mock_session.assert_called_once_with(
|
||||||
index_cache_size_bytes=2 << 20, metadata_cache_size_bytes=4 << 20
|
index_cache_size_bytes=2 << 20, metadata_cache_size_bytes=4 << 20
|
||||||
|
|
@ -466,15 +457,11 @@ class TestSessionAndConsistency:
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_read_consistency_interval_is_forwarded(self):
|
async def test_read_consistency_interval_is_forwarded(self):
|
||||||
config = AppConfig(
|
config = AppConfig(lancedb=LanceDBConfig(read_consistency_interval_seconds=5))
|
||||||
lancedb=LanceDBConfig(
|
|
||||||
uri="s3://bucket/path", read_consistency_interval_seconds=5
|
|
||||||
)
|
|
||||||
)
|
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
|
|
||||||
assert mock_connect.call_args.kwargs["read_consistency_interval"] == timedelta(
|
assert mock_connect.call_args.kwargs["read_consistency_interval"] == timedelta(
|
||||||
seconds=5
|
seconds=5
|
||||||
|
|
@ -483,14 +470,12 @@ class TestSessionAndConsistency:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_read_consistency_interval_omitted_when_disabled(self):
|
async def test_read_consistency_interval_omitted_when_disabled(self):
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(read_consistency_interval_seconds=None)
|
||||||
uri="s3://bucket/path", read_consistency_interval_seconds=None
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config)
|
await connect_lancedb("s3://bucket/path", config)
|
||||||
|
|
||||||
assert mock_connect.call_args.kwargs["read_consistency_interval"] is None
|
assert mock_connect.call_args.kwargs["read_consistency_interval"] is None
|
||||||
|
|
||||||
|
|
@ -500,7 +485,7 @@ class TestSessionAndConsistency:
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
|
||||||
) as mock_connect:
|
) as mock_connect:
|
||||||
await connect_lancedb(config, tmp_path / "db.lancedb")
|
await connect_lancedb(tmp_path / "db.lancedb", config)
|
||||||
|
|
||||||
assert mock_connect.call_args.kwargs["session"] is not None
|
assert mock_connect.call_args.kwargs["session"] is not None
|
||||||
assert mock_connect.call_args.kwargs["read_consistency_interval"] == timedelta(
|
assert mock_connect.call_args.kwargs["read_consistency_interval"] == timedelta(
|
||||||
|
|
|
||||||
|
|
@ -665,6 +665,18 @@ class TestMCPClientLifetime:
|
||||||
with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"):
|
with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"):
|
||||||
create_mcp_server(config=config)
|
create_mcp_server(config=config)
|
||||||
|
|
||||||
|
def test_the_public_factory_refuses_a_path_beside_a_configured_set(self, tmp_path):
|
||||||
|
"""A path and `lancedb.databases` both place the database."""
|
||||||
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
|
from haiku.rag.store.exceptions import AmbiguousDatabaseError
|
||||||
|
|
||||||
|
config = AppConfig(
|
||||||
|
lancedb=LanceDBConfig(databases={"alpha": str(tmp_path / "a")})
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(AmbiguousDatabaseError, match="alpha"):
|
||||||
|
create_mcp_server(tmp_path / "other.lancedb", config=config)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_the_command_hands_the_server_its_resolved_database(
|
async def test_the_command_hands_the_server_its_resolved_database(
|
||||||
self, monkeypatch
|
self, monkeypatch
|
||||||
|
|
@ -697,7 +709,7 @@ class TestMCPClientLifetime:
|
||||||
|
|
||||||
[ref] = seen["scope"].databases
|
[ref] = seen["scope"].databases
|
||||||
assert ref.name == "prod"
|
assert ref.name == "prod"
|
||||||
assert ref.uri == "s3://bucket/prod.lancedb"
|
assert ref.location == "s3://bucket/prod.lancedb"
|
||||||
# The caller's configuration, not one derived from the ref.
|
# The caller's configuration, not one derived from the ref.
|
||||||
assert seen["config"].lancedb.databases == {"prod": "s3://bucket/prod.lancedb"}
|
assert seen["config"].lancedb.databases == {"prod": "s3://bucket/prod.lancedb"}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -41,12 +41,18 @@ def _make_config() -> AppConfig:
|
||||||
unique_prefix = uuid4().hex[:8]
|
unique_prefix = uuid4().hex[:8]
|
||||||
return AppConfig(
|
return AppConfig(
|
||||||
lancedb=LanceDBConfig(
|
lancedb=LanceDBConfig(
|
||||||
uri=f"s3://{S3_BUCKET}/test-{unique_prefix}",
|
databases={"test": f"s3://{S3_BUCKET}/test-{unique_prefix}"},
|
||||||
storage_options=S3_STORAGE_OPTIONS,
|
storage_options=S3_STORAGE_OPTIONS,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _uri(config: AppConfig) -> str:
|
||||||
|
"""The one configured S3 location."""
|
||||||
|
[uri] = config.lancedb.databases.values()
|
||||||
|
return uri
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def config():
|
def config():
|
||||||
"""A config pointing at a unique S3 prefix, cleaned up after the test.
|
"""A config pointing at a unique S3 prefix, cleaned up after the test.
|
||||||
|
|
@ -58,7 +64,7 @@ def config():
|
||||||
config = _make_config()
|
config = _make_config()
|
||||||
yield config
|
yield config
|
||||||
|
|
||||||
bucket, _, prefix = config.lancedb.uri.removeprefix("s3://").partition("/")
|
bucket, _, prefix = _uri(config).removeprefix("s3://").partition("/")
|
||||||
store = make_s3_store(bucket, S3_STORAGE_OPTIONS)
|
store = make_s3_store(bucket, S3_STORAGE_OPTIONS)
|
||||||
paths = [obj["path"] for batch in store.list(prefix=f"{prefix}/") for obj in batch]
|
paths = [obj["path"] for batch in store.list(prefix=f"{prefix}/") for obj in batch]
|
||||||
if paths:
|
if paths:
|
||||||
|
|
@ -73,7 +79,7 @@ def _remote_scope(config: AppConfig) -> DatabaseScope:
|
||||||
"""
|
"""
|
||||||
scope = DatabaseScope.resolve(config)
|
scope = DatabaseScope.resolve(config)
|
||||||
[ref] = scope.databases
|
[ref] = scope.databases
|
||||||
assert ref.db_path is None and ref.uri.startswith("s3://")
|
assert ref.db_path is None and str(ref.location).startswith("s3://")
|
||||||
return scope
|
return scope
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -89,7 +95,7 @@ async def _remote_client(config: AppConfig):
|
||||||
async def test_store_connect_and_create(tmp_path, config):
|
async def test_store_connect_and_create(tmp_path, config):
|
||||||
from haiku.rag.store.info import get_database_stats
|
from haiku.rag.store.info import get_database_stats
|
||||||
|
|
||||||
async with Store(tmp_path / "unused", config=config, create=True) as store:
|
async with Store(_uri(config), config=config, create=True) as store:
|
||||||
stats = await get_database_stats(store.db)
|
stats = await get_database_stats(store.db)
|
||||||
assert stats["documents"]["exists"]
|
assert stats["documents"]["exists"]
|
||||||
assert stats["chunks"]["exists"]
|
assert stats["chunks"]["exists"]
|
||||||
|
|
@ -97,7 +103,7 @@ async def test_store_connect_and_create(tmp_path, config):
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_store_vacuum(tmp_path, config):
|
async def test_store_vacuum(tmp_path, config):
|
||||||
async with Store(tmp_path / "unused", config=config, create=True) as store:
|
async with Store(_uri(config), config=config, create=True) as store:
|
||||||
await store.vacuum()
|
await store.vacuum()
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -106,7 +112,7 @@ async def test_store_add_document(tmp_path, config):
|
||||||
from haiku.rag.store.info import get_database_stats
|
from haiku.rag.store.info import get_database_stats
|
||||||
from haiku.rag.store.schema import DocumentRecord
|
from haiku.rag.store.schema import DocumentRecord
|
||||||
|
|
||||||
async with Store(tmp_path / "unused", config=config, create=True) as store:
|
async with Store(_uri(config), config=config, create=True) as store:
|
||||||
doc = DocumentRecord(content="The quick brown fox jumps over the lazy dog.")
|
doc = DocumentRecord(content="The quick brown fox jumps over the lazy dog.")
|
||||||
await store.documents_table.add([doc])
|
await store.documents_table.add([doc])
|
||||||
|
|
||||||
|
|
@ -164,7 +170,7 @@ async def test_app_info(capsys, config):
|
||||||
|
|
||||||
out = capsys.readouterr().out
|
out = capsys.readouterr().out
|
||||||
assert "path:" in out
|
assert "path:" in out
|
||||||
assert config.lancedb.uri in out
|
assert _uri(config) in out
|
||||||
assert "documents: 1" in out
|
assert "documents: 1" in out
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -885,7 +885,7 @@ async def test_format_citations_rich_omits_the_database_for_one_database():
|
||||||
)
|
)
|
||||||
client = AsyncMock()
|
client = AsyncMock()
|
||||||
client.covers_multiple = False
|
client.covers_multiple = False
|
||||||
client.source_names = ()
|
client.source_names = ("papers",)
|
||||||
|
|
||||||
output = _render_rich(await format_citations_rich([citation], client))
|
output = _render_rich(await format_citations_rich([citation], client))
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue