Add ConnectionMode enum (LOCAL/CLOUD/OBJECT_STORAGE) and connect_lancedb() utility to support S3, GCS, Azure, and HDFS backends via storage_options.
This commit is contained in:
parent
23d2aee955
commit
110accb8e7
3 changed files with 243 additions and 59 deletions
|
|
@ -60,6 +60,7 @@ class LanceDBConfig(BaseModel):
|
||||||
uri: str = ""
|
uri: str = ""
|
||||||
api_key: str = ""
|
api_key: str = ""
|
||||||
region: str = ""
|
region: str = ""
|
||||||
|
storage_options: dict[str, str] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
class EmbeddingsConfig(BaseModel):
|
class EmbeddingsConfig(BaseModel):
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
|
from enum import Enum
|
||||||
from importlib import metadata
|
from importlib import metadata
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
@ -18,6 +19,42 @@ from haiku.rag.store.exceptions import MigrationRequiredError, ReadOnlyError
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
OBJECT_STORAGE_PREFIXES = ("s3://", "gs://", "az://", "hdfs://")
|
||||||
|
|
||||||
|
|
||||||
|
class ConnectionMode(Enum):
|
||||||
|
LOCAL = "local"
|
||||||
|
CLOUD = "cloud"
|
||||||
|
OBJECT_STORAGE = "object_storage"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def from_config(config: AppConfig) -> "ConnectionMode":
|
||||||
|
uri = config.lancedb.uri
|
||||||
|
if not uri:
|
||||||
|
return ConnectionMode.LOCAL
|
||||||
|
if uri.startswith("db://"):
|
||||||
|
return ConnectionMode.CLOUD
|
||||||
|
return ConnectionMode.OBJECT_STORAGE
|
||||||
|
|
||||||
|
|
||||||
|
def connect_lancedb(config: AppConfig, db_path: Path | None = None):
|
||||||
|
mode = ConnectionMode.from_config(config)
|
||||||
|
if mode == ConnectionMode.CLOUD:
|
||||||
|
return lancedb.connect(
|
||||||
|
uri=config.lancedb.uri,
|
||||||
|
api_key=config.lancedb.api_key,
|
||||||
|
region=config.lancedb.region,
|
||||||
|
)
|
||||||
|
elif mode == ConnectionMode.OBJECT_STORAGE:
|
||||||
|
kwargs: dict[str, Any] = {"uri": config.lancedb.uri}
|
||||||
|
if config.lancedb.storage_options:
|
||||||
|
kwargs["storage_options"] = config.lancedb.storage_options
|
||||||
|
return lancedb.connect(**kwargs)
|
||||||
|
else:
|
||||||
|
if db_path is None:
|
||||||
|
raise ValueError("No lancedb.uri configured and no db_path provided")
|
||||||
|
return lancedb.connect(db_path)
|
||||||
|
|
||||||
|
|
||||||
class DocumentRecord(LanceModel):
|
class DocumentRecord(LanceModel):
|
||||||
id: str = Field(default_factory=lambda: str(uuid4()))
|
id: str = Field(default_factory=lambda: str(uuid4()))
|
||||||
|
|
@ -97,7 +134,7 @@ class Store:
|
||||||
|
|
||||||
# Check if database exists (for local filesystem only)
|
# Check if database exists (for local filesystem only)
|
||||||
is_new_db = False
|
is_new_db = False
|
||||||
if not self._has_cloud_config():
|
if self._connection_mode == ConnectionMode.LOCAL:
|
||||||
if not db_path.exists():
|
if not db_path.exists():
|
||||||
if not create:
|
if not create:
|
||||||
raise FileNotFoundError(
|
raise FileNotFoundError(
|
||||||
|
|
@ -110,7 +147,7 @@ class Store:
|
||||||
Path.mkdir(db_path.parent, parents=True)
|
Path.mkdir(db_path.parent, parents=True)
|
||||||
|
|
||||||
# Connect to LanceDB
|
# Connect to LanceDB
|
||||||
self.db = self._connect_to_lancedb(db_path)
|
self.db = connect_lancedb(self._config, db_path)
|
||||||
|
|
||||||
# For existing databases, read stored vector dimension to create ChunkRecord
|
# For existing databases, read stored vector dimension to create ChunkRecord
|
||||||
# that can read existing chunks. For new databases, use config's dimension.
|
# that can read existing chunks. For new databases, use config's dimension.
|
||||||
|
|
@ -198,9 +235,7 @@ class Store:
|
||||||
"""
|
"""
|
||||||
self._assert_writable()
|
self._assert_writable()
|
||||||
|
|
||||||
if self._has_cloud_config() and str(self._config.lancedb.uri).startswith(
|
if self._connection_mode == ConnectionMode.CLOUD:
|
||||||
"db://"
|
|
||||||
):
|
|
||||||
return
|
return
|
||||||
|
|
||||||
# Skip if already running (non-blocking)
|
# Skip if already running (non-blocking)
|
||||||
|
|
@ -224,26 +259,9 @@ class Store:
|
||||||
# Handle resource errors gracefully
|
# Handle resource errors gracefully
|
||||||
logger.debug(f"Vacuum skipped due to resource constraints: {e}")
|
logger.debug(f"Vacuum skipped due to resource constraints: {e}")
|
||||||
|
|
||||||
def _connect_to_lancedb(self, db_path: Path):
|
@property
|
||||||
"""Establish connection to LanceDB (local, cloud, or object storage)."""
|
def _connection_mode(self) -> ConnectionMode:
|
||||||
# Check if we have cloud configuration
|
return ConnectionMode.from_config(self._config)
|
||||||
if self._has_cloud_config():
|
|
||||||
return lancedb.connect(
|
|
||||||
uri=self._config.lancedb.uri,
|
|
||||||
api_key=self._config.lancedb.api_key,
|
|
||||||
region=self._config.lancedb.region,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Local file system connection
|
|
||||||
return lancedb.connect(db_path)
|
|
||||||
|
|
||||||
def _has_cloud_config(self) -> bool:
|
|
||||||
"""Check if cloud configuration is complete."""
|
|
||||||
return bool(
|
|
||||||
self._config.lancedb.uri
|
|
||||||
and self._config.lancedb.api_key
|
|
||||||
and self._config.lancedb.region
|
|
||||||
)
|
|
||||||
|
|
||||||
def get_stats(self) -> dict:
|
def get_stats(self) -> dict:
|
||||||
"""Get comprehensive table statistics.
|
"""Get comprehensive table statistics.
|
||||||
|
|
@ -298,7 +316,7 @@ class Store:
|
||||||
it will be replaced (using replace=True parameter).
|
it will be replaced (using replace=True parameter).
|
||||||
Note: Index creation requires sufficient training data.
|
Note: Index creation requires sufficient training data.
|
||||||
"""
|
"""
|
||||||
if self._has_cloud_config():
|
if self._connection_mode == ConnectionMode.CLOUD:
|
||||||
return
|
return
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
|
||||||
|
|
@ -3,46 +3,211 @@ from unittest.mock import patch
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.config.models import AppConfig, LanceDBConfig
|
||||||
|
from haiku.rag.store.engine import ConnectionMode, Store, connect_lancedb
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
class TestConnectionMode:
|
||||||
async def test_lancedb_cloud_skips_optimization(temp_db_path):
|
def test_local_when_uri_empty(self):
|
||||||
"""Test that vacuum is skipped when using LanceDB Cloud (db:// URI)."""
|
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
||||||
# Create a store
|
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:
|
||||||
|
def test_local_passes_db_path(self, temp_db_path):
|
||||||
|
config = AppConfig(lancedb=LanceDBConfig(uri=""))
|
||||||
|
with patch("haiku.rag.store.engine.lancedb.connect") as mock_connect:
|
||||||
|
connect_lancedb(config, db_path=temp_db_path)
|
||||||
|
mock_connect.assert_called_once_with(temp_db_path)
|
||||||
|
|
||||||
|
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") as mock_connect:
|
||||||
|
connect_lancedb(config)
|
||||||
|
mock_connect.assert_called_once_with(
|
||||||
|
uri="db://my-database", api_key="test-key", region="us-west-2"
|
||||||
|
)
|
||||||
|
|
||||||
|
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") as mock_connect:
|
||||||
|
connect_lancedb(config)
|
||||||
|
mock_connect.assert_called_once_with(
|
||||||
|
uri="s3://bucket/path",
|
||||||
|
storage_options={
|
||||||
|
"endpoint": "http://minio:9000",
|
||||||
|
"region": "us-east-1",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_object_storage_without_storage_options(self):
|
||||||
|
config = AppConfig(lancedb=LanceDBConfig(uri="s3://bucket/path"))
|
||||||
|
with patch("haiku.rag.store.engine.lancedb.connect") as mock_connect:
|
||||||
|
connect_lancedb(config)
|
||||||
|
mock_connect.assert_called_once_with(uri="s3://bucket/path")
|
||||||
|
|
||||||
|
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"
|
||||||
|
):
|
||||||
|
connect_lancedb(config)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStoreConnectionMode:
|
||||||
|
def test_store_connection_mode_local(self, temp_db_path):
|
||||||
store = Store(temp_db_path, create=True)
|
store = Store(temp_db_path, create=True)
|
||||||
|
assert store._connection_mode == ConnectionMode.LOCAL
|
||||||
|
store.close()
|
||||||
|
|
||||||
# Mock all cloud config to simulate LanceDB Cloud usage
|
def test_store_connection_mode_cloud(self, temp_db_path):
|
||||||
|
store = Store(temp_db_path, create=True)
|
||||||
with (
|
with (
|
||||||
patch.object(Config.lancedb, "uri", "db://test-database"),
|
patch.object(Config.lancedb, "uri", "db://test-database"),
|
||||||
patch.object(Config.lancedb, "api_key", "test-api-key"),
|
patch.object(Config.lancedb, "api_key", "test-api-key"),
|
||||||
patch.object(Config.lancedb, "region", "us-east-1"),
|
patch.object(Config.lancedb, "region", "us-east-1"),
|
||||||
):
|
):
|
||||||
# Mock the optimize method to track if it's called
|
assert store._connection_mode == ConnectionMode.CLOUD
|
||||||
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
|
||||||
# Call vacuum - this should skip optimization for LanceDB Cloud
|
|
||||||
await store.vacuum()
|
|
||||||
|
|
||||||
# The optimize method should NOT have been called for LanceDB Cloud
|
|
||||||
mock_optimize.assert_not_called()
|
|
||||||
|
|
||||||
store.close()
|
store.close()
|
||||||
|
|
||||||
|
def test_store_connection_mode_object_storage(self, temp_db_path):
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_local_storage_calls_optimization(temp_db_path):
|
|
||||||
"""Test that vacuum calls optimization for local storage."""
|
|
||||||
# Create a store
|
|
||||||
store = Store(temp_db_path, create=True)
|
store = Store(temp_db_path, create=True)
|
||||||
|
with patch.object(Config.lancedb, "uri", "s3://bucket/path"):
|
||||||
# Ensure uri is empty (local storage)
|
assert store._connection_mode == ConnectionMode.OBJECT_STORAGE
|
||||||
with patch.object(Config.lancedb, "uri", ""):
|
store.close()
|
||||||
# Mock the optimize method to track if it's called
|
|
||||||
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
|
||||||
# Call vacuum - this should optimize all tables for local storage
|
class TestVacuumByConnectionMode:
|
||||||
await store.vacuum()
|
@pytest.mark.asyncio
|
||||||
|
async def test_cloud_skips_vacuum(self, temp_db_path):
|
||||||
# The optimize method SHOULD have been called for local storage
|
store = Store(temp_db_path, create=True)
|
||||||
mock_optimize.assert_called()
|
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") as mock_optimize:
|
||||||
|
await store.vacuum()
|
||||||
|
mock_optimize.assert_not_called()
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_object_storage_runs_vacuum(self, temp_db_path):
|
||||||
|
store = Store(temp_db_path, create=True)
|
||||||
|
with patch.object(Config.lancedb, "uri", "s3://bucket/path"):
|
||||||
|
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
||||||
|
await store.vacuum()
|
||||||
|
mock_optimize.assert_called()
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_local_runs_vacuum(self, temp_db_path):
|
||||||
|
store = Store(temp_db_path, create=True)
|
||||||
|
with patch.object(Config.lancedb, "uri", ""):
|
||||||
|
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
||||||
|
await store.vacuum()
|
||||||
|
mock_optimize.assert_called()
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestVectorIndexByConnectionMode:
|
||||||
|
def test_cloud_skips_index_creation(self, temp_db_path):
|
||||||
|
store = Store(temp_db_path, create=True)
|
||||||
|
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") as mock_count:
|
||||||
|
store._ensure_vector_index()
|
||||||
|
mock_count.assert_not_called()
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
def test_object_storage_runs_index_creation(self, temp_db_path):
|
||||||
|
store = Store(temp_db_path, create=True)
|
||||||
|
with patch.object(Config.lancedb, "uri", "s3://bucket/path"):
|
||||||
|
with patch.object(
|
||||||
|
store.chunks_table, "count_rows", return_value=0
|
||||||
|
) as mock_count:
|
||||||
|
store._ensure_vector_index()
|
||||||
|
mock_count.assert_called()
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestStoreSkipsPathValidationForRemote:
|
||||||
|
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"):
|
||||||
|
with patch.object(Store, "_init_tables"):
|
||||||
|
store = Store(
|
||||||
|
nonexistent,
|
||||||
|
config=config,
|
||||||
|
create=True,
|
||||||
|
skip_validation=True,
|
||||||
|
skip_migration_check=True,
|
||||||
|
)
|
||||||
|
store.close()
|
||||||
|
|
||||||
|
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"):
|
||||||
|
with patch.object(Store, "_init_tables"):
|
||||||
|
store = Store(
|
||||||
|
nonexistent,
|
||||||
|
config=config,
|
||||||
|
create=True,
|
||||||
|
skip_validation=True,
|
||||||
|
skip_migration_check=True,
|
||||||
|
)
|
||||||
store.close()
|
store.close()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue