diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c90e472..a516006 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,8 +6,10 @@ on: jobs: label-gate: uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/label-gate.yml@main - build-test: - uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test-postgres.yml@main + build-test-sqlite: + uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test.yml@main + build-test-postgres: + uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test-postgres-v2.yml@main with: postgres-image: pgvector/pgvector:pg16 postgres-user: test @@ -15,12 +17,12 @@ jobs: postgres-db: postgres setup-commands: | uv run omop-config configure omop_emb \ - --set test_emb_db.kind=generic \ - --set test_emb_db.connection.dialect=postgresql+psycopg \ - --set test_emb_db.connection.host=localhost \ - --set test_emb_db.connection.port=5432 \ - --set test_emb_db.connection.user=test \ - --set test_emb_db.connection.password=test \ - --set test_emb_db.connection.database_name=test_omop_emb \ - --set test_emb_db.connection.test_only=true \ - --set test_emb_db.schema_name=public + --set test_emb_db_pg.kind=generic \ + --set test_emb_db_pg.connection.dialect=postgresql+psycopg \ + --set test_emb_db_pg.connection.host=localhost \ + --set test_emb_db_pg.connection.port=5432 \ + --set test_emb_db_pg.connection.user=test \ + --set test_emb_db_pg.connection.password=test \ + --set test_emb_db_pg.connection.database_name=test_omop_emb \ + --set test_emb_db_pg.connection.test_only=true \ + --set test_emb_db_pg.schema_name=public diff --git a/Dockerfile b/Dockerfile deleted file mode 100644 index b98b272..0000000 --- a/Dockerfile +++ /dev/null @@ -1,3 +0,0 @@ -FROM python:3.12-slim -RUN pip install --no-cache-dir ".[pgvector,faiss-cpu]" -WORKDIR /workspace diff --git a/docs/usage/backend-selection.md b/docs/usage/backend-selection.md index 4ea5e38..b080e1c 100644 --- a/docs/usage/backend-selection.md +++ b/docs/usage/backend-selection.md @@ -10,7 +10,7 @@ | **pgvector** | `pgvector` | `omop-emb[pgvector]` + PostgreSQL | Scales to large corpora. HNSW indexing and `halfvec` storage. | !!! note "FAISS is a sidecar, not a backend" - FAISS (`omop-emb[faiss-cpu]`) is a read-acceleration layer that sits on top of sqlite-vec or pgvector. It is not a primary backend and has no `backend_type` of its own. See the [CLI reference](cli.md#faiss-sidecar) for how to export and use FAISS indices. + FAISS (`omop-emb[faiss-cpu]`) is a read-acceleration layer that sits on top of sqlite-vec or pgvector. It is not a primary backend and has no `backend_type` of its own. See the [CLI reference](cli.md#build-faiss-cache) for how to export and use FAISS indices. ## Selecting a backend diff --git a/docs/usage/installation.md b/docs/usage/installation.md index dc18fdf..e6d2c0e 100644 --- a/docs/usage/installation.md +++ b/docs/usage/installation.md @@ -42,7 +42,7 @@ The backend (sqlite-vec or pgvector) is selected by a `[vector_stores.*]` entry' ```bash omop-config connections add emb --dialect postgresql+psycopg --host localhost --database-name omop_emb -omop-config databases add emb_db --kind generic --connection emb +omop-config databases add generic emb_db --connection emb omop-config vector-stores add vector_store --backend-type pgvector --database emb_db omop-config configure omop_emb --vector-store-name vector_store ``` diff --git a/docs/usage/interface-guide.md b/docs/usage/interface-guide.md index 5846e2d..6f6771d 100644 --- a/docs/usage/interface-guide.md +++ b/docs/usage/interface-guide.md @@ -28,14 +28,19 @@ backend = resolve_backend_from_resolved_vector_store(resolved) Or construct one directly: ```python -from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend +from sqlalchemy import create_engine +from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine from omop_emb.backends.pgvector import PGVectorEmbeddingBackend # sqlite-vec -backend = SQLiteVecEmbeddingBackend.from_path(db_path="/data/omop_emb.db") +backend = SQLiteVecEmbeddingBackend( + emb_engine=create_sqlitevec_engine(create_engine("sqlite:///data/omop_emb.db")) +) # pgvector -backend = PGVectorEmbeddingBackend.from_db_url(db_url="postgresql+psycopg://user:pass@host:5432/db") +backend = PGVectorEmbeddingBackend( + emb_engine=create_engine("postgresql+psycopg://user:pass@host:5432/db") +) ``` --- @@ -186,9 +191,14 @@ results = reader.get_nearest_concepts(query_embedding=joint_vec[None, :], k=10) `get_nearest_concepts_from_query_texts` takes a `ModelBackend` directly: build one with `omop_llm.build_model_backend` (the reader has no default backend of its own to embed with): ```python -from omop_llm import build_model_backend +from omop_llm import Capabilities, build_model_backend -model_backend = build_model_backend("ollama", "nomic-embed-text:v1.5", base_url="http://localhost:11434") +model_backend = build_model_backend( + "ollama", + "nomic-embed-text:v1.5", + model_capabilities=Capabilities(embeddings=True), + base_url="http://localhost:11434", +) results = reader.get_nearest_concepts_from_query_texts( query_texts=("high blood pressure", "type 2 diabetes"), @@ -267,6 +277,7 @@ from omop_llm import build_model_backend model_backend = build_model_backend( "ollama", "nomic-embed-text:v1.5", + model_capabilities=Capabilities(embeddings=True), base_url="http://host.docker.internal:11434", ) @@ -279,6 +290,7 @@ print(model_backend.dimensions()) # auto-discovered via Ollama /api/show model_backend = build_model_backend( "openai", "text-embedding-3-large", + model_capabilities=Capabilities(embeddings=True), base_url="https://api.openai.com/v1", api_key="sk-...", ) diff --git a/pyproject.toml b/pyproject.toml index 6018c40..c118927 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -102,7 +102,7 @@ markers = [ "faiss: FAISS sidecar tests", ] testpaths = ["tests"] -addopts = "-v --tb=short" +addopts = "-v --tb=short -m 'not db_dialect'" filterwarnings = [ "ignore:datetime\\.datetime\\.utcfromtimestamp\\(\\) is deprecated:DeprecationWarning:dateutil.tz.tz", "ignore:builtin type .* has no __module__ attribute:DeprecationWarning", diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index 830d2ab..f355577 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -6,7 +6,12 @@ from datetime import datetime from typing import Any, Callable, Generic, Iterable, Mapping, Optional, Sequence, Tuple, TypeVar, Union from numpy import ndarray -from oa_configurator import ResolvedDatabase, ResolvedVectorStore +from oa_configurator import ( + SCHEMA_TRANSLATE_MAP_KEY, + Dialect, + ResolvedDatabase, + ResolvedVectorStore, +) from sqlalchemy import Engine from sqlalchemy.engine import make_url from sqlalchemy.orm import sessionmaker @@ -22,7 +27,12 @@ from omop_emb.backends.embedding_table import ConceptEmbeddingRecord from omop_emb.backends.index_config import IndexConfig, FlatIndexConfig -from omop_emb.model_registry import EmbeddingModelRecord, RegistryManager +from omop_emb.model_registry import ( + EmbeddingModelRecord, + REGISTRY_SCHEMA_KEY, + RegistryManager, + resolve_registry_physical_schema, +) from omop_emb.utils.embedding_utils import ( EmbeddingConceptFilter, NearestConceptMatch, @@ -168,7 +178,12 @@ class EmbeddingBackend(ABC, Generic[TEmbeddingTable]): DEFAULT_K_NEAREST = 10 - def __init__(self, emb_engine: Engine) -> None: + def __init__( + self, + emb_engine: Engine, + *, + resolved: ResolvedDatabase | None = None, + ) -> None: actual_dialect = emb_engine.dialect.name if actual_dialect != self.dialect: raise ValueError( @@ -176,7 +191,8 @@ def __init__(self, emb_engine: Engine) -> None: f"got '{actual_dialect}'." ) super().__init__() - self._registry = RegistryManager(emb_engine) + self._resolved = resolved + self._registry = RegistryManager(emb_engine, resolved=resolved) self._table_cache: dict[str, TEmbeddingTable] = {} self._initialise_store() @@ -278,7 +294,9 @@ def register_model( ModelRegistrationConflictError If the model is already registered with a different dimensionality. ValueError - If ``metadata`` contains a reserved key. + If ``metadata`` contains a reserved key, or if ``index_config`` is + not ``FlatIndexConfig()`` (non-FLAT indexes may only be built + after registration, not at registration time). """ if index_config is None: @@ -1036,20 +1054,29 @@ def resolve_backend( dialect = make_url(database.connection.url).get_backend_name() + # The model registry lives in its own reserved schema (MODEL_REGISTRY_SCHEMA), + # independent of database's own schema. create_engine() merges this extra + # key onto its own configured map rather than replacing it. + registry_schema = resolve_registry_physical_schema(database.connection.dialect_name) + registry_schema_translate_map = {REGISTRY_SCHEMA_KEY: registry_schema} + if resolved_backend == BackendType.SQLITEVEC: - from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend + from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine - if dialect != "sqlite": + if dialect != Dialect.SQLITE: raise RuntimeError( f"sqlitevec backend requires a sqlite-dialect database, got dialect: {dialect!r}." ) - db_path = make_url(database.connection.url).database - assert db_path is not None, "ConnectionConfig.build_url() always sets a database segment for sqlite" - logger.info(f"Using SQLiteVec backend with database file: {db_path}") - return SQLiteVecEmbeddingBackend.from_path(db_path) + emb_engine = create_sqlitevec_engine( + database.create_engine( + execution_options={SCHEMA_TRANSLATE_MAP_KEY: registry_schema_translate_map} + ) + ) + logger.info(f"Using SQLiteVec backend with engine: {emb_engine.url}") + return SQLiteVecEmbeddingBackend(emb_engine=emb_engine, resolved=database) if resolved_backend == BackendType.PGVECTOR: - if dialect != "postgresql": + if dialect != Dialect.POSTGRESQL: raise RuntimeError( "The resolved URL must point to a PostgreSQL database " f"(pgvector extension required), got dialect: {dialect!r}." @@ -1061,9 +1088,11 @@ def resolve_backend( "pgvector backend is not installed. " "Install it with: pip install omop-emb[pgvector]" ) from exc - emb_engine = database.create_engine() + emb_engine = database.create_engine( + execution_options={SCHEMA_TRANSLATE_MAP_KEY: registry_schema_translate_map} + ) logger.info(f"Using pgvector backend with engine: {emb_engine.url}") - return PGVectorEmbeddingBackend(emb_engine=emb_engine) + return PGVectorEmbeddingBackend(emb_engine=emb_engine, resolved=database) raise RuntimeError(f"Implementation for {resolved_backend.value} is not available.") diff --git a/src/omop_emb/backends/db_utils.py b/src/omop_emb/backends/db_utils.py index ceb8b13..8b561cf 100644 --- a/src/omop_emb/backends/db_utils.py +++ b/src/omop_emb/backends/db_utils.py @@ -3,6 +3,7 @@ from contextlib import contextmanager from typing import Any, Iterator, Sequence +from oa_configurator import Dialect from sqlalchemy import Column, Select, text from sqlalchemy.orm import Session from sqlalchemy.sql.base import ColumnCollection @@ -100,13 +101,13 @@ def temp_filter_table( connections safely. The table is truncated before each use and cleaned up when the connection is returned to the pool. """ - if dialect == "postgresql": + if dialect == Dialect.POSTGRESQL: session.execute( text( f'CREATE TEMPORARY TABLE "{table_name}" (id {col_type}) ON COMMIT DROP' ) ) - elif dialect == "sqlite": + elif dialect == Dialect.SQLITE: session.execute( text(f'CREATE TEMPORARY TABLE IF NOT EXISTS "{table_name}" (id {col_type})') ) diff --git a/src/omop_emb/backends/index_config.py b/src/omop_emb/backends/index_config.py index 4ffeab2..c96f4e7 100644 --- a/src/omop_emb/backends/index_config.py +++ b/src/omop_emb/backends/index_config.py @@ -126,9 +126,7 @@ def from_dict(cls, config_dict: dict[str, Any]) -> Self: Notes ----- Use this method to reconstruct an ``IndexConfig`` from the ORM - ``index_config`` JSON column. It is distinct from - :meth:`from_metadata`, which reads from a metadata dict that wraps the - config under ``"index_config"`` key. + ``index_config`` JSON column. """ if not is_dataclass(cls): raise TypeError(f"Must be called on a dataclass, not {cls.__name__}.") diff --git a/src/omop_emb/backends/pgvector/pg_backend.py b/src/omop_emb/backends/pgvector/pg_backend.py index 876fd49..08acc99 100644 --- a/src/omop_emb/backends/pgvector/pg_backend.py +++ b/src/omop_emb/backends/pgvector/pg_backend.py @@ -5,7 +5,8 @@ from typing import Mapping, Optional, Sequence, Tuple from numpy import ndarray -from sqlalchemy import Engine, select, text, create_engine +from oa_configurator import Dialect, ResolvedDatabase +from sqlalchemy import Engine, select, text try: from pgvector.sqlalchemy import Vector # noqa: F401 @@ -78,25 +79,14 @@ class PGVectorEmbeddingBackend(EmbeddingBackend[type[PGEmbeddingTable]]): ``__init__``. """ - def __init__(self, emb_engine: Engine) -> None: + def __init__( + self, + emb_engine: Engine, + *, + resolved: ResolvedDatabase | None = None, + ) -> None: self._index_managers: dict[str, PGVectorBaseIndexManager] = {} - super().__init__(emb_engine=emb_engine) - - @classmethod - def from_db_url(cls, db_url: str) -> PGVectorEmbeddingBackend: - """Create a pgvector embedding backend from a database URL. - - Parameters - ---------- - db_url : str - Database URL in SQLAlchemy format, e.g. ``postgresql://user:pass@host:port/dbname``. - - Returns - ------- - PGVectorEmbeddingBackend - """ - engine = create_engine(db_url, echo=False) - return cls(emb_engine=engine) + super().__init__(emb_engine=emb_engine, resolved=resolved) # ------------------------------------------------------------------ # Backend identity @@ -108,7 +98,7 @@ def backend_type(self) -> BackendType: @property def dialect(self) -> str: - return "postgresql" + return Dialect.POSTGRESQL # ------------------------------------------------------------------ # Store lifecycle @@ -135,7 +125,9 @@ def _create_storage_table( self, model_record: EmbeddingModelRecord ) -> type[PGEmbeddingTable]: return create_pg_embedding_table( - engine=self.emb_engine, model_record=model_record + engine=self.emb_engine, + model_record=model_record, + resolved=self._resolved, ) def _delete_storage_table(self, model_record: EmbeddingModelRecord) -> None: @@ -181,7 +173,11 @@ def register_model( Raises ------ ValueError - If ``dimensions`` exceeds the pgvector halfvec limit of 4 000. + If ``dimensions`` exceeds the pgvector halfvec limit of 4 000, or + for any reason :meth:`EmbeddingBackend.register_model` raises + (a reserved metadata key, or a non-``FlatIndexConfig`` index). + ModelRegistrationConflictError + If the model is already registered with a different dimensionality. """ vector_column_type_for_dimensions(dimensions) # validates halfvec limit return super().register_model( diff --git a/src/omop_emb/backends/pgvector/pg_index_manager.py b/src/omop_emb/backends/pgvector/pg_index_manager.py index a3d57c1..4f73455 100644 --- a/src/omop_emb/backends/pgvector/pg_index_manager.py +++ b/src/omop_emb/backends/pgvector/pg_index_manager.py @@ -17,6 +17,7 @@ import logging from typing import Generic, TypeVar +from oa_configurator import qualified, physical_schema_of from sqlalchemy import Engine, inspect, text from omop_emb.config import IndexType, MetricType, VectorColumnType @@ -77,7 +78,10 @@ def index_config(self) -> C: def has_index(self, metric_type: MetricType) -> bool: with self._engine.connect() as conn: existing = { - idx["name"] for idx in inspect(conn).get_indexes(self._tablename) + idx["name"] + for idx in inspect(conn).get_indexes( + self._tablename, schema=physical_schema_of(conn) + ) } return self._index_name(metric_type) in existing @@ -101,7 +105,9 @@ def drop_index(self, metric_type: MetricType) -> None: name = self._index_name(metric_type) existed = self.has_index(metric_type) with self._engine.begin() as conn: - conn.execute(text(f"DROP INDEX IF EXISTS {name}")) + conn.execute( + text(f"DROP INDEX IF EXISTS {qualified(conn, name, physical_schema=physical_schema_of(conn))}") + ) if existed: logger.info(f"Dropped pgvector index '{name}'.") @@ -190,9 +196,10 @@ def supported_index_type(self) -> IndexType: def _create_index_ddl(self, metric_type: MetricType) -> str: ops = self._ops_for_metric(metric_type) cfg = self.index_config + table_ref = qualified(self._engine, self._tablename, physical_schema=physical_schema_of(self._engine)) return ( f"CREATE INDEX {self._index_name(metric_type)} " - f"ON {self._tablename} " + f"ON {table_ref} " f"USING hnsw ({self._embedding_column} {ops}) " f"WITH (m = {cfg.num_neighbors}, ef_construction = {cfg.ef_construction})" ) diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index f639211..5131b83 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -15,6 +15,15 @@ from typing import List, Optional, Sequence, Union from numpy import ndarray +from oa_configurator import ( + Dialect, + ResolvedDatabase, + Role, + ensure_schema, + guard_schema_provenance_for, + qualified, + physical_schema_of, +) from sqlalchemy import Engine, Integer, Row, Select, func, inspect as sa_inspect, literal, select, text, TextClause from sqlalchemy.sql import cast, column, values from sqlalchemy.sql.elements import ColumnElement @@ -31,11 +40,14 @@ def table_exists(engine: Engine, table_name: str) -> bool: - """Return ``True`` if *table_name* exists in the current Postgres schema.""" - return sa_inspect(engine).has_table(table_name) + """Return ``True`` if *table_name* exists in the engine's configured schema.""" + return sa_inspect(engine).has_table(table_name, schema=physical_schema_of(engine)) def create_pg_embedding_table( - engine: Engine, model_record: EmbeddingModelRecord + engine: Engine, + model_record: EmbeddingModelRecord, + *, + resolved: ResolvedDatabase | None = None, ) -> type[PGEmbeddingTable]: """Create a pgvector embedding table and return its ORM class. @@ -44,6 +56,10 @@ def create_pg_embedding_table( engine : Engine SQLAlchemy engine for the pgvector database. model_record : EmbeddingModelRecord + resolved : ResolvedDatabase, optional + Enables the schema-provenance guard around the ``create_all()`` + call. Omitted by callers with no resolved config behind their + engine, in which case the guard no-ops. Returns ------- @@ -57,7 +73,15 @@ def create_pg_embedding_table( base class; this function always issues DDL. """ table_cls = pg_embedding_table_descriptor(model_record) - EmbeddingTableBase.metadata.create_all(engine, tables=[table_cls.__table__]) # ty: ignore[invalid-argument-type] + with engine.begin() as connection: + # Ensure that the schema exists before creating the table + physical_schema = physical_schema_of(connection, schema_tag=Role.PRIMARY) + ensure_schema(connection, physical_schema) + guard = guard_schema_provenance_for( + connection, resolved, schema_tag=Role.PRIMARY, tables=[table_cls.__table__] # ty: ignore[invalid-argument-type] + ) + with guard: + EmbeddingTableBase.metadata.create_all(connection, tables=[table_cls.__table__]) # ty: ignore[invalid-argument-type] return table_cls @@ -71,7 +95,9 @@ def drop_pg_embedding_table(engine: Engine, model_record: EmbeddingModelRecord) """ tablename = model_record.storage_identifier with engine.begin() as conn: - conn.execute(text(f'DROP TABLE IF EXISTS "{tablename}"')) + conn.execute( + text(f"DROP TABLE IF EXISTS {qualified(conn, tablename, physical_schema=physical_schema_of(conn))}") + ) logger.info(f"Dropped embedding table '{tablename}'.") @@ -175,7 +201,7 @@ def query_nearest_concept_ids( metric_type: MetricType, k: int, concept_filter: Optional[EmbeddingConceptFilter] = None, - dialect: str = "postgresql", + dialect: str = Dialect.POSTGRESQL, ) -> Sequence[Row]: """Run a pgvector ANN query returning the nearest concept IDs per query. @@ -263,7 +289,7 @@ def query_concept_ids_matching_filter( session: Session, embedding_table: type[PGEmbeddingTable], concept_filter: EmbeddingConceptFilter, - dialect: str = "postgresql", + dialect: str = Dialect.POSTGRESQL, ) -> set[int]: """Return every ``concept_id`` satisfying *concept_filter*.""" setup_concept_filter_temps(session, concept_filter, dialect) @@ -277,7 +303,7 @@ def query_concept_filter_metadata( session: Session, embedding_table: type[PGEmbeddingTable], concept_filter: EmbeddingConceptFilter, - dialect: str = "postgresql", + dialect: str = Dialect.POSTGRESQL, ) -> Sequence[Row]: """Return filter metadata columns (raw rows) for every concept ID satisfying concept_filter. @@ -373,7 +399,7 @@ def pg_embedding_table_descriptor(model_record: EmbeddingModelRecord) -> type[PG (PGEmbeddingTable,), { "__tablename__": tablename, - "__table_args__": {"extend_existing": True}, + "__table_args__": {"schema": Role.PRIMARY.value, "extend_existing": True}, "__module__": __name__, EMBEDDING_COLUMN_NAME: emb_col, }, diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index 3e1cc9d..dd431d4 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -12,7 +12,12 @@ from collections.abc import Iterator from dataclasses import dataclass -from oa_configurator import ResolvedVectorStore +from oa_configurator import ( + Dialect, + ResolvedVectorStore, + qualified, + schema_if_supported +) from sqlalchemy import Engine, event, inspect, select from omop_emb.backends.base_backend import ( @@ -22,6 +27,7 @@ from omop_emb.backends.embedding_table import concept_metadata_table_descriptor from omop_emb.config import BackendType, parse_backend_type from omop_emb.model_registry import EmbeddingModelRecord, RegistryManager +from omop_emb.utils.cdm import streamed @dataclass(frozen=True) @@ -43,11 +49,11 @@ def __init__( engine: Engine, *, backend_type: str | BackendType, - schema: str | None, + physical_schema: str | None, ) -> None: self._engine = engine self.backend_type = parse_backend_type(backend_type) - self.schema = schema + self.physical_schema = physical_schema self._registry = RegistryManager.read_only(engine) def __enter__(self) -> ReadOnlyEmbeddingStore: @@ -92,24 +98,23 @@ def iter_stored_embeddings( if record is None: return inspector = inspect(self._engine) - if not inspector.has_table(record.storage_identifier, schema=self.schema): + if not inspector.has_table(record.storage_identifier, schema=self.physical_schema): return - schema = ( - None - if self._engine.dialect.name == "sqlite" and self.schema == "main" - else self.schema - ) + physical_schema = schema_if_supported(self.physical_schema, self._engine) table = concept_metadata_table_descriptor( record.storage_identifier, - schema=schema, + schema=physical_schema, + ) + statement = streamed( + select( + table.c.concept_id, + table.c.domain_id, + table.c.vocabulary_id, + table.c.is_standard, + table.c.is_valid, + ), + batch_size, ) - statement = select( - table.c.concept_id, - table.c.domain_id, - table.c.vocabulary_id, - table.c.is_standard, - table.c.is_valid, - ).execution_options(stream_results=True, yield_per=batch_size) with self._engine.connect() as connection: rows = connection.execute(statement).mappings() for row in rows: @@ -124,7 +129,7 @@ def iter_stored_embeddings( def physical_indexes(self, model_name: str) -> tuple[str, ...]: """Return existing PostgreSQL indexes without creating or changing them.""" - if self._engine.dialect.name != "postgresql": + if self._engine.dialect.name != Dialect.POSTGRESQL: return () record = self.model(model_name) if record is None: @@ -133,8 +138,7 @@ def physical_indexes(self, model_name: str) -> tuple[str, ...]: return tuple( str(item["name"]) for item in inspect(self._engine).get_indexes( - record.storage_identifier, - schema=self.schema, + record.storage_identifier, schema=self.physical_schema, ) if str(item["name"]).startswith(expected_prefix) ) @@ -142,10 +146,8 @@ def physical_indexes(self, model_name: str) -> tuple[str, ...]: def drop_index_sql(self, model_name: str) -> tuple[str, ...]: """Return reviewed index-removal statements without executing them.""" - quote = self._engine.dialect.identifier_preparer.quote - schema_prefix = f"{quote(self.schema)}." if self.schema else "" return tuple( - f"DROP INDEX IF EXISTS {schema_prefix}{quote(name)};" + f"DROP INDEX IF EXISTS {qualified(self._engine, name, physical_schema=self.physical_schema)};" for name in self.physical_indexes(model_name) ) @@ -166,7 +168,7 @@ def inspect_resolved_vector_store( """ engine = resolved.database.create_engine() - if resolved.backend_type == "sqlitevec": + if resolved.backend_type == BackendType.SQLITEVEC: try: import sqlite_vec except ImportError as exc: # pragma: no cover - optional extra @@ -185,7 +187,7 @@ def _load_sqlite_vec(dbapi_connection, _connection_record): return ReadOnlyEmbeddingStore( engine, backend_type=resolved.backend_type, - schema=resolved.database.schema_name, + physical_schema=resolved.database.schema_name, ) except Exception: engine.dispose() diff --git a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py index eedbae7..6d92d3c 100644 --- a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py +++ b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py @@ -11,7 +11,8 @@ import numpy as np from numpy import ndarray -from sqlalchemy import Engine, MetaData, Table, create_engine, event, text +from oa_configurator import Dialect, ResolvedDatabase +from sqlalchemy import Engine, MetaData, Table, event, text try: import sqlite_vec @@ -50,20 +51,22 @@ logger = logging.getLogger(__name__) -def create_sqlitevec_engine(db_path: str) -> Engine: - """Create a SQLAlchemy engine for sqlite-vec with the extension pre-loaded. +def create_sqlitevec_engine(engine: Engine) -> Engine: + """Attach the sqlite-vec extension-loading connect listener to *engine*. Parameters ---------- - db_path : str - File path to the SQLite database, or ``':memory:'`` for an in-memory - database. + engine : Engine + An already-built SQLite engine (e.g. ``database.create_engine()`` + from an oa-configurator ``ResolvedDatabase``, so it carries whatever + ``schema_translate_map``/pool settings the resolver configured, + rather than a bare path reconstructed from its URL). Returns ------- Engine + The same *engine*, with the extension listener attached. """ - engine = create_engine(f"sqlite:///{db_path}", echo=False) @event.listens_for(engine, "connect") def _load_sqlite_vec(dbapi_conn, _connection_record): @@ -90,24 +93,14 @@ class SQLiteVecEmbeddingBackend(EmbeddingBackend[Table]): ``CREATE VIRTUAL TABLE``. """ - def __init__(self, emb_engine: Engine) -> None: + def __init__( + self, + emb_engine: Engine, + *, + resolved: ResolvedDatabase | None = None, + ) -> None: self._sqlite_vec_metadata = MetaData() - super().__init__(emb_engine=emb_engine) - - @classmethod - def from_path(cls, db_path: str) -> "SQLiteVecEmbeddingBackend": - """Construct a backend from a database file path. - - Parameters - ---------- - db_path : str - File path to the SQLite database, or ``':memory:'`` for testing. - - Returns - ------- - SQLiteVecEmbeddingBackend - """ - return cls(emb_engine=create_sqlitevec_engine(db_path)) + super().__init__(emb_engine=emb_engine, resolved=resolved) # ------------------------------------------------------------------ # Backend identity @@ -119,7 +112,7 @@ def backend_type(self) -> BackendType: @property def dialect(self) -> str: - return "sqlite" + return Dialect.SQLITE # ------------------------------------------------------------------ # Storage table management diff --git a/src/omop_emb/backends/sqlitevec/sqlitevec_sql.py b/src/omop_emb/backends/sqlitevec/sqlitevec_sql.py index f493828..de8743f 100644 --- a/src/omop_emb/backends/sqlitevec/sqlitevec_sql.py +++ b/src/omop_emb/backends/sqlitevec/sqlitevec_sql.py @@ -7,6 +7,7 @@ import numpy as np from numpy import ndarray +from oa_configurator import Dialect from sqlalchemy import ( Column, Engine, @@ -158,7 +159,7 @@ def dml_upsert_rows( table: Table, records: Sequence[ConceptEmbeddingRecord], embeddings: ndarray, - dialect: str = "sqlite", + dialect: str = Dialect.SQLITE, ) -> None: """Upsert embedding rows into a vec0 table. @@ -249,7 +250,7 @@ def query_knn_batch( metric_type: MetricType, k: int, concept_filter: Optional[EmbeddingConceptFilter] = None, - dialect: str = "sqlite", + dialect: str = Dialect.SQLITE, ) -> list[Sequence[Row]]: """Run KNN queries against a vec0 table, one per vector in *query_vectors*. @@ -303,7 +304,7 @@ def query_concept_ids_matching_filter( session: Session, table: Table, concept_filter: EmbeddingConceptFilter, - dialect: str = "sqlite", + dialect: str = Dialect.SQLITE, ) -> set[int]: """Return every ``concept_id`` in `table` satisfying `concept_filter`. Used to build an exact FAISS pre-filter set, not for ranking. @@ -356,7 +357,7 @@ def query_embeddings_by_ids( session: Session, table: Table, concept_ids: Sequence[int], - dialect: str = "sqlite", + dialect: str = Dialect.SQLITE, ) -> dict[int, list[float]]: """Fetch embedding vectors for a set of concept IDs. @@ -419,7 +420,7 @@ def query_concept_filter_metadata( session: Session, table: Table, concept_filter: EmbeddingConceptFilter, - dialect: str = "sqlite", + dialect: str = Dialect.SQLITE, ) -> Sequence[Row]: """Return filter metadata columns (raw rows) for every concept ID satisfying concept_filter. diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index cade3d3..170435b 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -16,8 +16,17 @@ Resolver, ResolvedVectorStore, VectorStoreConfig, + register_reserved_schema, + register_reserved_schema_tag, ) +# Guaranteed to be imported and registered if there is a config +MODEL_REGISTRY_SCHEMA: str = "omop_emb_registry" +REGISTRY_SCHEMA_KEY: str = "registry" + +register_reserved_schema(MODEL_REGISTRY_SCHEMA, owner="omop_emb") +register_reserved_schema_tag(REGISTRY_SCHEMA_KEY, owner="omop_emb") + class OmopEmbConfig(PackageConfigBase): """oa-configurator config class for omop-emb. @@ -37,7 +46,18 @@ class OmopEmbConfig(PackageConfigBase): extra_logging_namespaces: ClassVar[tuple[str, ...]] = ("orm_loader", "omop_alchemy") cdm_db: Annotated[str, RefTo(CDMDatabaseConfig)] = "cdm_db" - test_emb_db: Annotated[str | None, RefTo(GenericDatabaseConfig, is_test=True)] = None + test_emb_db_pg: Annotated[str | None, RefTo(GenericDatabaseConfig, is_test=True)] = Field( + default=None, + description="Real PostgreSQL test database, for Postgres-only integration testing.", + ) + test_emb_db_sqlite: Annotated[str | None, RefTo(GenericDatabaseConfig, is_test=True)] = Field( + default=None, + description=( + "Disposable SQLite test database; left unconfigured by design " + "since isolated_test_database(..., dialect='sqlite') provisions " + "one without needing a config entry." + ), + ) embedding_model_name: Annotated[str, RefTo(ModelConfig)] = Field( default="embedding-model", description=( @@ -254,6 +274,31 @@ def parse_metric_type(value: str | MetricType) -> MetricType: } +def _supported_indices(backend: BackendType) -> Dict[IndexType, Tuple[MetricType, ...]]: + """The index/metric support map for *backend*. + + Parameters + ---------- + backend : BackendType + + Returns + ------- + Dict[IndexType, Tuple[MetricType, ...]] + + Raises + ------ + ValueError + If *backend* isn't registered in SUPPORTED_INDICES_AND_METRICS_PER_BACKEND. + """ + try: + return SUPPORTED_INDICES_AND_METRICS_PER_BACKEND[backend] + except KeyError: + raise ValueError( + f"Unsupported backend {backend!r}. Supported: " + f"{sorted(b.value for b in SUPPORTED_INDICES_AND_METRICS_PER_BACKEND)}." + ) from None + + def is_supported_index_metric_combination_for_backend( backend: BackendType, index: IndexType, metric: MetricType ) -> bool: @@ -269,8 +314,7 @@ def is_supported_index_metric_combination_for_backend( ------- bool """ - supported = SUPPORTED_INDICES_AND_METRICS_PER_BACKEND.get(backend, {}) - return metric in supported.get(index, ()) + return metric in _supported_indices(backend).get(index, ()) def is_index_type_supported_for_backend(backend: BackendType, index: IndexType) -> bool: @@ -285,7 +329,7 @@ def is_index_type_supported_for_backend(backend: BackendType, index: IndexType) ------- bool """ - return index in SUPPORTED_INDICES_AND_METRICS_PER_BACKEND.get(backend, {}) + return index in _supported_indices(backend) def get_supported_index_types_for_backend( @@ -301,7 +345,7 @@ def get_supported_index_types_for_backend( ------- tuple[IndexType, ...] """ - return tuple(SUPPORTED_INDICES_AND_METRICS_PER_BACKEND.get(backend, {}).keys()) + return tuple(_supported_indices(backend).keys()) def get_supported_metrics_for_backend(backend: BackendType) -> Tuple[MetricType, ...]: @@ -316,6 +360,6 @@ def get_supported_metrics_for_backend(backend: BackendType) -> Tuple[MetricType, tuple[MetricType, ...] """ seen: set[MetricType] = set() - for metrics in SUPPORTED_INDICES_AND_METRICS_PER_BACKEND.get(backend, {}).values(): + for metrics in _supported_indices(backend).values(): seen.update(metrics) return tuple(seen) diff --git a/src/omop_emb/interface.py b/src/omop_emb/interface.py index 1fae4d1..fea5ecd 100644 --- a/src/omop_emb/interface.py +++ b/src/omop_emb/interface.py @@ -115,7 +115,8 @@ class EmbeddingReaderInterface: ``is_standard``, and ``is_active`` are populated directly from the embedding table by the backend. model : str - Model name in canonical form. + Model name that is expected to be canonicalized by the constructor. + A possibly-canonical name is accepted as we cannot know if canonical or not. provider_type : str, optional omop-llm provider key. Defaults to ``'ollama'``. k : int diff --git a/src/omop_emb/model_registry/__init__.py b/src/omop_emb/model_registry/__init__.py index 276eb84..34b37f3 100644 --- a/src/omop_emb/model_registry/__init__.py +++ b/src/omop_emb/model_registry/__init__.py @@ -1,13 +1,17 @@ from omop_emb.model_registry.model_registry_types import EmbeddingModelRecord from omop_emb.model_registry.model_registry_manager import RegistryManager from omop_emb.model_registry.model_registry_orm import ( + REGISTRY_SCHEMA_KEY, ModelRegistry, - ensure_registry_schema, + ensure_registry_table, + resolve_registry_physical_schema, ) __all__ = [ "EmbeddingModelRecord", "RegistryManager", "ModelRegistry", - "ensure_registry_schema", + "ensure_registry_table", + "REGISTRY_SCHEMA_KEY", + "resolve_registry_physical_schema", ] diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index eba015d..9919e99 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -5,6 +5,7 @@ from datetime import datetime, timezone from typing import Mapping, Optional +from oa_configurator import SCHEMA_TRANSLATE_MAP_KEY, ResolvedDatabase from sqlalchemy import Engine, inspect, select, update from sqlalchemy.orm import Session, sessionmaker @@ -14,8 +15,10 @@ index_config_from_orm_row, ) from omop_emb.model_registry.model_registry_orm import ( + REGISTRY_SCHEMA_KEY, ModelRegistry, - ensure_registry_schema, + ensure_registry_table, + resolve_registry_physical_schema, ) from omop_emb.model_registry.model_registry_types import EmbeddingModelRecord from omop_emb.utils.errors import ModelRegistrationConflictError @@ -33,17 +36,41 @@ class RegistryManager: ---------- embedding_engine : Engine SQLAlchemy embedding engine connected to the embedding store. + + Notes + ----- + The registry table is created under the ``registry`` schema tag (a + ``schema_translate_map`` key, resolved to the physical schema named by + ``MODEL_REGISTRY_SCHEMA``), independent of whichever schema the embedding + store itself resolves to. """ - def __init__(self, embedding_engine: Engine, *, initialize: bool = True) -> None: - self._embedding_engine = embedding_engine + def __init__( + self, + embedding_engine: Engine, + *, + initialize: bool = True, + resolved: ResolvedDatabase | None = None, + ) -> None: + # A caller building through base_backend.py's own factory already injected + # REGISTRY_SCHEMA_KEY at engine-construction time + existing_schema_translate_map = embedding_engine.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {} + if REGISTRY_SCHEMA_KEY in existing_schema_translate_map: + self._embedding_engine = embedding_engine + else: + self._embedding_engine = embedding_engine.execution_options( + schema_translate_map={ + **existing_schema_translate_map, + REGISTRY_SCHEMA_KEY: resolve_registry_physical_schema(embedding_engine), + } + ) self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize - self._registry_available = inspect(embedding_engine).has_table( - ModelRegistry.__tablename__ + self._registry_available = inspect(self._embedding_engine).has_table( + ModelRegistry.__tablename__, schema=resolve_registry_physical_schema(self._embedding_engine) ) if initialize: - ensure_registry_schema(embedding_engine) + ensure_registry_table(self._embedding_engine, resolved=resolved) self._registry_available = True @classmethod @@ -148,6 +175,8 @@ def register_model( ------ ModelRegistrationConflictError If the model is already registered with a different configuration. + ValueError + If ``metadata`` contains a reserved key. """ self._require_writable() _validate_metadata_keys(metadata) diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index 8cd72d4..8ae9a43 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -1,7 +1,18 @@ from __future__ import annotations +import warnings +from contextlib import nullcontext from typing import Any, Optional +from oa_configurator import ( + SCHEMA_TRANSLATE_MAP_KEY, + Dialect, + ResolvedDatabase, + ensure_schema, + guard_schema_provenance, + qualified, + supports_schemas, +) from sqlalchemy import ( DateTime, Engine, @@ -18,6 +29,8 @@ from omop_llm import supported_providers from omop_emb.config import ( + MODEL_REGISTRY_SCHEMA, + REGISTRY_SCHEMA_KEY, IndexType, MetricType, ) @@ -78,6 +91,7 @@ class ModelRegistry(ModelRegistryBase): """ __tablename__ = "model_registry" + __table_args__ = {"schema": REGISTRY_SCHEMA_KEY} model_name = mapped_column(String, primary_key=True) @@ -175,15 +189,53 @@ def _validate_and_sync_index_config( return index_config.to_dict() -def ensure_registry_schema(engine: Engine) -> None: - """Create or upgrade the model registry table. +def resolve_registry_physical_schema(bindable) -> str | None: + """MODEL_REGISTRY_SCHEMA on a dialect with real schema support, else None.""" + return MODEL_REGISTRY_SCHEMA if supports_schemas(bindable) else None + + +def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = None) -> None: + """Create or upgrade the model registry table, in its own reserved schema. + + Notes + ----- + Shared across every resolved database on the same connection; `database_name= + MODEL_REGISTRY_SCHEMA` (not `resolved.name`) keeps them tracked as one identity, + avoiding a false "already populated" on the second one. Parameters ---------- engine : Engine - SQLAlchemy engine connected to the registry database. + SQLAlchemy engine connected to the registry database. Not required + to already carry a REGISTRY_SCHEMA_KEY entry in its schema_translate_map. + resolved : ResolvedDatabase, optional + Enables the schema-provenance guard around the ``create_all()`` + call. Omitted by callers with no resolved config behind their + engine, in which case the guard no-ops. """ - ModelRegistryBase.metadata.create_all(engine, tables=[ModelRegistry.__table__]) # ty: ignore[invalid-argument-type] + with engine.begin() as connection: + connection = connection.execution_options( + schema_translate_map={ + **(connection.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), + REGISTRY_SCHEMA_KEY: resolve_registry_physical_schema(connection), + } + ) + ensure_schema(connection, MODEL_REGISTRY_SCHEMA) + registry_schema = resolve_registry_physical_schema(connection) + guard = ( + guard_schema_provenance( + connection, + database_name=MODEL_REGISTRY_SCHEMA, + test_only=resolved.connection.test_only, + schema_tag=REGISTRY_SCHEMA_KEY, + physical_schema=registry_schema, + tables=[ModelRegistry.__table__], # ty: ignore[invalid-argument-type] + ) + if resolved is not None and registry_schema is not None + else nullcontext() + ) + with guard: + ModelRegistryBase.metadata.create_all(connection, tables=[ModelRegistry.__table__]) # ty: ignore[invalid-argument-type] _migrate_legacy_provider_type_column(engine) @@ -199,7 +251,9 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: The migration is deliberately idempotent so normal backend construction can safely run it for both existing and newly-created registries. """ - columns = inspect(engine).get_columns(ModelRegistry.__tablename__) + columns = inspect(engine).get_columns( + ModelRegistry.__tablename__, schema=resolve_registry_physical_schema(engine) + ) provider_column = next( (column for column in columns if column["name"] == "provider_type"), None, @@ -209,17 +263,28 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: legacy_length = getattr(provider_column["type"], "length", None) with engine.begin() as connection: - if engine.dialect.name == "postgresql" and legacy_length is not None: + registry_schema = resolve_registry_physical_schema(connection) + if engine.dialect.name == Dialect.POSTGRESQL and legacy_length is not None: + warnings.warn( + "Widening a legacy fixed-length provider_type column. This " + "migration path is deprecated and will be removed once no " + "pre-omop-llm registry remains.", + DeprecationWarning, + stacklevel=2, + ) connection.execute( text( - "ALTER TABLE model_registry " + f"ALTER TABLE {qualified(connection, ModelRegistry.__tablename__, physical_schema=registry_schema)} " "ALTER COLUMN provider_type TYPE VARCHAR " "USING provider_type::text" ) ) + # Raw text(), not update(): update() against the full mapped table + # would pull in updated_at's onupdate=func.now() default, which the + # legacy partial table (provider_type only) doesn't have. connection.execute( text( - "UPDATE model_registry " + f"UPDATE {qualified(connection, ModelRegistry.__tablename__, physical_schema=registry_schema)} " "SET provider_type = lower(provider_type) " "WHERE provider_type IS NOT NULL " "AND provider_type <> lower(provider_type)" diff --git a/src/omop_emb/population.py b/src/omop_emb/population.py index eac0236..74a62d8 100644 --- a/src/omop_emb/population.py +++ b/src/omop_emb/population.py @@ -14,6 +14,7 @@ from omop_alchemy.cdm.query import ConceptFilter from omop_emb.backends.read_only import ReadOnlyEmbeddingStore, StoredEmbedding +from omop_emb.utils.cdm import streamed @dataclass(frozen=True) @@ -190,10 +191,9 @@ def _iter_current_concepts( Concept.is_valid_expr().label("is_valid"), ) ) - .execution_options(stream_results=True, yield_per=batch_size) ) with Session(cdm_engine) as session: - yield from session.execute(statement) + yield from session.execute(streamed(statement, batch_size)) def _iter_stored_embeddings( diff --git a/src/omop_emb/storage/embedding_bundle.py b/src/omop_emb/storage/embedding_bundle.py index bf25caa..c6cc22c 100644 --- a/src/omop_emb/storage/embedding_bundle.py +++ b/src/omop_emb/storage/embedding_bundle.py @@ -137,10 +137,11 @@ class ExportMetadata: (the bundle's HDF5 attributes) and :meth:`to_json`/:meth:`from_json` (the FAISS cache's per-index JSON sidecar). Each direction populates every field even though it only reads back some of them: e.g. the - bundle path never reads ``index_config`` back, the FAISS path never - reads ``provider_type`` back: both are free to obtain (already on - the ``EmbeddingModelRecord`` fetched at write time), so it isn't worth - two separate types over. + FAISS path never reads ``provider_type`` back (the bundle path does + read ``index_config`` back, via ``import_bundle(rebuild_index=True)``); + both directions' unused fields are free to obtain (already on the + ``EmbeddingModelRecord`` fetched at write time), so it isn't worth two + separate types over. """ model_name: str diff --git a/src/omop_emb/utils/cdm.py b/src/omop_emb/utils/cdm.py index f45643d..9cb1c6a 100644 --- a/src/omop_emb/utils/cdm.py +++ b/src/omop_emb/utils/cdm.py @@ -16,6 +16,12 @@ logger = logging.getLogger(__name__) +def streamed(stmt: Select, batch_size: int) -> Select: + """Return *stmt* with server-side cursor streaming enabled. + """ + return stmt.execution_options(stream_results=True, yield_per=batch_size) + + def concept_embedding_projection() -> Select: """Select CDM fields and canonical derived flags used for embeddings.""" @@ -90,9 +96,7 @@ def iter_cdm_concepts_for_filter( if concept_filter is not None: query = concept_filter.apply(query) with cdm_session(cdm_engine) as session: - yield from session.execute( - query.execution_options(stream_results=True, yield_per=chunk_size) - ) + yield from session.execute(streamed(query, chunk_size)) def count_missing_concepts( @@ -111,9 +115,7 @@ def count_missing_concepts( query = concept_filter.apply(query) count = 0 with cdm_session(cdm_engine) as session: - for row in session.execute( - query.execution_options(stream_results=True, yield_per=chunk_size) - ): + for row in session.execute(streamed(query, chunk_size)): if row.concept_id not in embedded_ids: count += 1 return count diff --git a/src/omop_emb/utils/errors.py b/src/omop_emb/utils/errors.py index 3597956..ece7a1a 100644 --- a/src/omop_emb/utils/errors.py +++ b/src/omop_emb/utils/errors.py @@ -6,15 +6,25 @@ class EmbeddingBackendError(RuntimeError): class UnknownEmbeddingBackendError(EmbeddingBackendError): - """Raised when a requested backend name is not recognized.""" + """Error type for an unrecognized backend name. + + Not currently raised anywhere in this package: ``resolve_backend()`` + raises a plain ``RuntimeError`` for an unknown backend name instead. + """ class EmbeddingBackendDependencyError(EmbeddingBackendError, ImportError): - """Raised when a backend was requested but its optional dependencies are missing.""" + """Error type for a backend requested without its optional dependencies installed. + + Not currently raised anywhere in this package. + """ class EmbeddingBackendConfigurationError(EmbeddingBackendError): - """Raised when backend selection or configuration is internally inconsistent.""" + """Error type for an internally inconsistent backend selection or configuration. + + Not currently raised anywhere in this package. + """ class ModelRegistrationConflictError(Exception): diff --git a/tests/conftest.py b/tests/conftest.py index 5c7a4a3..b019ee8 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,8 +2,6 @@ from __future__ import annotations -from typing import Iterator - import numpy as np import pytest import sqlalchemy as sa @@ -58,7 +56,7 @@ # --------------------------------------------------------------------------- # PostgreSQL config (integration tests only) # -# Resolved via OA_Configurator resource 'test_emb_db' in ~/.config/omop/config.toml. +# Resolved via OA_Configurator resource 'test_emb_db_pg' in ~/.config/omop/config.toml. # Run: omop-config configure omop_emb (answer Y when asked to configure test database). # --------------------------------------------------------------------------- @@ -69,11 +67,16 @@ @pytest.fixture -def svec_engine(): - """In-memory SQLiteVec engine, fresh per test.""" - engine = create_sqlitevec_engine(":memory:") - yield engine - engine.dispose() +def svec_engine(request): + """Fresh SQLiteVec engine per test, via oa-configurator's canonical + dialect-agnostic test-database entrypoint rather than a hand-built + ``sa.create_engine()``.""" + from oa_configurator.testing import isolated_test_database + + with isolated_test_database( + OmopEmbConfig, "test_emb_db_sqlite", dialect="sqlite", request=request + ) as db: + yield create_sqlitevec_engine(db.connection.engine) @pytest.fixture @@ -87,29 +90,26 @@ def svec_backend(svec_engine) -> SQLiteVecEmbeddingBackend: # --------------------------------------------------------------------------- -@pytest.fixture(scope="session") -def pg_engine() -> Iterator[sa.Engine]: - """Session-scoped PostgreSQL engine. Skipped when test_emb_db is not configured.""" - from oa_configurator.pytest_plugin import ( - create_fresh_test_db, - drop_test_db, - ensure_test_user_exists, - require_pg_extension, - resolve_test_database, - ) - - raw_url = resolve_test_database(OmopEmbConfig, "test_emb_db") - ensure_test_user_exists(raw_url) - url = create_fresh_test_db(raw_url, extensions=["vector"]) - require_pg_extension(url, "vector") # defensive: verify installation succeeded - engine = sa.create_engine(url, echo=False, future=True) - try: - with engine.connect() as conn: - conn.execute(sa.text("SELECT 1")) - yield engine - finally: - engine.dispose() - drop_test_db(raw_url) +@pytest.fixture +def pg_db(request): + """Canonical isolated PostgreSQL test database (Phase 0 of the + schema_translate_map fix).""" + from oa_configurator.testing import isolated_test_database + + with isolated_test_database( + OmopEmbConfig, "test_emb_db_pg", extensions=["vector"], request=request + ) as db: + yield db + + +@pytest.fixture +def pg_engine(pg_db) -> sa.Engine: + """Real, committing engine for ``PGVectorEmbeddingBackend`` (needs + ``.begin()``/``.connect()`` semantics a bare ``Connection`` can't give). + A thin shim over ``pg_db.connection.engine``; isolation comes from + ``pg_backend``'s teardown (drops each model's table), not a rollback. + """ + return pg_db.committing_engine @pytest.fixture diff --git a/tests/test_cli_pgvector_snapshot.py b/tests/test_cli_pgvector_snapshot.py index 6699b0f..52d3f6c 100644 --- a/tests/test_cli_pgvector_snapshot.py +++ b/tests/test_cli_pgvector_snapshot.py @@ -10,7 +10,6 @@ from .conftest import EMBEDDING_DIM, MODEL_NAME, PROVIDER_TYPE -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration def test_list_registered_models_empty(pg_backend) -> None: @@ -21,7 +20,6 @@ def test_list_registered_models_empty(pg_backend) -> None: assert results == () -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration def test_list_registered_models_after_registration(pg_backend) -> None: diff --git a/tests/test_config.py b/tests/test_config.py index 803ee33..dd476bd 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1,6 +1,45 @@ -"""Placeholder: backend factory tests removed. +"""Tests for omop_emb.config's per-backend index/metric support accessors.""" -``omop_emb.backends.factory`` was not carried forward in the refactor. -Backend construction is done directly via ``SQLiteVecEmbeddingBackend(emb_engine=...)`` -or ``PGVectorEmbeddingBackend(emb_engine=...)``. -""" +from __future__ import annotations + +import pytest + +from omop_emb.config import ( + BackendType, + IndexType, + MetricType, + SUPPORTED_INDICES_AND_METRICS_PER_BACKEND, + get_supported_index_types_for_backend, + get_supported_metrics_for_backend, + is_index_type_supported_for_backend, + is_supported_index_metric_combination_for_backend, +) + + +@pytest.mark.unit +class TestUnregisteredBackendRaises: + """An unregistered backend must raise, not silently report 'supports + nothing'. Simulated via monkeypatch since BackendType has no member + left unregistered today.""" + + def test_is_supported_index_metric_combination_for_backend(self, monkeypatch): + monkeypatch.delitem(SUPPORTED_INDICES_AND_METRICS_PER_BACKEND, BackendType.SQLITEVEC) + with pytest.raises(ValueError, match="Unsupported backend"): + is_supported_index_metric_combination_for_backend( + BackendType.SQLITEVEC, IndexType.FLAT, MetricType.L2 + ) + + def test_is_index_type_supported_for_backend(self, monkeypatch): + monkeypatch.delitem(SUPPORTED_INDICES_AND_METRICS_PER_BACKEND, BackendType.SQLITEVEC) + with pytest.raises(ValueError, match="Unsupported backend"): + is_index_type_supported_for_backend(BackendType.SQLITEVEC, IndexType.FLAT) + + def test_get_supported_index_types_for_backend(self, monkeypatch): + monkeypatch.delitem(SUPPORTED_INDICES_AND_METRICS_PER_BACKEND, BackendType.SQLITEVEC) + with pytest.raises(ValueError, match="Unsupported backend"): + get_supported_index_types_for_backend(BackendType.SQLITEVEC) + + def test_get_supported_metrics_for_backend(self, monkeypatch): + monkeypatch.delitem(SUPPORTED_INDICES_AND_METRICS_PER_BACKEND, BackendType.SQLITEVEC) + with pytest.raises(ValueError, match="Unsupported backend"): + get_supported_metrics_for_backend(BackendType.SQLITEVEC) diff --git a/tests/test_embedding_bundle.py b/tests/test_embedding_bundle.py index 2d4bd87..4dfee90 100644 --- a/tests/test_embedding_bundle.py +++ b/tests/test_embedding_bundle.py @@ -11,13 +11,11 @@ import h5py import numpy as np import pytest +from sqlalchemy import create_engine from omop_emb.backends.base_backend import ConceptEmbeddingRecord from omop_emb.backends.index_config import FlatIndexConfig -from omop_emb.backends.sqlitevec import ( - SQLiteVecEmbeddingBackend, - create_sqlitevec_engine, -) +from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine from omop_emb.config import MetricType from omop_emb.storage import embedding_bundle @@ -28,7 +26,7 @@ def _make_backend() -> SQLiteVecEmbeddingBackend: - engine = create_sqlitevec_engine(":memory:") + engine = create_sqlitevec_engine(create_engine("sqlite:///:memory:")) return SQLiteVecEmbeddingBackend(emb_engine=engine) diff --git a/tests/test_faiss_cache.py b/tests/test_faiss_cache.py index d266524..438a60d 100644 --- a/tests/test_faiss_cache.py +++ b/tests/test_faiss_cache.py @@ -8,14 +8,12 @@ import numpy as np import pytest +from sqlalchemy import create_engine import faiss from omop_emb.backends.base_backend import ConceptEmbeddingRecord from omop_emb.backends.index_config import FlatIndexConfig, HNSWIndexConfig -from omop_emb.backends.sqlitevec import ( - SQLiteVecEmbeddingBackend, - create_sqlitevec_engine, -) +from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine from omop_emb.config import MetricType from omop_emb.storage import embedding_bundle from omop_emb.storage.faiss.faiss_cache import FAISSCache @@ -72,7 +70,7 @@ def _make_svec_backend(dim: int) -> SQLiteVecEmbeddingBackend: - engine = create_sqlitevec_engine(":memory:") + engine = create_sqlitevec_engine(create_engine("sqlite:///:memory:")) return SQLiteVecEmbeddingBackend(emb_engine=engine) diff --git a/tests/test_pgvector.py b/tests/test_pgvector.py index 79cccc4..43d8d09 100644 --- a/tests/test_pgvector.py +++ b/tests/test_pgvector.py @@ -13,9 +13,14 @@ "pgvector", reason="omop-emb[pgvector] not installed: skipping pgvector tests" ) +import sqlalchemy as sa +from oa_configurator import Role +from oa_configurator.testing import isolated_test_schema + from omop_emb.backends.index_config import FlatIndexConfig, HNSWIndexConfig from omop_emb.backends.pgvector import PGVectorEmbeddingBackend from omop_emb.config import IndexType, MetricType +from omop_emb.model_registry import RegistryManager from .conftest import ( CONCEPT_EMBEDDINGS, @@ -28,7 +33,6 @@ from .shared_backend_tests import SharedBackendTests -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration class TestPGVectorBackend(SharedBackendTests): @@ -39,7 +43,6 @@ def backend(self, pg_backend: PGVectorEmbeddingBackend): return pg_backend -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration class TestPGVectorHNSWBackend: @@ -137,3 +140,97 @@ def test_rebuild_index(self, pg_backend: PGVectorEmbeddingBackend): k=1, ) assert results[0][0].concept_id == HYPERTENSION_ID + + +@pytest.mark.pgvector +@pytest.mark.integration +class TestPGVectorNonDefaultSchema: + """Every method here defaulted to the public schema in existing coverage, + so a bug that silently ignored schema_translate_map would still pass + every other test in this file. This is what actually catches that. + """ + + HNSW_CONFIG = HNSWIndexConfig( + metric_type=MetricType.L2, num_neighbors=4, ef_search=8, ef_construction=16 + ) + + @pytest.fixture + def scoped_backend(self, pg_engine): + with isolated_test_schema(pg_engine, prefix="emb_schema") as schema: + scoped_engine = pg_engine.execution_options( + schema_translate_map={Role.PRIMARY.value: schema} + ) + backend = PGVectorEmbeddingBackend(emb_engine=scoped_engine) + yield backend, schema + + def test_table_and_index_lifecycle_stays_in_the_configured_schema( + self, scoped_backend, pg_engine + ): + backend, schema = scoped_backend + record = backend.register_model( + model_name=MODEL_NAME, + provider_type=PROVIDER_TYPE, + index_config=FlatIndexConfig(), + dimensions=EMBEDDING_DIM, + ) + backend.upsert_embeddings( + model_name=MODEL_NAME, + metric_type=MetricType.L2, + records=list(CONCEPT_RECORDS), + embeddings=CONCEPT_EMBEDDINGS, + ) + + # table_exists() sees it in the configured schema... + assert backend._storage_table_exists(record) is True + # ...and a bare inspector scoped to "public" doesn't. + assert sa.inspect(pg_engine).has_table( + record.storage_identifier, schema="public" + ) is False + + # get_indexes()/drop_index(): rebuild to HNSW, confirm the index lands + # in the configured schema, then drop it. + backend.rebuild_index(model_name=MODEL_NAME, index_config=self.HNSW_CONFIG) + manager = backend.get_index_manager(record.storage_identifier) + assert manager.has_index(MetricType.L2) is True + indexes_in_schema = sa.inspect(pg_engine).get_indexes( + record.storage_identifier, schema=schema + ) + assert any( + idx["name"] == manager._index_name(MetricType.L2) for idx in indexes_in_schema + ) + manager.drop_index(MetricType.L2) + assert manager.has_index(MetricType.L2) is False + + # drop_pg_embedding_table(): drops from the configured schema, not public. + backend.delete_model(model_name=MODEL_NAME) + assert backend._storage_table_exists(record) is False + + def test_model_registry_lives_in_its_own_reserved_schema( + self, scoped_backend, pg_engine + ): + backend, schema = scoped_backend + backend.register_model( + model_name=MODEL_NAME, + provider_type=PROVIDER_TYPE, + index_config=FlatIndexConfig(), + dimensions=EMBEDDING_DIM, + ) + + scoped_engine = pg_engine.execution_options( + schema_translate_map={Role.PRIMARY.value: schema} + ) + registry = RegistryManager.read_only(scoped_engine) + assert registry.registry_available is True + assert len(registry.get_registered_models(model_name=MODEL_NAME)) == 1 + + # A registry built with a completely different primary-tagged schema + # sees the exact same row: the registry is decoupled from it entirely. + public_engine = pg_engine.execution_options( + schema_translate_map={Role.PRIMARY.value: "public"} + ) + public_registry = RegistryManager.read_only(public_engine) + assert public_registry.registry_available is True + assert len(public_registry.get_registered_models(model_name=MODEL_NAME)) == 1 + + # And genuinely never created under the storage schema's own name. + assert sa.inspect(pg_engine).has_table("model_registry", schema=schema) is False diff --git a/tests/test_pgvector_index_manager.py b/tests/test_pgvector_index_manager.py index c7703e8..6468c64 100644 --- a/tests/test_pgvector_index_manager.py +++ b/tests/test_pgvector_index_manager.py @@ -31,7 +31,19 @@ ) -@pytest.fixture(scope="module") +@pytest.fixture +def unconnected_pg_engine() -> sa.Engine: + """A real Engine, never connect()ed or begin()'d. + + For tests that only build DDL strings or exercise no-op methods and + never actually touch a database: sa.create_engine() never touches the + network on its own, but still gives qualified() a real Postgres + dialect to quote against, unlike passing None. + """ + return sa.create_engine("postgresql+psycopg://unused:unused@localhost/unused") + + +@pytest.fixture def hnsw_table(pg_engine): with pg_engine.begin() as conn: conn.execute( @@ -71,9 +83,9 @@ def test_supported_index_type(self): mgr._index_config = FlatIndexConfig() assert mgr.supported_index_type == IndexType.FLAT - def test_has_index_always_true(self): + def test_has_index_always_true(self, unconnected_pg_engine): mgr = PGVectorFlatIndexManager( - emb_engine=None, + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", # type: ignore[arg-type] index_config=FlatIndexConfig(), @@ -82,9 +94,9 @@ def test_has_index_always_true(self): assert mgr.has_index(MetricType.L2) is True assert mgr.has_index(MetricType.COSINE) is True - def test_create_index_noop(self): + def test_create_index_noop(self, unconnected_pg_engine): mgr = PGVectorFlatIndexManager( - emb_engine=None, + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", # type: ignore[arg-type] index_config=FlatIndexConfig(), @@ -92,9 +104,9 @@ def test_create_index_noop(self): ) mgr.create_index(MetricType.L2) - def test_drop_index_noop(self): + def test_drop_index_noop(self, unconnected_pg_engine): mgr = PGVectorFlatIndexManager( - emb_engine=None, + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", # type: ignore[arg-type] index_config=FlatIndexConfig(), @@ -102,9 +114,9 @@ def test_drop_index_noop(self): ) mgr.drop_index(MetricType.L2) - def test_create_index_ddl_returns_none(self): + def test_create_index_ddl_returns_none(self, unconnected_pg_engine): mgr = PGVectorFlatIndexManager( - emb_engine=None, + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", # type: ignore[arg-type] index_config=FlatIndexConfig(), @@ -112,10 +124,10 @@ def test_create_index_ddl_returns_none(self): ) assert mgr._create_index_ddl(MetricType.L2) is None - def test_wrong_index_config_raises(self): + def test_wrong_index_config_raises(self, unconnected_pg_engine): with pytest.raises(ValueError, match="index_type"): PGVectorFlatIndexManager( - emb_engine=None, # type: ignore + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", index_config=HNSWIndexConfig(metric_type=MetricType.L2), # type: ignore @@ -131,9 +143,9 @@ def test_wrong_index_config_raises(self): @pytest.mark.unit class TestPGVectorHNSWIndexManagerDDL: @pytest.fixture - def mgr(self) -> PGVectorHNSWIndexManager: + def mgr(self, unconnected_pg_engine) -> PGVectorHNSWIndexManager: m = PGVectorHNSWIndexManager.__new__(PGVectorHNSWIndexManager) - m._engine = None # type: ignore + m._engine = unconnected_pg_engine m._tablename = "my_table" m._embedding_column = "embedding" m._index_config = HNSWIndexConfig( @@ -172,19 +184,19 @@ def test_ddl_hamming_generates_bit_ops_ddl(self, mgr): def test_supported_index_type(self, mgr): assert mgr.supported_index_type == IndexType.HNSW - def test_wrong_config_type_raises(self): + def test_wrong_config_type_raises(self, unconnected_pg_engine): with pytest.raises(ValueError, match="index_type"): PGVectorHNSWIndexManager( - emb_engine=None, # type: ignore + emb_engine=unconnected_pg_engine, tablename="t", embedding_column="e", index_config=FlatIndexConfig(), # type: ignore dimensions=4, ) - def test_halfvec_ddl_uses_halfvec_ops(self): + def test_halfvec_ddl_uses_halfvec_ops(self, unconnected_pg_engine): m = PGVectorHNSWIndexManager.__new__(PGVectorHNSWIndexManager) - m._engine = None # type: ignore + m._engine = unconnected_pg_engine m._tablename = "my_table" m._embedding_column = "embedding" m._index_config = HNSWIndexConfig( @@ -252,7 +264,6 @@ def test_get_similarity_raises_valueerror_for_hamming(self): # --------------------------------------------------------------------------- -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration class TestPGVectorHNSWIndexManagerIntegration: diff --git a/tests/test_read_only_population.py b/tests/test_read_only_population.py index 71863b8..f066daf 100644 --- a/tests/test_read_only_population.py +++ b/tests/test_read_only_population.py @@ -9,7 +9,7 @@ from omop_emb.backends import ReadOnlyEmbeddingStore, StoredEmbedding from omop_emb.backends.embedding_table import concept_metadata_table_descriptor from omop_emb.backends.index_config import FlatIndexConfig -from omop_emb.model_registry import RegistryManager, ensure_registry_schema +from omop_emb.model_registry import RegistryManager, ensure_registry_table from omop_emb.population import PopulationScope, plan_population @@ -35,7 +35,7 @@ def test_read_only_registry_does_not_create_schema() -> None: store = ReadOnlyEmbeddingStore( engine, backend_type="sqlitevec", - schema="main", + physical_schema="main", ) assert store.initialized is False @@ -46,11 +46,11 @@ def test_read_only_registry_does_not_create_schema() -> None: def test_explicit_registry_initialization_is_visible_to_read_only_store() -> None: engine = create_engine("sqlite:///:memory:") - ensure_registry_schema(engine) + ensure_registry_table(engine) store = ReadOnlyEmbeddingStore( engine, backend_type="sqlitevec", - schema="main", + physical_schema="main", ) assert store.initialized is True @@ -93,7 +93,7 @@ def capture(_connection, _cursor, statement, _parameters, _context, _many): with ReadOnlyEmbeddingStore( engine, backend_type="sqlitevec", - schema="main", + physical_schema="main", ) as store: assert store.stored_embeddings("test-model") == ( StoredEmbedding(7, "Condition", "SNOMED", True, True), @@ -112,7 +112,7 @@ def capture(_connection, _cursor, statement, _parameters, _context, _many): def test_read_only_registry_rejects_mutation() -> None: engine = create_engine("sqlite:///:memory:") - ensure_registry_schema(engine) + ensure_registry_table(engine) registry = RegistryManager.read_only(engine) with pytest.raises(RuntimeError, match="opened read-only"): diff --git a/tests/test_registry.py b/tests/test_registry.py index fcc716f..89a0b01 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -5,9 +5,11 @@ import pytest import sqlalchemy as sa +from oa_configurator import ensure_schema + from omop_emb.backends.index_config import FlatIndexConfig, HNSWIndexConfig -from omop_emb.config import IndexType, MetricType -from omop_emb.model_registry import RegistryManager, ensure_registry_schema +from omop_emb.config import MODEL_REGISTRY_SCHEMA, IndexType, MetricType +from omop_emb.model_registry import RegistryManager, ensure_registry_table from omop_emb.utils.errors import ModelRegistrationConflictError from .conftest import EMBEDDING_DIM, MODEL_NAME, PROVIDER_TYPE @@ -242,37 +244,40 @@ def test_legacy_provider_name_is_normalized_in_sqlite(svec_engine): ) == "ollama" -@pytest.mark.requires_database("test_emb_db") @pytest.mark.pgvector @pytest.mark.integration def test_legacy_provider_column_is_widened_in_postgres(pg_engine): + with pg_engine.begin() as connection: + ensure_schema(connection, MODEL_REGISTRY_SCHEMA) + connection.execute(sa.text(f"DROP TABLE IF EXISTS {MODEL_REGISTRY_SCHEMA}.model_registry CASCADE")) + connection.execute( + sa.text(f"CREATE TABLE {MODEL_REGISTRY_SCHEMA}.model_registry (provider_type VARCHAR(6))") + ) + connection.execute( + sa.text(f"INSERT INTO {MODEL_REGISTRY_SCHEMA}.model_registry (provider_type) VALUES ('OLLAMA')") + ) try: - with pg_engine.begin() as connection: - connection.execute(sa.text("DROP TABLE IF EXISTS model_registry CASCADE")) - connection.execute( - sa.text("CREATE TABLE model_registry (provider_type VARCHAR(6))") - ) - connection.execute( - sa.text("INSERT INTO model_registry (provider_type) VALUES ('OLLAMA')") - ) - RegistryManager(pg_engine) provider_column = next( column - for column in sa.inspect(pg_engine).get_columns("model_registry") + for column in sa.inspect(pg_engine).get_columns( + "model_registry", schema=MODEL_REGISTRY_SCHEMA + ) if column["name"] == "provider_type" ) assert getattr(provider_column["type"], "length", None) is None with pg_engine.begin() as connection: assert connection.scalar( - sa.text("SELECT provider_type FROM model_registry") + sa.text(f"SELECT provider_type FROM {MODEL_REGISTRY_SCHEMA}.model_registry") ) == "ollama" connection.execute( - sa.text("INSERT INTO model_registry (provider_type) VALUES ('anthropic')") + sa.text( + f"INSERT INTO {MODEL_REGISTRY_SCHEMA}.model_registry (provider_type) VALUES ('anthropic')" + ) ) finally: with pg_engine.begin() as connection: - connection.execute(sa.text("DROP TABLE IF EXISTS model_registry CASCADE")) - ensure_registry_schema(pg_engine) + connection.execute(sa.text(f"DROP TABLE IF EXISTS {MODEL_REGISTRY_SCHEMA}.model_registry CASCADE")) + ensure_registry_table(pg_engine) diff --git a/tests/test_resolve_backend.py b/tests/test_resolve_backend.py new file mode 100644 index 0000000..a0f4be3 --- /dev/null +++ b/tests/test_resolve_backend.py @@ -0,0 +1,56 @@ +"""resolve_backend() is the real production entry point every CLI/config- +driven caller goes through to build an EmbeddingBackend. This suite +tests regressions in resolve_backend() itself (e.g. wrong dialect detection, +a dropped registry_schema_translate_map key, wiring a different backend class) +""" + +from __future__ import annotations + +import pytest + +from omop_emb.backends.base_backend import resolve_backend +from omop_emb.backends.pgvector import PGVectorEmbeddingBackend +from omop_emb.config import BackendType + +from .conftest import EMBEDDING_DIM, MODEL_NAME, PROVIDER_TYPE + +pytestmark = [pytest.mark.postgresql, pytest.mark.db_dialect] + + +def test_resolve_backend_builds_a_working_pgvector_backend(pg_db): + backend = resolve_backend(BackendType.PGVECTOR, database=pg_db.resolved) + try: + assert isinstance(backend, PGVectorEmbeddingBackend) + + record = backend.register_model( + model_name=MODEL_NAME, provider_type=PROVIDER_TYPE, dimensions=EMBEDDING_DIM + ) + assert record.model_name == MODEL_NAME + + registered = {r.model_name for r in backend.get_registered_models()} + assert MODEL_NAME in registered + finally: + for record in backend.get_registered_models(): + try: + backend.delete_model(model_name=record.model_name) + except Exception: + pass + + +def test_resolve_backend_accepts_the_backend_type_as_a_plain_string(pg_db): + """Callers resolving from a `[vector_stores.*]` config entry pass a plain + string, not the BackendType enum member.""" + backend = resolve_backend("pgvector", database=pg_db.resolved) + try: + assert isinstance(backend, PGVectorEmbeddingBackend) + finally: + for record in backend.get_registered_models(): + try: + backend.delete_model(model_name=record.model_name) + except Exception: + pass + + +def test_resolve_backend_rejects_an_unknown_backend_type(pg_db): + with pytest.raises(RuntimeError, match="Unknown backend"): + resolve_backend("not-a-real-backend", database=pg_db.resolved) diff --git a/tests/test_schema_provenance_guard.py b/tests/test_schema_provenance_guard.py new file mode 100644 index 0000000..b3efd85 --- /dev/null +++ b/tests/test_schema_provenance_guard.py @@ -0,0 +1,148 @@ +"""schema-provenance guard wired into PGVectorEmbeddingBackend's +registry-schema and storage-table creation (ensure_registry_table(), +create_pg_embedding_table()). + +Only the "fires on a genuinely reconfigured schema" case is covered here. +The guard's own agree/no-op/test_only semantics are already exhaustively +covered at the primitive level in oa-configurator's own test suite; what's +worth proving per consuming repo is that this call site is actually wired +to it, and a wiring mistake would show up here too. + +pg_engine is a real, committing engine, so every provenance row this test +writes is a genuine commit. cleanup_after_test deletes this test's own +schema_provenance rows at teardown (see Phase 10.12 in the plan). +""" + +from __future__ import annotations + +import dataclasses +import uuid + +import pytest +from oa_configurator import SchemaDriftError, record_schema_provenance +from oa_configurator.domains.resources.sql import SCHEMA_PROVENANCE_SCHEMA, _schema_provenance_table +from oa_configurator.testing import delete_rows_on_cleanup, isolated_test_schema + +from omop_emb.backends.pgvector.pg_backend import PGVectorEmbeddingBackend +from omop_emb.config import MODEL_REGISTRY_SCHEMA +from omop_emb.model_registry import REGISTRY_SCHEMA_KEY + +from .conftest import EMBEDDING_DIM, MODEL_NAME, PROVIDER_TYPE + +pytestmark = [pytest.mark.postgresql, pytest.mark.db_dialect] + + +def _resolved(pg_db, *, database_name: str, schema: str): + """pg_db.resolved with a unique name (the guard's own key includes it), + schema_name overridden, and connection.test_only forced False so the + guard doesn't no-op against pg_db's own test-only marking. + """ + return dataclasses.replace( + pg_db.resolved, + name=database_name, + schema_name=schema, + connection=dataclasses.replace(pg_db.resolved.connection, test_only=False), + ) + + +def _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) -> None: + """The registry's own provenance row is keyed by database_name=MODEL_REGISTRY_SCHEMA + The phyical registry tables outlives any single test run, so a delete-only reset + leaves "table already has rows, but no provenance record". + Overwrite it with a known-correct baseline instead, via record_schema_provenance + (which always overwrites, no "already populated" check), then register cleanup. + """ + table = _schema_provenance_table(SCHEMA_PROVENANCE_SCHEMA) + with pg_engine.begin() as connection: + record_schema_provenance( + connection, + database_name=MODEL_REGISTRY_SCHEMA, + schema_tag=REGISTRY_SCHEMA_KEY, + new_physical_schema=MODEL_REGISTRY_SCHEMA, + reason="test setup: establish a known-correct baseline", + ) + delete_rows_on_cleanup( + cleanup_after_test, pg_engine, table, table.c.database_name == MODEL_REGISTRY_SCHEMA + ) + + +def test_backend_construction_is_unaffected_by_the_primary_schema_changing( + pg_db, pg_engine, cleanup_after_test +): + _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) + with ( + isolated_test_schema(pg_engine, prefix="emb_guard_a") as schema_a, + isolated_test_schema(pg_engine, prefix="emb_guard_b") as schema_b, + ): + database_name = f"emb_guard_db_{uuid.uuid4().hex[:8]}" + resolved_a = _resolved(pg_db, database_name=database_name, schema=schema_a) + engine_a = pg_engine.execution_options(schema_translate_map={"primary": schema_a}) + backend_a = PGVectorEmbeddingBackend(emb_engine=engine_a, resolved=resolved_a) + assert backend_a is not None + + # A different database_name too: the registry provenance row must + # be shared correctly across distinct database entries pointed at + # the same connection, not just across two calls with one name. + resolved_b = _resolved( + pg_db, database_name=f"emb_guard_db_{uuid.uuid4().hex[:8]}", schema=schema_b + ) + engine_b = pg_engine.execution_options(schema_translate_map={"primary": schema_b}) + backend_b = PGVectorEmbeddingBackend(emb_engine=engine_b, resolved=resolved_b) + assert backend_b is not None + + +def test_registry_schema_guard_fires_on_genuine_registry_drift(pg_db, pg_engine, cleanup_after_test): + table = _schema_provenance_table(SCHEMA_PROVENANCE_SCHEMA) + delete_rows_on_cleanup( + cleanup_after_test, pg_engine, table, table.c.database_name == MODEL_REGISTRY_SCHEMA + ) + resolved = _resolved( + pg_db, + database_name=f"emb_guard_db_{uuid.uuid4().hex[:8]}", + schema="unrelated_primary_schema", + ) + + with pg_engine.begin() as connection: + record_schema_provenance( + connection, + database_name=MODEL_REGISTRY_SCHEMA, + schema_tag=REGISTRY_SCHEMA_KEY, + new_physical_schema="a_previous_registry_schema_that_is_not_current", + reason="test: force a stale baseline", + ) + + engine = pg_engine.execution_options(schema_translate_map={"primary": "unrelated_primary_schema"}) + with pytest.raises(SchemaDriftError): + PGVectorEmbeddingBackend(emb_engine=engine, resolved=resolved) + + +def test_primary_schema_guard_fires_on_genuine_primary_drift(pg_db, pg_engine, cleanup_after_test): + """Tests embedding storage table's own guard. Only triggered + once a model is actually registered, not at backend construction time. + Restores the PRIMARY schema-tag drift coverage a prior rewrite replaced instead + of adding alongside""" + _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) + table = _schema_provenance_table(SCHEMA_PROVENANCE_SCHEMA) + database_name = f"emb_guard_db_{uuid.uuid4().hex[:8]}" + delete_rows_on_cleanup( + cleanup_after_test, pg_engine, table, table.c.database_name == database_name + ) + + with ( + isolated_test_schema(pg_engine, prefix="emb_guard_primary_a") as schema_a, + isolated_test_schema(pg_engine, prefix="emb_guard_primary_b") as schema_b, + ): + resolved_a = _resolved(pg_db, database_name=database_name, schema=schema_a) + engine_a = pg_engine.execution_options(schema_translate_map={"primary": schema_a}) + backend_a = PGVectorEmbeddingBackend(emb_engine=engine_a, resolved=resolved_a) + backend_a.register_model( + model_name=MODEL_NAME, provider_type=PROVIDER_TYPE, dimensions=EMBEDDING_DIM + ) + + resolved_b = _resolved(pg_db, database_name=database_name, schema=schema_b) + engine_b = pg_engine.execution_options(schema_translate_map={"primary": schema_b}) + # Constructing backend_b already reloads the model registered above + # (shared registry schema, same database_name) and tries to load its + # storage table under the new schema, which is where the guard fires + with pytest.raises(SchemaDriftError): + PGVectorEmbeddingBackend(emb_engine=engine_b, resolved=resolved_b) diff --git a/tests/test_sqlitevec.py b/tests/test_sqlitevec.py index aa509e6..78cce53 100644 --- a/tests/test_sqlitevec.py +++ b/tests/test_sqlitevec.py @@ -7,9 +7,10 @@ import numpy as np import pytest +from sqlalchemy import create_engine from omop_emb.backends.index_config import FlatIndexConfig, HNSWIndexConfig -from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend +from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine from omop_emb.config import MetricType from .conftest import ( @@ -129,9 +130,11 @@ def test_one_table_per_model(self, svec_backend: SQLiteVecEmbeddingBackend): assert r1.storage_identifier == r2.storage_identifier assert len(svec_backend.get_registered_models(model_name=MODEL_NAME)) == 1 - def test_from_path_constructor(self, tmp_path): + def test_file_backed_engine_constructor(self, tmp_path): + """A real file-backed (not just in-memory) engine works end to end.""" db_file = str(tmp_path / "test.db") - backend = SQLiteVecEmbeddingBackend.from_path(db_file) + engine = create_sqlitevec_engine(create_engine(f"sqlite:///{db_file}")) + backend = SQLiteVecEmbeddingBackend(emb_engine=engine) record = backend.register_model( model_name=MODEL_NAME, provider_type=PROVIDER_TYPE,