Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion src/omop_emb/backends/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

from typing import TYPE_CHECKING

from .base_backend import EmbeddingBackend, resolve_backend
from .base_backend import EmbeddingBackend, resolve_backend, resolve_backend_from_config
from .sqlitevec import SQLiteVecEmbeddingBackend

if TYPE_CHECKING:
Expand All @@ -11,6 +11,7 @@
__all__ = [
"EmbeddingBackend",
"resolve_backend",
"resolve_backend_from_config",
"SQLiteVecEmbeddingBackend",
"PGVectorEmbeddingBackend",
]
Expand Down
49 changes: 29 additions & 20 deletions src/omop_emb/backends/base_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from functools import wraps
import logging
from datetime import datetime
from typing import Any, Callable, Generic, Iterable, Mapping, Optional, Sequence, Tuple, TypeVar, Union
from typing import Any, Callable, Generic, Iterable, Mapping, Optional, Sequence, Tuple, TypeVar
from numpy import ndarray
from sqlalchemy import Engine
from sqlalchemy.orm import sessionmaker
Expand Down Expand Up @@ -993,46 +993,45 @@ def validate_embeddings_and_records(


def resolve_backend(
backend_type: Optional[Union[str, BackendType]] = None,
backend_type: str | BackendType,
*,
sqlite_path: Optional[str] = None,
) -> EmbeddingBackend:
"""Return the configured embedding backend via oa-configurator.
"""Return the embedding backend for *backend_type*.

When backend_type is omitted, reads from the active oa-configurator config.
Connection details are resolved via the oa-configurator Resolver:

- ``sqlitevec``: uses sqlite_path from config (required; must be set explicitly).
- ``sqlitevec``: uses *sqlite_path* (required; must be set explicitly).
- ``pgvector``: uses the ``emb_db`` resource from oa-configurator.

Callers that have a loaded ``OmopEmbConfig`` should prefer
:func:`resolve_backend_from_config` over calling this directly.
"""
cfg = OmopEmbConfig.get_config()
if backend_type is None:
backend_str = cfg.backend
else:
backend_str = (
backend_type if isinstance(backend_type, str) else backend_type.value
)
backend_str = (
backend_type if isinstance(backend_type, str) else backend_type.value
)

try:
resolved_backend = BackendType(backend_str.lower())
resolved_backend_type = BackendType(backend_str.lower())
except ValueError:
raise RuntimeError(
f"Unknown backend {backend_str!r}. "
f"Supported: {[b.value for b in BackendType]}."
)

if resolved_backend == BackendType.SQLITEVEC:
if resolved_backend_type == BackendType.SQLITEVEC:
from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend

if not cfg.sqlite_path:
if not sqlite_path:
raise RuntimeError(
"sqlitevec backend requires 'sqlite_path' to be configured. "
"Set it via `omop-config configure omop-emb`. "
"To use an ephemeral in-memory store intentionally, set sqlite_path = ':memory:'."
)
path = cfg.sqlite_path
logger.info(f"Using SQLiteVec backend with database file: {path}")
return SQLiteVecEmbeddingBackend.from_path(path)
logger.info(f"Using SQLiteVec backend with database file: {sqlite_path}")
return SQLiteVecEmbeddingBackend.from_path(sqlite_path)

if resolved_backend == BackendType.PGVECTOR:
if resolved_backend_type == BackendType.PGVECTOR:
engine = resolve_omop_emb_engine()
if engine.dialect.name != "postgresql":
raise RuntimeError(
Expand All @@ -1049,4 +1048,14 @@ def resolve_backend(
logger.info(f"Using pgvector backend with engine: {engine.url}")
return PGVectorEmbeddingBackend(emb_engine=engine)

raise RuntimeError(f"Implementation for {resolved_backend.value} is not available.")
raise RuntimeError(f"Implementation for {resolved_backend_type.value} is not available.")


def resolve_backend_from_config(cfg: OmopEmbConfig) -> EmbeddingBackend:
"""Resolve the embedding backend using an already-loaded config object.

The one place config is read for backend selection; call sites that
already have *cfg* in scope should prefer this over calling
:func:`resolve_backend` directly.
"""
return resolve_backend(cfg.backend, sqlite_path=cfg.sqlite_path)
7 changes: 4 additions & 3 deletions src/omop_emb/cli/cli_diagnostics.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@
import sqlalchemy as sa
import typer

from omop_emb.backends import resolve_backend
from omop_emb.config import MetricType, resolve_omop_cdm_engine
from omop_emb.backends import resolve_backend_from_config
from omop_emb.config import MetricType, load_omop_emb_config, resolve_omop_cdm_engine
from omop_emb.interface import list_registered_models

logger = logging.getLogger(__name__)
Expand All @@ -17,7 +17,8 @@
name="health-check", help="Verify backend connectivity and list registered models."
)
def health_check():
backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
typer.echo(f"Backend: {backend.backend_type.value} | connected.")

# CDM connectivity is optional for the health check
Expand Down
59 changes: 36 additions & 23 deletions src/omop_emb/cli/cli_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,13 @@

from omop_emb.utils.cdm import check_concept_cdm
from omop_emb.backends.index_config import index_config_from_index_type
from omop_emb.backends import resolve_backend
from omop_emb.backends import resolve_backend_from_config
from omop_emb.config import (
IndexType,
MetricType,
OmopEmbConfig,
ProviderType,
load_omop_emb_config,
provider_type_examples,
resolve_omop_cdm_engine,
)
Expand Down Expand Up @@ -53,17 +54,26 @@ def consolidate_queries(
raise ValueError("No queries provided.")


def _get_config() -> OmopEmbConfig:
"""Load and return the active OmopEmbConfig, converting a missing-file
error into an actionable setup message.
"""
try:
return OmopEmbConfig.get_config()
except FileNotFoundError:
raise RuntimeError(
"No omop-emb configuration file found. "
"Run `omop-config configure omop-emb` to set it up."
)
def _build_embedding_client(
cfg: OmopEmbConfig,
*,
model: str,
api_base: str,
api_key: str,
provider_type: ProviderType,
embedding_batch_size: int = 32,
) -> EmbeddingClient:
"""Construct an EmbeddingClient, applying config-derived dim/prefixes."""
return EmbeddingClient(
model=model,
api_base=api_base,
api_key=api_key,
embedding_batch_size=embedding_batch_size,
provider_type=provider_type,
embedding_dim=cfg.embedding_dim,
document_embedding_prefix=cfg.document_embedding_prefix,
query_embedding_prefix=cfg.query_embedding_prefix,
)


def _render_search_results(
Expand Down Expand Up @@ -167,21 +177,22 @@ def add_embeddings(
``create-index`` afterwards to upgrade to an HNSW approximate index.
"""

cfg = _get_config()
cfg = load_omop_emb_config()
resolved_api_base = api_base or cfg.api_base
resolved_api_key = api_key or cfg.api_key
resolved_provider = provider or cfg.provider_type
resolved_model = model or cfg.embedding_model

backend = resolve_backend()
backend = resolve_backend_from_config(cfg)
omop_cdm_engine = resolve_omop_cdm_engine()

embedding_client = EmbeddingClient(
embedding_client = _build_embedding_client(
cfg,
model=resolved_model,
api_base=resolved_api_base,
api_key=resolved_api_key,
embedding_batch_size=batch_size,
provider_type=resolved_provider,
embedding_batch_size=batch_size,
)
# FLAT registration: metric_type=COSINE is used only for upsert validation;
# FLAT accepts any backend-supported metric, so COSINE is always valid here.
Expand Down Expand Up @@ -322,14 +333,15 @@ def create_index(
locked in and all subsequent queries must use the same metric.
"""

cfg = _get_config()
cfg = load_omop_emb_config()
resolved_api_base = api_base or cfg.api_base
resolved_api_key = api_key or cfg.api_key
resolved_provider = provider or cfg.provider_type
resolved_model = model or cfg.embedding_model

backend = resolve_backend()
embedding_client = EmbeddingClient(
backend = resolve_backend_from_config(cfg)
embedding_client = _build_embedding_client(
cfg,
model=resolved_model,
api_base=resolved_api_base,
api_key=resolved_api_key,
Expand Down Expand Up @@ -611,15 +623,15 @@ def search(
] = None,
):

cfg = _get_config()
cfg = load_omop_emb_config()
resolved_api_base = api_base or cfg.api_base
resolved_api_key = api_key or cfg.api_key
resolved_provider = provider or cfg.provider_type
resolved_faiss_cache_dir = faiss_cache_dir or cfg.faiss_cache_dir
resolved_model = model or cfg.embedding_model

queries_generator = consolidate_queries(queries=queries, queries_file=queries_file)
backend = resolve_backend()
backend = resolve_backend_from_config(cfg)

# CDM enrichment is optional for search
try:
Expand All @@ -630,12 +642,13 @@ def search(
"CDM engine not configured; concept names will not be enriched in results."
)

embedding_client = EmbeddingClient(
embedding_client = _build_embedding_client(
cfg,
model=resolved_model,
api_base=resolved_api_base,
api_key=resolved_api_key,
embedding_batch_size=batch_size,
provider_type=resolved_provider,
embedding_batch_size=batch_size,
)
embedding_reader = EmbeddingReaderInterface(
model=embedding_client.canonical_model_name,
Expand Down
7 changes: 4 additions & 3 deletions src/omop_emb/cli/cli_legacy.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,9 +12,9 @@
import typer
from tqdm import tqdm

from omop_emb.backends import resolve_backend
from omop_emb.backends import resolve_backend_from_config
from omop_emb.backends.index_config import IndexConfig, index_config_from_index_type
from omop_emb.config import IndexType, MetricType, ProviderType
from omop_emb.config import IndexType, MetricType, ProviderType, load_omop_emb_config
from omop_emb.storage import embedding_bundle

if TYPE_CHECKING:
Expand Down Expand Up @@ -249,7 +249,8 @@ def import_legacy_faiss_cache(
err=True,
)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
try:
from omop_emb.storage.faiss import FAISSCache
except ImportError as e:
Expand Down
24 changes: 16 additions & 8 deletions src/omop_emb/cli/cli_maintenance.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,13 @@
from typing import Annotated, Optional

import typer
from omop_emb.backends import resolve_backend
from omop_emb.backends import resolve_backend_from_config
from omop_emb.backends.index_config import index_config_from_index_type
from omop_emb.config import (
IndexType,
MetricType,
ProviderType,
load_omop_emb_config,
)
from omop_emb.embeddings.embedding_providers import get_provider_from_provider_type
from omop_emb.interface import list_registered_models
Expand Down Expand Up @@ -38,7 +39,8 @@ def list_models(
] = None,
):

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
records = list_registered_models(
backend=backend,
provider_type=provider_type,
Expand Down Expand Up @@ -128,7 +130,8 @@ def rebuild_index(
embedding_provider = get_provider_from_provider_type(provider_type)
model = embedding_provider.canonical_model_name(model)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)

record = backend.get_registered_model(model_name=model)
if record is None:
Expand Down Expand Up @@ -196,7 +199,8 @@ def delete_model(
abort=True,
)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
record = backend.get_registered_model(model_name=model)
if record is None:
typer.echo(
Expand Down Expand Up @@ -254,7 +258,8 @@ def export_bundle_cmd(
embedding_provider = get_provider_from_provider_type(provider_type)
model = embedding_provider.canonical_model_name(model)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)

meta, h5_path = embedding_bundle.export_bundle(
backend=backend,
Expand Down Expand Up @@ -339,7 +344,8 @@ def build_faiss_cache(
embedding_provider = get_provider_from_provider_type(provider_type)
model = embedding_provider.canonical_model_name(model)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)

index_config = index_config_from_index_type(
index_type,
Expand Down Expand Up @@ -407,7 +413,8 @@ def check_faiss_cache(
embedding_provider = get_provider_from_provider_type(provider_type)
model = embedding_provider.canonical_model_name(model)

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
record = backend.get_registered_model(model_name=model)
if record is None:
typer.echo(
Expand Down Expand Up @@ -475,7 +482,8 @@ def import_bundle_cmd(
] = False,
):

backend = resolve_backend()
cfg = load_omop_emb_config()
backend = resolve_backend_from_config(cfg)
try:
imported = embedding_bundle.import_bundle(
backend=backend,
Expand Down
Loading
Loading