haiku.rag/tests/test_lancedb_connection.py
Yiorgis Gozadinos f96a428ef1
Fix defects found reviewing the coverage work
check_source_accessible narrowed its handler to ValueError, but Path.exists
re-raises errno values outside its ignored set (EACCES, ENAMETOOLONG). Those
were swallowed before and now escaped into the rebuild sweep the guard exists
to protect. Catch OSError too.

Restore the arity guard in _common_path_prefix: without it an empty list
raises from min() and a single label yields a prefix covering the whole path.

Two tests would have hung rather than failed on regression (the vacuum skip
and the protected-wait cancellation); both are now bounded. The import
vacuum test raced against the done-callback that discards the task, and now
spies on the call instead, with a negative control.

Replace assertions that could not fail: blank-query search against an empty
corpus, a batch flush counted against an empty table, a picture description
asserting its own input state, and an FS scheme check with nothing on disk to
resolve. The get_model matrix asserted only the returned type across 26
cases and now pins the per-provider settings. The three batching tests now
count flushes, which revealed embed-only writes through chunks_table.add
rather than _flush_rebuild_batch.
2026-07-27 10:44:32 +03:00

400 lines
16 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):
import logging
from haiku.rag.store import engine as engine_module
from tests.conftest import capture_logs
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):
with capture_logs(engine_module.logger, logging.WARNING) as records:
await store._ensure_vector_index()
assert any("index build failed" in r.getMessage() for r in records)
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):
import asyncio
async with Store(temp_db_path, create=True) as store:
async with store._vacuum_lock:
# Bounded: a regression here blocks on the held lock, and the
# timeout turns that deadlock into a clean failure.
await asyncio.wait_for(store.vacuum(), timeout=5)
@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")