haiku.rag/tests/test_lancedb_connection.py
Yiorgis Gozadinos f6acb65e95
Reach and enforce 100% coverage
Cover the remaining paths in the client, context, downloads, title
generation, document tools and store models, and add fail_under=100 so
uncovered lines fail CI.

Six lines that no test can reach get a pragma with its reason: the docling
import guard, the nameless PDF attachment, the FS symlink OSError guard that
resolve(strict=False) absorbs, the docling bbox and LanceDB document-id
shape guards, the tag-retention branch vacuum makes unreachable, and Monty's
Rust-thread print callback.

Fix test_find_config_file_user_config, which wrote its config into the cwd it
had chdir'd to, so the cwd branch answered first and the user-directory
lookup it names was never exercised.
2026-07-26 20:11:02 +03:00

390 lines
15 KiB
Python

from unittest.mock import AsyncMock, patch
import pytest
from haiku.rag.config import Config
from haiku.rag.config.models import AppConfig, LanceDBConfig
from haiku.rag.store.engine import ConnectionMode, Store, connect_lancedb
class TestConnectionMode:
def test_local_when_uri_empty(self):
config = AppConfig(lancedb=LanceDBConfig(uri=""))
assert ConnectionMode.from_config(config) == ConnectionMode.LOCAL
def test_cloud_when_db_uri(self):
config = AppConfig(
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):
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
def test_object_storage_gs(self):
config = AppConfig(lancedb=LanceDBConfig(uri="gs://bucket/path"))
assert ConnectionMode.from_config(config) == ConnectionMode.OBJECT_STORAGE
def test_object_storage_az(self):
config = AppConfig(lancedb=LanceDBConfig(uri="az://container/path"))
assert ConnectionMode.from_config(config) == 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:
@pytest.mark.asyncio
async def test_local_passes_absolute_db_path(self, temp_db_path):
config = AppConfig(lancedb=LanceDBConfig(uri=""))
with patch(
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
) as mock_connect:
await connect_lancedb(config, db_path=temp_db_path)
mock_connect.assert_called_once_with(temp_db_path.absolute())
@pytest.mark.asyncio
async def test_local_resolves_relative_db_path(self, tmp_path, monkeypatch):
from pathlib import Path
monkeypatch.chdir(tmp_path)
relative = Path("db/rag.lancedb")
config = AppConfig(lancedb=LanceDBConfig(uri=""))
with patch(
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
) as mock_connect:
await connect_lancedb(config, db_path=relative)
mock_connect.assert_called_once_with(relative.absolute())
@pytest.mark.asyncio
async def test_cloud_passes_uri_api_key_region(self):
config = AppConfig(
lancedb=LanceDBConfig(
uri="db://my-database", 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(config)
mock_connect.assert_called_once_with(
uri="db://my-database", api_key="test-key", region="us-west-2"
)
@pytest.mark.asyncio
async def test_object_storage_passes_uri_and_storage_options(self):
config = AppConfig(
lancedb=LanceDBConfig(
uri="s3://bucket/path",
storage_options={
"endpoint": "http://minio:9000",
"region": "us-east-1",
},
)
)
with patch(
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
) as mock_connect:
await connect_lancedb(config)
mock_connect.assert_called_once_with(
uri="s3://bucket/path",
storage_options={
"endpoint": "http://minio:9000",
"region": "us-east-1",
},
)
@pytest.mark.asyncio
async def test_object_storage_without_storage_options(self):
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
with patch(
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
) as mock_connect:
await connect_lancedb(config)
mock_connect.assert_called_once_with(uri="s3://bucket/path")
@pytest.mark.asyncio
async def test_local_without_db_path_raises(self):
config = AppConfig(lancedb=LanceDBConfig(uri=""))
with pytest.raises(
ValueError, match="No lancedb.uri configured and no db_path provided"
):
await connect_lancedb(config)
class TestStoreConnectionMode:
@pytest.mark.asyncio
async def test_store_connection_mode_local(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
assert store._connection_mode == ConnectionMode.LOCAL
@pytest.mark.asyncio
async def test_store_connection_mode_cloud(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with (
patch.object(Config.lancedb, "uri", "db://test-database"),
patch.object(Config.lancedb, "api_key", "test-api-key"),
patch.object(Config.lancedb, "region", "us-east-1"),
):
assert store._connection_mode == ConnectionMode.CLOUD
@pytest.mark.asyncio
async def test_store_connection_mode_object_storage(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with patch.object(Config.lancedb, "uri", "s3://bucket/path"):
assert store._connection_mode == ConnectionMode.OBJECT_STORAGE
class TestVacuumByConnectionMode:
@pytest.mark.asyncio
async def test_cloud_skips_vacuum(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with (
patch.object(Config.lancedb, "uri", "db://test-database"),
patch.object(Config.lancedb, "api_key", "test-api-key"),
patch.object(Config.lancedb, "region", "us-east-1"),
):
with patch.object(
store.chunks_table, "optimize", new_callable=AsyncMock
) as mock_optimize:
await store.vacuum()
mock_optimize.assert_not_called()
@pytest.mark.asyncio
async def test_object_storage_runs_vacuum(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with patch.object(Config.lancedb, "uri", "s3://bucket/path"):
with patch.object(
store.chunks_table, "optimize", new_callable=AsyncMock
) as mock_optimize:
await store.vacuum()
mock_optimize.assert_called()
@pytest.mark.asyncio
async def test_local_runs_vacuum(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with patch.object(Config.lancedb, "uri", ""):
with patch.object(
store.chunks_table, "optimize", new_callable=AsyncMock
) as mock_optimize:
await store.vacuum()
mock_optimize.assert_called()
class TestVectorIndexByConnectionMode:
@pytest.mark.asyncio
async def test_cloud_skips_index_creation(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with (
patch.object(Config.lancedb, "uri", "db://test-database"),
patch.object(Config.lancedb, "api_key", "test-api-key"),
patch.object(Config.lancedb, "region", "us-east-1"),
):
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(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 Store(
nonexistent,
config=config,
create=True,
skip_validation=True,
skip_migration_check=True,
) as store:
assert store is not None
@pytest.mark.asyncio
async def test_skips_path_check_for_object_storage(self, tmp_path):
nonexistent = tmp_path / "does_not_exist" / "db.lancedb"
config = AppConfig(
lancedb=LanceDBConfig(
uri="s3://bucket/path",
storage_options={"endpoint": "http://localhost:9000"},
)
)
with patch(
"haiku.rag.store.engine.lancedb.connect_async", new_callable=AsyncMock
):
with patch.object(Store, "_init_tables", new_callable=AsyncMock):
async with Store(
nonexistent,
config=config,
create=True,
skip_validation=True,
skip_migration_check=True,
) as store:
assert store is not None
class TestInitFailureCleanup:
@pytest.mark.asyncio
async def test_store_aenter_closes_connection_on_init_failure(
self, temp_db_path, monkeypatch
):
"""If _initialize raises after connect, __aenter__ must close the
AsyncConnection so it doesn't leak (no __aexit__ runs in that case)."""
mock_conn = AsyncMock()
mock_conn.close = lambda: mock_conn.close_calls.append(True) # type: ignore[attr-defined]
mock_conn.close_calls = [] # type: ignore[attr-defined]
async def fake_connect(*args, **kwargs):
return mock_conn
async def failing_init_tables(self, is_new_db):
raise RuntimeError("simulated table init failure")
monkeypatch.setattr("haiku.rag.store.engine.connect_lancedb", fake_connect)
monkeypatch.setattr(Store, "_init_tables", failing_init_tables)
with pytest.raises(RuntimeError, match="simulated table init failure"):
async with Store(temp_db_path, create=True) as store:
assert store is not None
assert mock_conn.close_calls == [True], (
"AsyncConnection.close() was not called on init failure"
)
@pytest.mark.asyncio
async def test_client_aenter_closes_store_on_init_failure(
self, temp_db_path, monkeypatch
):
"""HaikuRAG.__aenter__ must close the store if _initialize fails."""
from haiku.rag.client import HaikuRAG
close_calls: list[bool] = []
original_close = Store.close
def tracking_close(self):
close_calls.append(True)
original_close(self)
async def failing_init(self):
# Set db so close() has something to close
self.db = AsyncMock()
self.db.close = lambda: None
raise RuntimeError("simulated initialize failure")
monkeypatch.setattr(Store, "_initialize", failing_init)
monkeypatch.setattr(Store, "close", tracking_close)
with pytest.raises(RuntimeError, match="simulated initialize failure"):
async with HaikuRAG(temp_db_path, create=True):
pass
assert close_calls, "Store.close() was not called when _initialize raised"
class TestVectorIndexCreation:
"""_ensure_vector_index needs 256 rows of training data before it builds."""
@staticmethod
async def _seed_chunks(store, count: int) -> None:
import random
records = [
store.ChunkRecord(
document_id="doc-1",
content=f"row {i}",
content_fts=f"row {i}",
metadata="{}",
order=i,
vector=[random.random() for _ in range(store.embedder.vector_dim)],
)
for i in range(count)
]
await store.chunks_table.add(records)
@pytest.mark.asyncio
async def test_builds_index_once_enough_rows_exist(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
await self._seed_chunks(store, 256)
await store._ensure_vector_index()
indexes = await store.chunks_table.list_indices()
assert any("vector" in idx.columns for idx in indexes)
@pytest.mark.asyncio
async def test_index_failure_is_warned_not_raised(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
await self._seed_chunks(store, 256)
async def boom(*_args, **_kwargs):
raise RuntimeError("index build failed")
with patch.object(store.chunks_table, "create_index", boom):
await store._ensure_vector_index()
indexes = await store.chunks_table.list_indices()
assert not any("vector" in idx.columns for idx in indexes)
class TestStoreMiscellany:
@pytest.mark.asyncio
async def test_create_makes_missing_parent_directories(self, tmp_path):
nested = tmp_path / "a" / "b" / "db.lancedb"
async with Store(nested, create=True) as store:
assert store._is_new_db is True
assert nested.exists()
@pytest.mark.asyncio
async def test_stored_vector_dim_is_none_for_corrupt_settings(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
await store.settings_table.update(
{"settings": "not json at all"}, where="id = 'settings'"
)
assert await store._get_stored_vector_dim() is None
@pytest.mark.asyncio
async def test_vacuum_skips_when_already_running(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
async with store._vacuum_lock:
# Returns immediately rather than blocking on the held lock.
await store.vacuum()
@pytest.mark.asyncio
async def test_history_rejects_unknown_table(self, temp_db_path):
async with Store(temp_db_path, create=True) as store:
with pytest.raises(ValueError, match="Unknown table"):
await store.list_table_versions("not_a_table")