From a850d2616286e835d6a6f19ea87ff237f058df5d Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Tue, 1 Sep 2026 22:56:12 +0000 Subject: [PATCH 01/16] Test per dialect with new method, schema-aware changes from oa-configurator, adaption of the config, update CI --- .github/workflows/ci.yml | 20 ++-- docs/usage/interface-guide.md | 11 ++- pyproject.toml | 2 +- src/omop_emb/backends/base_backend.py | 9 +- src/omop_emb/backends/pgvector/pg_backend.py | 18 +--- .../backends/pgvector/pg_index_manager.py | 13 ++- src/omop_emb/backends/pgvector/pg_sql.py | 7 +- src/omop_emb/backends/read_only.py | 18 ++-- .../backends/sqlitevec/sqlitevec_backend.py | 31 ++---- src/omop_emb/config.py | 13 ++- .../model_registry/model_registry_manager.py | 5 +- .../model_registry/model_registry_orm.py | 19 +++- src/omop_emb/population.py | 4 +- src/omop_emb/utils/cdm.py | 14 +-- tests/conftest.py | 62 ++++++------ tests/test_cli_pgvector_snapshot.py | 2 - tests/test_embedding_bundle.py | 8 +- tests/test_faiss_cache.py | 8 +- tests/test_pgvector.py | 96 ++++++++++++++++++- tests/test_pgvector_index_manager.py | 3 +- tests/test_registry.py | 1 - tests/test_sqlitevec.py | 9 +- 22 files changed, 236 insertions(+), 137 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c90e472..002b4cc 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,6 +6,8 @@ on: jobs: label-gate: uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/label-gate.yml@main + build-test-sqlite: + uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test.yml@main build-test: uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test-postgres.yml@main with: @@ -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/docs/usage/interface-guide.md b/docs/usage/interface-guide.md index 5846e2d..11b794e 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") +) ``` --- 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..4811ea1 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -1037,16 +1037,15 @@ def resolve_backend( dialect = make_url(database.connection.url).get_backend_name() 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": 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()) + logger.info(f"Using SQLiteVec backend with engine: {emb_engine.url}") + return SQLiteVecEmbeddingBackend(emb_engine=emb_engine) if resolved_backend == BackendType.PGVECTOR: if dialect != "postgresql": diff --git a/src/omop_emb/backends/pgvector/pg_backend.py b/src/omop_emb/backends/pgvector/pg_backend.py index 876fd49..59f3edd 100644 --- a/src/omop_emb/backends/pgvector/pg_backend.py +++ b/src/omop_emb/backends/pgvector/pg_backend.py @@ -5,7 +5,7 @@ from typing import Mapping, Optional, Sequence, Tuple from numpy import ndarray -from sqlalchemy import Engine, select, text, create_engine +from sqlalchemy import Engine, select, text try: from pgvector.sqlalchemy import Vector # noqa: F401 @@ -82,22 +82,6 @@ def __init__(self, emb_engine: Engine) -> 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) - # ------------------------------------------------------------------ # Backend identity # ------------------------------------------------------------------ diff --git a/src/omop_emb/backends/pgvector/pg_index_manager.py b/src/omop_emb/backends/pgvector/pg_index_manager.py index a3d57c1..9939689 100644 --- a/src/omop_emb/backends/pgvector/pg_index_manager.py +++ b/src/omop_emb/backends/pgvector/pg_index_manager.py @@ -17,7 +17,8 @@ import logging from typing import Generic, TypeVar -from sqlalchemy import Engine, inspect, text +from oa_configurator import qualified, schema_inspect +from sqlalchemy import Engine, text from omop_emb.config import IndexType, MetricType, VectorColumnType from omop_emb.backends.index_config import IndexConfig, FlatIndexConfig, HNSWIndexConfig @@ -77,7 +78,7 @@ 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 schema_inspect(conn).get_indexes(self._tablename) } return self._index_name(metric_type) in existing @@ -101,7 +102,7 @@ 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)}")) if existed: logger.info(f"Dropped pgvector index '{name}'.") @@ -190,9 +191,13 @@ 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 + # self._engine is None only in pure-DDL-string unit tests that never + # open a connection; qualified() needs a real bindable for its + # schema, so fall back to the bare name in that case only. + table_ref = qualified(self._engine, self._tablename) if self._engine is not None else self._tablename 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..b319eb5 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -15,6 +15,7 @@ from typing import List, Optional, Sequence, Union from numpy import ndarray +from oa_configurator import qualified, schema_inspect 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,8 +32,8 @@ 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 schema_inspect(engine).has_table(table_name) def create_pg_embedding_table( engine: Engine, model_record: EmbeddingModelRecord @@ -71,7 +72,7 @@ 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)}")) logger.info(f"Dropped embedding table '{tablename}'.") diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index 3e1cc9d..dda5197 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -22,6 +22,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) @@ -103,13 +104,16 @@ def iter_stored_embeddings( record.storage_identifier, schema=schema, ) - 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) + 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, + ) with self._engine.connect() as connection: rows = connection.execute(statement).mappings() for row in rows: diff --git a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py index eedbae7..ae5ab47 100644 --- a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py +++ b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py @@ -11,7 +11,7 @@ import numpy as np from numpy import ndarray -from sqlalchemy import Engine, MetaData, Table, create_engine, event, text +from sqlalchemy import Engine, MetaData, Table, event, text try: import sqlite_vec @@ -50,20 +50,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): @@ -94,21 +96,6 @@ def __init__(self, emb_engine: Engine) -> 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)) - # ------------------------------------------------------------------ # Backend identity # ------------------------------------------------------------------ diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index cade3d3..0c61559 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -37,7 +37,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=( diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index eba015d..44ac278 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -5,7 +5,8 @@ from datetime import datetime, timezone from typing import Mapping, Optional -from sqlalchemy import Engine, inspect, select, update +from oa_configurator import schema_inspect +from sqlalchemy import Engine, select, update from sqlalchemy.orm import Session, sessionmaker from omop_emb.backends.index_config import ( @@ -39,7 +40,7 @@ def __init__(self, embedding_engine: Engine, *, initialize: bool = True) -> None self._embedding_engine = embedding_engine self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize - self._registry_available = inspect(embedding_engine).has_table( + self._registry_available = schema_inspect(embedding_engine).has_table( ModelRegistry.__tablename__ ) if initialize: diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index 8cd72d4..a32bbe4 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -1,7 +1,9 @@ from __future__ import annotations +import warnings from typing import Any, Optional +from oa_configurator import qualified, schema_inspect from sqlalchemy import ( DateTime, Engine, @@ -10,7 +12,6 @@ JSON, String, func, - inspect, text, ) from sqlalchemy.orm import DeclarativeBase, mapped_column, validates, Mapped @@ -199,7 +200,7 @@ 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 = schema_inspect(engine).get_columns(ModelRegistry.__tablename__) provider_column = next( (column for column in columns if column["name"] == "provider_type"), None, @@ -210,16 +211,26 @@ 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: + 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__)} " "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__)} " "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/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/tests/conftest.py b/tests/conftest.py index 5c7a4a3..b896b0f 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.connection.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_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..bfb2f01 100644 --- a/tests/test_pgvector.py +++ b/tests/test_pgvector.py @@ -13,9 +13,13 @@ "pgvector", reason="omop-emb[pgvector] not installed: skipping pgvector tests" ) +from oa_configurator import schema_inspect +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 +32,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 +42,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 +139,93 @@ 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={None: 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 schema_inspect(pg_engine, schema="public").has_table( + record.storage_identifier + ) 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 = schema_inspect(pg_engine, schema=schema).get_indexes( + record.storage_identifier + ) + 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_lookup_stays_in_the_configured_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={None: 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 pointed at "public" must not see it: proves the lookup is + # genuinely schema-scoped, not incidentally finding it via search_path. + public_engine = pg_engine.execution_options(schema_translate_map={None: "public"}) + public_registry = RegistryManager.read_only(public_engine) + found_in_public = ( + public_registry.get_registered_models(model_name=MODEL_NAME) + if public_registry.registry_available + else () + ) + assert found_in_public == () diff --git a/tests/test_pgvector_index_manager.py b/tests/test_pgvector_index_manager.py index c7703e8..5f1a9ec 100644 --- a/tests/test_pgvector_index_manager.py +++ b/tests/test_pgvector_index_manager.py @@ -31,7 +31,7 @@ ) -@pytest.fixture(scope="module") +@pytest.fixture def hnsw_table(pg_engine): with pg_engine.begin() as conn: conn.execute( @@ -252,7 +252,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_registry.py b/tests/test_registry.py index fcc716f..4590efe 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -242,7 +242,6 @@ 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): 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, From 7017bafbd11b3884e0c15605ca39046171a1e425 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Thu, 3 Sep 2026 05:58:17 +0000 Subject: [PATCH 02/16] Drop docker --- Dockerfile | 3 --- 1 file changed, 3 deletions(-) delete mode 100644 Dockerfile 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 From d8e11caed7faef5e7d5213bf4fb1fa522cc349f2 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Thu, 3 Sep 2026 05:58:50 +0000 Subject: [PATCH 03/16] Update CI --- .github/workflows/ci.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 002b4cc..a516006 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,8 +8,8 @@ jobs: uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/label-gate.yml@main build-test-sqlite: uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test.yml@main - build-test: - uses: AustralianCancerDataNetwork/cava-devops/.github/workflows/build-test-postgres.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 From ce5c7056404a0a882a9fee0079c64df57a80838f Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Fri, 4 Sep 2026 06:16:08 +0000 Subject: [PATCH 04/16] Carry ResolvedDatabase downstream to other methods adhering to oa-configurator, guard schemas properly --- src/omop_emb/backends/base_backend.py | 30 ++++++--- src/omop_emb/backends/pgvector/pg_backend.py | 14 ++++- src/omop_emb/backends/pgvector/pg_sql.py | 15 ++++- .../backends/sqlitevec/sqlitevec_backend.py | 10 ++- src/omop_emb/config.py | 6 ++ src/omop_emb/model_registry/__init__.py | 4 +- tests/test_schema_provenance_guard.py | 62 +++++++++++++++++++ 7 files changed, 124 insertions(+), 17 deletions(-) create mode 100644 tests/test_schema_provenance_guard.py diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index 4811ea1..c0ba258 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -6,12 +6,13 @@ 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 ResolvedDatabase, ResolvedVectorStore, supports_schemas from sqlalchemy import Engine from sqlalchemy.engine import make_url from sqlalchemy.orm import sessionmaker from omop_emb.config import ( + MODEL_REGISTRY_SCHEMA, BackendType, MetricType, IndexType, @@ -168,7 +169,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 +182,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() @@ -1036,6 +1043,13 @@ 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 -- added on top of database's own + # translate map, not passed as a bare override, since create_engine()'s + # execution_options replaces the whole map rather than merging it. + registry_schema = MODEL_REGISTRY_SCHEMA if supports_schemas(database.connection.dialect_name) else None + schema_translate_map = {**database.schema_translate_map(), "registry": registry_schema} + if resolved_backend == BackendType.SQLITEVEC: from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine @@ -1043,9 +1057,11 @@ def resolve_backend( raise RuntimeError( f"sqlitevec backend requires a sqlite-dialect database, got dialect: {dialect!r}." ) - emb_engine = create_sqlitevec_engine(database.create_engine()) + emb_engine = create_sqlitevec_engine( + database.create_engine(execution_options={"schema_translate_map": schema_translate_map}) + ) logger.info(f"Using SQLiteVec backend with engine: {emb_engine.url}") - return SQLiteVecEmbeddingBackend(emb_engine=emb_engine) + return SQLiteVecEmbeddingBackend(emb_engine=emb_engine, resolved=database) if resolved_backend == BackendType.PGVECTOR: if dialect != "postgresql": @@ -1060,9 +1076,9 @@ 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": 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/pgvector/pg_backend.py b/src/omop_emb/backends/pgvector/pg_backend.py index 59f3edd..c821948 100644 --- a/src/omop_emb/backends/pgvector/pg_backend.py +++ b/src/omop_emb/backends/pgvector/pg_backend.py @@ -5,6 +5,7 @@ from typing import Mapping, Optional, Sequence, Tuple from numpy import ndarray +from oa_configurator import ResolvedDatabase from sqlalchemy import Engine, select, text try: @@ -78,9 +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) + super().__init__(emb_engine=emb_engine, resolved=resolved) # ------------------------------------------------------------------ # Backend identity @@ -119,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: diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index b319eb5..f081efd 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -15,7 +15,7 @@ from typing import List, Optional, Sequence, Union from numpy import ndarray -from oa_configurator import qualified, schema_inspect +from oa_configurator import ResolvedDatabase, Role, guard_schema_provenance, qualified, schema_inspect 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 @@ -36,7 +36,10 @@ def table_exists(engine: Engine, table_name: str) -> bool: return schema_inspect(engine).has_table(table_name) 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. @@ -45,6 +48,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 ------- @@ -58,7 +65,9 @@ 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: + with guard_schema_provenance(connection, resolved, role=Role.PRIMARY): + EmbeddingTableBase.metadata.create_all(connection, tables=[table_cls.__table__]) # ty: ignore[invalid-argument-type] return table_cls diff --git a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py index ae5ab47..9a75685 100644 --- a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py +++ b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py @@ -11,6 +11,7 @@ import numpy as np from numpy import ndarray +from oa_configurator import ResolvedDatabase from sqlalchemy import Engine, MetaData, Table, event, text try: @@ -92,9 +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) + super().__init__(emb_engine=emb_engine, resolved=resolved) # ------------------------------------------------------------------ # Backend identity diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index 0c61559..392075f 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -16,8 +16,14 @@ Resolver, ResolvedVectorStore, VectorStoreConfig, + register_reserved_schema, ) +# Guaranteed to be imported and registered if there is a config +MODEL_REGISTRY_SCHEMA: str = "omop_emb_registry" + +register_reserved_schema(MODEL_REGISTRY_SCHEMA, owner="omop_emb") + class OmopEmbConfig(PackageConfigBase): """oa-configurator config class for omop-emb. diff --git a/src/omop_emb/model_registry/__init__.py b/src/omop_emb/model_registry/__init__.py index 276eb84..e6173ce 100644 --- a/src/omop_emb/model_registry/__init__.py +++ b/src/omop_emb/model_registry/__init__.py @@ -2,12 +2,12 @@ from omop_emb.model_registry.model_registry_manager import RegistryManager from omop_emb.model_registry.model_registry_orm import ( ModelRegistry, - ensure_registry_schema, + ensure_registry_table, ) __all__ = [ "EmbeddingModelRecord", "RegistryManager", "ModelRegistry", - "ensure_registry_schema", + "ensure_registry_table", ] diff --git a/tests/test_schema_provenance_guard.py b/tests/test_schema_provenance_guard.py new file mode 100644 index 0000000..94e0e12 --- /dev/null +++ b/tests/test_schema_provenance_guard.py @@ -0,0 +1,62 @@ +"""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 +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 + +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 test_backend_construction_guard_fires_on_reconfigured_schema(pg_db, pg_engine, cleanup_after_test): + database_name = f"emb_guard_db_{uuid.uuid4().hex[:8]}" + table = _schema_provenance_table(SCHEMA_PROVENANCE_SCHEMA) + 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_a") as schema_a, + isolated_test_schema(pg_engine, prefix="emb_guard_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={None: schema_a}) + backend_a = PGVectorEmbeddingBackend(emb_engine=engine_a, resolved=resolved_a) + assert backend_a is not None + + resolved_b = _resolved(pg_db, database_name=database_name, schema=schema_b) + engine_b = pg_engine.execution_options(schema_translate_map={None: schema_b}) + with pytest.raises(SchemaDriftError): + PGVectorEmbeddingBackend(emb_engine=engine_b, resolved=resolved_b) From a9017d6800c5457737a11eaea9c252207b3a5565 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Sun, 6 Sep 2026 22:55:44 +0000 Subject: [PATCH 05/16] Schema-aware Registry --- src/omop_emb/backends/base_backend.py | 4 +- src/omop_emb/model_registry/__init__.py | 2 + .../model_registry/model_registry_manager.py | 35 ++++++++++--- .../model_registry/model_registry_orm.py | 51 ++++++++++++++++--- tests/test_pgvector.py | 17 +++---- tests/test_read_only_population.py | 6 +-- tests/test_registry.py | 38 ++++++++------ 7 files changed, 107 insertions(+), 46 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index c0ba258..14e4cf4 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -23,7 +23,7 @@ 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 from omop_emb.utils.embedding_utils import ( EmbeddingConceptFilter, NearestConceptMatch, @@ -1048,7 +1048,7 @@ def resolve_backend( # translate map, not passed as a bare override, since create_engine()'s # execution_options replaces the whole map rather than merging it. registry_schema = MODEL_REGISTRY_SCHEMA if supports_schemas(database.connection.dialect_name) else None - schema_translate_map = {**database.schema_translate_map(), "registry": registry_schema} + schema_translate_map = {**database.schema_translate_map(), REGISTRY_SCHEMA_KEY: registry_schema} if resolved_backend == BackendType.SQLITEVEC: from omop_emb.backends.sqlitevec import SQLiteVecEmbeddingBackend, create_sqlitevec_engine diff --git a/src/omop_emb/model_registry/__init__.py b/src/omop_emb/model_registry/__init__.py index e6173ce..d61ac02 100644 --- a/src/omop_emb/model_registry/__init__.py +++ b/src/omop_emb/model_registry/__init__.py @@ -1,6 +1,7 @@ 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_table, ) @@ -10,4 +11,5 @@ "RegistryManager", "ModelRegistry", "ensure_registry_table", + "REGISTRY_SCHEMA_KEY", ] diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index 44ac278..0a2bcc6 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -5,7 +5,7 @@ from datetime import datetime, timezone from typing import Mapping, Optional -from oa_configurator import schema_inspect +from oa_configurator import ResolvedDatabase, schema_inspect from sqlalchemy import Engine, select, update from sqlalchemy.orm import Session, sessionmaker @@ -15,8 +15,10 @@ index_config_from_orm_row, ) from omop_emb.model_registry.model_registry_orm import ( + REGISTRY_SCHEMA_KEY, ModelRegistry, - ensure_registry_schema, + _registry_schema, + ensure_registry_table, ) from omop_emb.model_registry.model_registry_types import EmbeddingModelRecord from omop_emb.utils.errors import ModelRegistrationConflictError @@ -34,17 +36,34 @@ class RegistryManager: ---------- embedding_engine : Engine SQLAlchemy embedding engine connected to the embedding store. + + Notes + ----- + The registry table is created in a schema named ``registry`` for dialects + supporting schema registration. Allows schema-independent access to the registry table + from any schema in the same database. """ - 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: + self._embedding_engine = embedding_engine.execution_options( + schema_translate_map={ + **(embedding_engine.get_execution_options().get("schema_translate_map") or {}), + REGISTRY_SCHEMA_KEY: _registry_schema(embedding_engine), + } + ) self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize - self._registry_available = schema_inspect(embedding_engine).has_table( - ModelRegistry.__tablename__ - ) + self._registry_available = schema_inspect( + self._embedding_engine, schema=_registry_schema(self._embedding_engine) + ).has_table(ModelRegistry.__tablename__) if initialize: - ensure_registry_schema(embedding_engine) + ensure_registry_table(self._embedding_engine, resolved=resolved) self._registry_available = True @classmethod diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index a32bbe4..2b0f199 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -3,7 +3,15 @@ import warnings from typing import Any, Optional -from oa_configurator import qualified, schema_inspect +from oa_configurator import ( + ResolvedDatabase, + Role, + ensure_schema, + guard_schema_provenance, + qualified, + schema_inspect, + supports_schemas, +) from sqlalchemy import ( DateTime, Engine, @@ -19,11 +27,17 @@ from omop_llm import supported_providers from omop_emb.config import ( + MODEL_REGISTRY_SCHEMA, IndexType, MetricType, ) from omop_emb.backends.index_config import IndexConfig +# Schema name for the model registry table. Dialects with schema support +# store the registry table in a dedicated schema to allow schema-independent +# access to the registry table from any schema in the same database. +REGISTRY_SCHEMA_KEY = "registry" + class ModelRegistryBase(DeclarativeBase): """Dedicated declarative base for local model registry metadata.""" @@ -79,6 +93,7 @@ class ModelRegistry(ModelRegistryBase): """ __tablename__ = "model_registry" + __table_args__ = {"schema": REGISTRY_SCHEMA_KEY} model_name = mapped_column(String, primary_key=True) @@ -176,15 +191,34 @@ 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 _registry_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. 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") or {}), + REGISTRY_SCHEMA_KEY: _registry_schema(connection), + } + ) + ensure_schema(connection, MODEL_REGISTRY_SCHEMA) + with guard_schema_provenance(connection, resolved, role=Role.PRIMARY): + ModelRegistryBase.metadata.create_all(connection, tables=[ModelRegistry.__table__]) # ty: ignore[invalid-argument-type] _migrate_legacy_provider_type_column(engine) @@ -200,7 +234,7 @@ 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 = schema_inspect(engine).get_columns(ModelRegistry.__tablename__) + columns = schema_inspect(engine, schema=_registry_schema(engine)).get_columns(ModelRegistry.__tablename__) provider_column = next( (column for column in columns if column["name"] == "provider_type"), None, @@ -210,6 +244,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: legacy_length = getattr(provider_column["type"], "length", None) with engine.begin() as connection: + registry_schema = _registry_schema(connection) if engine.dialect.name == "postgresql" and legacy_length is not None: warnings.warn( "Widening a legacy fixed-length provider_type column. This " @@ -220,7 +255,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: ) connection.execute( text( - f"ALTER TABLE {qualified(connection, ModelRegistry.__tablename__)} " + f"ALTER TABLE {qualified(connection, ModelRegistry.__tablename__, schema=registry_schema)} " "ALTER COLUMN provider_type TYPE VARCHAR " "USING provider_type::text" ) @@ -230,7 +265,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: # legacy partial table (provider_type only) doesn't have. connection.execute( text( - f"UPDATE {qualified(connection, ModelRegistry.__tablename__)} " + f"UPDATE {qualified(connection, ModelRegistry.__tablename__, schema=registry_schema)} " "SET provider_type = lower(provider_type) " "WHERE provider_type IS NOT NULL " "AND provider_type <> lower(provider_type)" diff --git a/tests/test_pgvector.py b/tests/test_pgvector.py index bfb2f01..388d4ed 100644 --- a/tests/test_pgvector.py +++ b/tests/test_pgvector.py @@ -203,7 +203,7 @@ def test_table_and_index_lifecycle_stays_in_the_configured_schema( backend.delete_model(model_name=MODEL_NAME) assert backend._storage_table_exists(record) is False - def test_model_registry_lookup_stays_in_the_configured_schema( + def test_model_registry_lives_in_its_own_reserved_schema( self, scoped_backend, pg_engine ): backend, schema = scoped_backend @@ -219,13 +219,12 @@ def test_model_registry_lookup_stays_in_the_configured_schema( assert registry.registry_available is True assert len(registry.get_registered_models(model_name=MODEL_NAME)) == 1 - # A registry pointed at "public" must not see it: proves the lookup is - # genuinely schema-scoped, not incidentally finding it via search_path. + # A registry built with a completely different None-key schema sees + # the exact same row: the registry is decoupled from it entirely. public_engine = pg_engine.execution_options(schema_translate_map={None: "public"}) public_registry = RegistryManager.read_only(public_engine) - found_in_public = ( - public_registry.get_registered_models(model_name=MODEL_NAME) - if public_registry.registry_available - else () - ) - assert found_in_public == () + 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 schema_inspect(pg_engine, schema=schema).has_table("model_registry") is False diff --git a/tests/test_read_only_population.py b/tests/test_read_only_population.py index 71863b8..de40615 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 @@ -46,7 +46,7 @@ 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", @@ -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 4590efe..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 @@ -245,33 +247,37 @@ def test_legacy_provider_name_is_normalized_in_sqlite(svec_engine): @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) From 01a886f6e4a3ab49b62b16c546dfe4369a44d67d Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Mon, 7 Sep 2026 05:17:16 +0000 Subject: [PATCH 06/16] Incorporate dialect and schema_translate_map changes and keys --- src/omop_emb/backends/base_backend.py | 27 ++++++++++++------- src/omop_emb/backends/db_utils.py | 5 ++-- src/omop_emb/backends/pgvector/pg_backend.py | 4 +-- src/omop_emb/backends/pgvector/pg_sql.py | 8 +++--- src/omop_emb/backends/read_only.py | 8 +++--- .../backends/sqlitevec/sqlitevec_backend.py | 4 +-- .../backends/sqlitevec/sqlitevec_sql.py | 11 ++++---- .../model_registry/model_registry_manager.py | 4 +-- .../model_registry/model_registry_orm.py | 6 +++-- tests/conftest.py | 2 +- 10 files changed, 46 insertions(+), 33 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index 14e4cf4..ad198aa 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -6,7 +6,13 @@ 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, supports_schemas +from oa_configurator import ( + SCHEMA_TRANSLATE_MAP_KEY, + Dialect, + ResolvedDatabase, + ResolvedVectorStore, + supports_schemas, +) from sqlalchemy import Engine from sqlalchemy.engine import make_url from sqlalchemy.orm import sessionmaker @@ -1044,27 +1050,28 @@ 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 -- added on top of database's own - # translate map, not passed as a bare override, since create_engine()'s - # execution_options replaces the whole map rather than merging it. + # independent of database's own schema. create_engine() merges this extra + # key onto its own configured map rather than replacing it. registry_schema = MODEL_REGISTRY_SCHEMA if supports_schemas(database.connection.dialect_name) else None - schema_translate_map = {**database.schema_translate_map(), REGISTRY_SCHEMA_KEY: registry_schema} + registry_schema_translate_map = {REGISTRY_SCHEMA_KEY: registry_schema} if resolved_backend == BackendType.SQLITEVEC: 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}." ) emb_engine = create_sqlitevec_engine( - database.create_engine(execution_options={"schema_translate_map": schema_translate_map}) + 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}." @@ -1076,7 +1083,9 @@ def resolve_backend( "pgvector backend is not installed. " "Install it with: pip install omop-emb[pgvector]" ) from exc - emb_engine = database.create_engine(execution_options={"schema_translate_map": schema_translate_map}) + 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, resolved=database) 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/pgvector/pg_backend.py b/src/omop_emb/backends/pgvector/pg_backend.py index c821948..b3114f6 100644 --- a/src/omop_emb/backends/pgvector/pg_backend.py +++ b/src/omop_emb/backends/pgvector/pg_backend.py @@ -5,7 +5,7 @@ from typing import Mapping, Optional, Sequence, Tuple from numpy import ndarray -from oa_configurator import ResolvedDatabase +from oa_configurator import Dialect, ResolvedDatabase from sqlalchemy import Engine, select, text try: @@ -98,7 +98,7 @@ def backend_type(self) -> BackendType: @property def dialect(self) -> str: - return "postgresql" + return Dialect.POSTGRESQL # ------------------------------------------------------------------ # Store lifecycle diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index f081efd..af6b797 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -15,7 +15,7 @@ from typing import List, Optional, Sequence, Union from numpy import ndarray -from oa_configurator import ResolvedDatabase, Role, guard_schema_provenance, qualified, schema_inspect +from oa_configurator import Dialect, ResolvedDatabase, Role, guard_schema_provenance, qualified, schema_inspect 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 @@ -185,7 +185,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. @@ -273,7 +273,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) @@ -287,7 +287,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. diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index dda5197..2557d6b 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -12,7 +12,7 @@ from collections.abc import Iterator from dataclasses import dataclass -from oa_configurator import ResolvedVectorStore +from oa_configurator import Dialect, ResolvedVectorStore from sqlalchemy import Engine, event, inspect, select from omop_emb.backends.base_backend import ( @@ -97,7 +97,7 @@ def iter_stored_embeddings( return schema = ( None - if self._engine.dialect.name == "sqlite" and self.schema == "main" + if self._engine.dialect.name == Dialect.SQLITE and self.schema == "main" else self.schema ) table = concept_metadata_table_descriptor( @@ -128,7 +128,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: @@ -170,7 +170,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 diff --git a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py index 9a75685..6d92d3c 100644 --- a/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py +++ b/src/omop_emb/backends/sqlitevec/sqlitevec_backend.py @@ -11,7 +11,7 @@ import numpy as np from numpy import ndarray -from oa_configurator import ResolvedDatabase +from oa_configurator import Dialect, ResolvedDatabase from sqlalchemy import Engine, MetaData, Table, event, text try: @@ -112,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/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index 0a2bcc6..45ed1dd 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -5,7 +5,7 @@ from datetime import datetime, timezone from typing import Mapping, Optional -from oa_configurator import ResolvedDatabase, schema_inspect +from oa_configurator import SCHEMA_TRANSLATE_MAP_KEY, ResolvedDatabase, schema_inspect from sqlalchemy import Engine, select, update from sqlalchemy.orm import Session, sessionmaker @@ -53,7 +53,7 @@ def __init__( ) -> None: self._embedding_engine = embedding_engine.execution_options( schema_translate_map={ - **(embedding_engine.get_execution_options().get("schema_translate_map") or {}), + **(embedding_engine.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), REGISTRY_SCHEMA_KEY: _registry_schema(embedding_engine), } ) diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index 2b0f199..e93df10 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -4,6 +4,8 @@ from typing import Any, Optional from oa_configurator import ( + SCHEMA_TRANSLATE_MAP_KEY, + Dialect, ResolvedDatabase, Role, ensure_schema, @@ -212,7 +214,7 @@ def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = with engine.begin() as connection: connection = connection.execution_options( schema_translate_map={ - **(connection.get_execution_options().get("schema_translate_map") or {}), + **(connection.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), REGISTRY_SCHEMA_KEY: _registry_schema(connection), } ) @@ -245,7 +247,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: legacy_length = getattr(provider_column["type"], "length", None) with engine.begin() as connection: registry_schema = _registry_schema(connection) - if engine.dialect.name == "postgresql" and legacy_length is not None: + 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 " diff --git a/tests/conftest.py b/tests/conftest.py index b896b0f..b019ee8 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -109,7 +109,7 @@ def pg_engine(pg_db) -> sa.Engine: 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.connection.engine + return pg_db.committing_engine @pytest.fixture From b3fc640f6e72e01dc8a8191221cb6b5e7765c22e Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Tue, 15 Sep 2026 00:17:47 +0000 Subject: [PATCH 07/16] Properly pass down schema, tag tables with schema, share the registry schema across databases --- src/omop_emb/backends/base_backend.py | 5 +- .../backends/pgvector/pg_index_manager.py | 5 +- src/omop_emb/backends/pgvector/pg_sql.py | 15 +++- src/omop_emb/backends/read_only.py | 15 ++-- .../model_registry/model_registry_orm.py | 22 +++++- tests/test_pgvector.py | 19 +++-- tests/test_pgvector_index_manager.py | 44 +++++++----- tests/test_schema_provenance_guard.py | 71 ++++++++++++++++--- 8 files changed, 144 insertions(+), 52 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index ad198aa..15401b1 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -11,14 +11,12 @@ Dialect, ResolvedDatabase, ResolvedVectorStore, - supports_schemas, ) from sqlalchemy import Engine from sqlalchemy.engine import make_url from sqlalchemy.orm import sessionmaker from omop_emb.config import ( - MODEL_REGISTRY_SCHEMA, BackendType, MetricType, IndexType, @@ -30,6 +28,7 @@ from omop_emb.backends.embedding_table import ConceptEmbeddingRecord from omop_emb.backends.index_config import IndexConfig, FlatIndexConfig from omop_emb.model_registry import EmbeddingModelRecord, REGISTRY_SCHEMA_KEY, RegistryManager +from omop_emb.model_registry.model_registry_orm import _registry_schema from omop_emb.utils.embedding_utils import ( EmbeddingConceptFilter, NearestConceptMatch, @@ -1052,7 +1051,7 @@ def resolve_backend( # 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 = MODEL_REGISTRY_SCHEMA if supports_schemas(database.connection.dialect_name) else None + registry_schema = _registry_schema(database.connection.dialect_name) registry_schema_translate_map = {REGISTRY_SCHEMA_KEY: registry_schema} if resolved_backend == BackendType.SQLITEVEC: diff --git a/src/omop_emb/backends/pgvector/pg_index_manager.py b/src/omop_emb/backends/pgvector/pg_index_manager.py index 9939689..0bb5b89 100644 --- a/src/omop_emb/backends/pgvector/pg_index_manager.py +++ b/src/omop_emb/backends/pgvector/pg_index_manager.py @@ -191,10 +191,7 @@ 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 - # self._engine is None only in pure-DDL-string unit tests that never - # open a connection; qualified() needs a real bindable for its - # schema, so fall back to the bare name in that case only. - table_ref = qualified(self._engine, self._tablename) if self._engine is not None else self._tablename + table_ref = qualified(self._engine, self._tablename) return ( f"CREATE INDEX {self._index_name(metric_type)} " f"ON {table_ref} " diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index af6b797..fd07e34 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -15,7 +15,16 @@ from typing import List, Optional, Sequence, Union from numpy import ndarray -from oa_configurator import Dialect, ResolvedDatabase, Role, guard_schema_provenance, qualified, schema_inspect +from oa_configurator import ( + Dialect, + ResolvedDatabase, + Role, + ensure_schema, + guard_schema_provenance, + qualified, + schema_inspect, + 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 @@ -66,6 +75,8 @@ def create_pg_embedding_table( """ table_cls = pg_embedding_table_descriptor(model_record) with engine.begin() as connection: + # Ensure that the schema exists before creating the table + ensure_schema(connection, schema_of(connection, role=Role.PRIMARY)) with guard_schema_provenance(connection, resolved, role=Role.PRIMARY): EmbeddingTableBase.metadata.create_all(connection, tables=[table_cls.__table__]) # ty: ignore[invalid-argument-type] return table_cls @@ -383,7 +394,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 2557d6b..f5e6e61 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -12,8 +12,8 @@ from collections.abc import Iterator from dataclasses import dataclass -from oa_configurator import Dialect, ResolvedVectorStore -from sqlalchemy import Engine, event, inspect, select +from oa_configurator import Dialect, ResolvedVectorStore, qualified, schema_inspect +from sqlalchemy import Engine, event, select from omop_emb.backends.base_backend import ( EmbeddingBackend, @@ -92,8 +92,8 @@ def iter_stored_embeddings( record = self.model(model_name) if record is None: return - inspector = inspect(self._engine) - if not inspector.has_table(record.storage_identifier, schema=self.schema): + inspector = schema_inspect(self._engine, schema=self.schema) + if not inspector.has_table(record.storage_identifier): return schema = ( None @@ -136,9 +136,8 @@ def physical_indexes(self, model_name: str) -> tuple[str, ...]: expected_prefix = f"idx_{record.storage_identifier}_" return tuple( str(item["name"]) - for item in inspect(self._engine).get_indexes( + for item in schema_inspect(self._engine, schema=self.schema).get_indexes( record.storage_identifier, - schema=self.schema, ) if str(item["name"]).startswith(expected_prefix) ) @@ -146,10 +145,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, schema=self.schema)};" for name in self.physical_indexes(model_name) ) diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index e93df10..41c80ec 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -1,13 +1,13 @@ 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, - Role, ensure_schema, guard_schema_provenance, qualified, @@ -200,6 +200,16 @@ def _registry_schema(bindable) -> str | 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 + ----- + The registry table is shared across all resolved databases that share the same connection + (e.g. a single Postgres instance with CDM and Vector Store). Without specifying + `shared_as=MODEL_REGISTRY_SCHEMA`, each entry's own resolved.name would be tracked as a + separate identity yet the schema already exists in the database. A second configured + database with the same connection would raise a false-positive finding of "already populated". + Parameters ---------- @@ -219,7 +229,15 @@ def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = } ) ensure_schema(connection, MODEL_REGISTRY_SCHEMA) - with guard_schema_provenance(connection, resolved, role=Role.PRIMARY): + registry_schema = _registry_schema(connection) + guard = ( + guard_schema_provenance( + connection, resolved, role=registry_schema, shared_as=MODEL_REGISTRY_SCHEMA + ) + if 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) diff --git a/tests/test_pgvector.py b/tests/test_pgvector.py index 388d4ed..aeacf47 100644 --- a/tests/test_pgvector.py +++ b/tests/test_pgvector.py @@ -13,7 +13,7 @@ "pgvector", reason="omop-emb[pgvector] not installed: skipping pgvector tests" ) -from oa_configurator import schema_inspect +from oa_configurator import Role, schema_inspect from oa_configurator.testing import isolated_test_schema from omop_emb.backends.index_config import FlatIndexConfig, HNSWIndexConfig @@ -146,7 +146,8 @@ def test_rebuild_index(self, pg_backend: PGVectorEmbeddingBackend): 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.""" + 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 @@ -156,7 +157,7 @@ class TestPGVectorNonDefaultSchema: 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={None: schema} + schema_translate_map={Role.PRIMARY.value: schema} ) backend = PGVectorEmbeddingBackend(emb_engine=scoped_engine) yield backend, schema @@ -214,14 +215,18 @@ def test_model_registry_lives_in_its_own_reserved_schema( dimensions=EMBEDDING_DIM, ) - scoped_engine = pg_engine.execution_options(schema_translate_map={None: schema}) + 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 None-key schema sees - # the exact same row: the registry is decoupled from it entirely. - public_engine = pg_engine.execution_options(schema_translate_map={None: "public"}) + # A registry built with a completely different primary-role 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 diff --git a/tests/test_pgvector_index_manager.py b/tests/test_pgvector_index_manager.py index 5f1a9ec..6468c64 100644 --- a/tests/test_pgvector_index_manager.py +++ b/tests/test_pgvector_index_manager.py @@ -31,6 +31,18 @@ ) +@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: @@ -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( diff --git a/tests/test_schema_provenance_guard.py b/tests/test_schema_provenance_guard.py index 94e0e12..5806874 100644 --- a/tests/test_schema_provenance_guard.py +++ b/tests/test_schema_provenance_guard.py @@ -19,11 +19,12 @@ import uuid import pytest -from oa_configurator import SchemaDriftError +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 pytestmark = [pytest.mark.postgresql, pytest.mark.db_dialect] @@ -41,22 +42,74 @@ def _resolved(pg_db, *, database_name: str, schema: str): ) -def test_backend_construction_guard_fires_on_reconfigured_schema(pg_db, pg_engine, cleanup_after_test): - database_name = f"emb_guard_db_{uuid.uuid4().hex[:8]}" +def _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) -> None: + """The registry's own provenance row is keyed by shared_as=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, + pg_db.resolved, + role=MODEL_REGISTRY_SCHEMA, + new_schema=MODEL_REGISTRY_SCHEMA, + reason="test setup: establish a known-correct baseline", + shared_as=MODEL_REGISTRY_SCHEMA, + ) delete_rows_on_cleanup( - cleanup_after_test, pg_engine, table, table.c.database_name == database_name + 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={None: 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 - resolved_b = _resolved(pg_db, database_name=database_name, schema=schema_b) - engine_b = pg_engine.execution_options(schema_translate_map={None: schema_b}) - with pytest.raises(SchemaDriftError): - PGVectorEmbeddingBackend(emb_engine=engine_b, resolved=resolved_b) + # 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, + resolved, + role=MODEL_REGISTRY_SCHEMA, + new_schema="a_previous_registry_schema_that_is_not_current", + reason="test: force a stale baseline", + shared_as=MODEL_REGISTRY_SCHEMA, + ) + + engine = pg_engine.execution_options(schema_translate_map={"primary": "unrelated_primary_schema"}) + with pytest.raises(SchemaDriftError): + PGVectorEmbeddingBackend(emb_engine=engine, resolved=resolved) From fb4cfbcc42996cb4b25e1cefe8d13b5e8c9792c9 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Tue, 15 Sep 2026 04:37:18 +0000 Subject: [PATCH 08/16] remove underscore-prefixing for public API method --- src/omop_emb/backends/base_backend.py | 10 +++++++--- src/omop_emb/model_registry/__init__.py | 2 ++ src/omop_emb/model_registry/model_registry_manager.py | 6 +++--- src/omop_emb/model_registry/model_registry_orm.py | 10 +++++----- 4 files changed, 17 insertions(+), 11 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index 15401b1..c0bbd66 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -27,8 +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, REGISTRY_SCHEMA_KEY, RegistryManager -from omop_emb.model_registry.model_registry_orm import _registry_schema +from omop_emb.model_registry import ( + EmbeddingModelRecord, + REGISTRY_SCHEMA_KEY, + RegistryManager, + resolve_registry_schema, +) from omop_emb.utils.embedding_utils import ( EmbeddingConceptFilter, NearestConceptMatch, @@ -1051,7 +1055,7 @@ def resolve_backend( # 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 = _registry_schema(database.connection.dialect_name) + registry_schema = resolve_registry_schema(database.connection.dialect_name) registry_schema_translate_map = {REGISTRY_SCHEMA_KEY: registry_schema} if resolved_backend == BackendType.SQLITEVEC: diff --git a/src/omop_emb/model_registry/__init__.py b/src/omop_emb/model_registry/__init__.py index d61ac02..231b0d4 100644 --- a/src/omop_emb/model_registry/__init__.py +++ b/src/omop_emb/model_registry/__init__.py @@ -4,6 +4,7 @@ REGISTRY_SCHEMA_KEY, ModelRegistry, ensure_registry_table, + resolve_registry_schema, ) __all__ = [ @@ -12,4 +13,5 @@ "ModelRegistry", "ensure_registry_table", "REGISTRY_SCHEMA_KEY", + "resolve_registry_schema", ] diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index 45ed1dd..b480e7c 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -17,8 +17,8 @@ from omop_emb.model_registry.model_registry_orm import ( REGISTRY_SCHEMA_KEY, ModelRegistry, - _registry_schema, ensure_registry_table, + resolve_registry_schema, ) from omop_emb.model_registry.model_registry_types import EmbeddingModelRecord from omop_emb.utils.errors import ModelRegistrationConflictError @@ -54,13 +54,13 @@ def __init__( self._embedding_engine = embedding_engine.execution_options( schema_translate_map={ **(embedding_engine.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), - REGISTRY_SCHEMA_KEY: _registry_schema(embedding_engine), + REGISTRY_SCHEMA_KEY: resolve_registry_schema(embedding_engine), } ) self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize self._registry_available = schema_inspect( - self._embedding_engine, schema=_registry_schema(self._embedding_engine) + self._embedding_engine, schema=resolve_registry_schema(self._embedding_engine) ).has_table(ModelRegistry.__tablename__) if initialize: ensure_registry_table(self._embedding_engine, resolved=resolved) diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index 41c80ec..0b9799f 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -193,7 +193,7 @@ def _validate_and_sync_index_config( return index_config.to_dict() -def _registry_schema(bindable) -> str | None: +def resolve_registry_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 @@ -225,11 +225,11 @@ def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = connection = connection.execution_options( schema_translate_map={ **(connection.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), - REGISTRY_SCHEMA_KEY: _registry_schema(connection), + REGISTRY_SCHEMA_KEY: resolve_registry_schema(connection), } ) ensure_schema(connection, MODEL_REGISTRY_SCHEMA) - registry_schema = _registry_schema(connection) + registry_schema = resolve_registry_schema(connection) guard = ( guard_schema_provenance( connection, resolved, role=registry_schema, shared_as=MODEL_REGISTRY_SCHEMA @@ -254,7 +254,7 @@ 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 = schema_inspect(engine, schema=_registry_schema(engine)).get_columns(ModelRegistry.__tablename__) + columns = schema_inspect(engine, schema=resolve_registry_schema(engine)).get_columns(ModelRegistry.__tablename__) provider_column = next( (column for column in columns if column["name"] == "provider_type"), None, @@ -264,7 +264,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: legacy_length = getattr(provider_column["type"], "length", None) with engine.begin() as connection: - registry_schema = _registry_schema(connection) + registry_schema = resolve_registry_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 " From 74e783d03381435f877162ec53e87588fcf9f9dc Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Tue, 15 Sep 2026 04:40:47 +0000 Subject: [PATCH 09/16] Increase test coverage --- tests/test_resolve_backend.py | 56 +++++++++++++++++++++++++++ tests/test_schema_provenance_guard.py | 34 ++++++++++++++++ 2 files changed, 90 insertions(+) create mode 100644 tests/test_resolve_backend.py 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 index 5806874..85ce80c 100644 --- a/tests/test_schema_provenance_guard.py +++ b/tests/test_schema_provenance_guard.py @@ -26,6 +26,8 @@ from omop_emb.backends.pgvector.pg_backend import PGVectorEmbeddingBackend from omop_emb.config import MODEL_REGISTRY_SCHEMA +from .conftest import EMBEDDING_DIM, MODEL_NAME, PROVIDER_TYPE + pytestmark = [pytest.mark.postgresql, pytest.mark.db_dialect] @@ -113,3 +115,35 @@ def test_registry_schema_guard_fires_on_genuine_registry_drift(pg_db, pg_engine, 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-role 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) From d1a072aaca69bdeba59a6721ec10f91aa57b1ca8 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Wed, 16 Sep 2026 03:05:15 +0000 Subject: [PATCH 10/16] Correctly raise on an unsupporter backend, include tests for it --- src/omop_emb/config.py | 34 ++++++++++++++++++++++++----- tests/test_config.py | 49 +++++++++++++++++++++++++++++++++++++----- 2 files changed, 73 insertions(+), 10 deletions(-) diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index 392075f..fc62a43 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -271,6 +271,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: @@ -286,8 +311,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: @@ -302,7 +326,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( @@ -318,7 +342,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, ...]: @@ -333,6 +357,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/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) From cbbab3f0f7048f8ab14ecf18f6f3ffaba537d2dc Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Wed, 23 Sep 2026 00:57:55 +0000 Subject: [PATCH 11/16] Remove Role following oa-configurator changes --- .../backends/pgvector/pg_index_manager.py | 15 ++++++---- src/omop_emb/backends/pgvector/pg_sql.py | 17 +++++++---- src/omop_emb/backends/read_only.py | 14 ++++----- .../model_registry/model_registry_manager.py | 10 +++---- .../model_registry/model_registry_orm.py | 29 ++++++++++--------- tests/test_pgvector.py | 15 +++++----- tests/test_schema_provenance_guard.py | 21 +++++++------- 7 files changed, 67 insertions(+), 54 deletions(-) diff --git a/src/omop_emb/backends/pgvector/pg_index_manager.py b/src/omop_emb/backends/pgvector/pg_index_manager.py index 0bb5b89..4ac8f0a 100644 --- a/src/omop_emb/backends/pgvector/pg_index_manager.py +++ b/src/omop_emb/backends/pgvector/pg_index_manager.py @@ -17,8 +17,8 @@ import logging from typing import Generic, TypeVar -from oa_configurator import qualified, schema_inspect -from sqlalchemy import Engine, text +from oa_configurator import qualified, schema_of +from sqlalchemy import Engine, inspect, text from omop_emb.config import IndexType, MetricType, VectorColumnType from omop_emb.backends.index_config import IndexConfig, FlatIndexConfig, HNSWIndexConfig @@ -78,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 schema_inspect(conn).get_indexes(self._tablename) + idx["name"] + for idx in inspect(conn).get_indexes( + self._tablename, schema=schema_of(conn) + ) } return self._index_name(metric_type) in existing @@ -102,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 {qualified(conn, name)}")) + conn.execute( + text(f"DROP INDEX IF EXISTS {qualified(conn, name, physical_schema=schema_of(conn))}") + ) if existed: logger.info(f"Dropped pgvector index '{name}'.") @@ -191,7 +196,7 @@ 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) + table_ref = qualified(self._engine, self._tablename, physical_schema=schema_of(self._engine)) return ( f"CREATE INDEX {self._index_name(metric_type)} " f"ON {table_ref} " diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index fd07e34..63a0378 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -20,9 +20,8 @@ ResolvedDatabase, Role, ensure_schema, - guard_schema_provenance, + guard_schema_provenance_for, qualified, - schema_inspect, schema_of, ) from sqlalchemy import Engine, Integer, Row, Select, func, inspect as sa_inspect, literal, select, text, TextClause @@ -42,7 +41,7 @@ def table_exists(engine: Engine, table_name: str) -> bool: """Return ``True`` if *table_name* exists in the engine's configured schema.""" - return schema_inspect(engine).has_table(table_name) + return sa_inspect(engine).has_table(table_name, schema=schema_of(engine)) def create_pg_embedding_table( engine: Engine, @@ -76,8 +75,12 @@ def create_pg_embedding_table( table_cls = pg_embedding_table_descriptor(model_record) with engine.begin() as connection: # Ensure that the schema exists before creating the table - ensure_schema(connection, schema_of(connection, role=Role.PRIMARY)) - with guard_schema_provenance(connection, resolved, role=Role.PRIMARY): + physical_schema = schema_of(connection, schema_tag=Role.PRIMARY) + ensure_schema(connection, physical_schema) + guard = guard_schema_provenance_for( + connection, resolved, role=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 @@ -92,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 {qualified(conn, tablename)}")) + conn.execute( + text(f"DROP TABLE IF EXISTS {qualified(conn, tablename, physical_schema=schema_of(conn))}") + ) logger.info(f"Dropped embedding table '{tablename}'.") diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index f5e6e61..079306b 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -12,8 +12,8 @@ from collections.abc import Iterator from dataclasses import dataclass -from oa_configurator import Dialect, ResolvedVectorStore, qualified, schema_inspect -from sqlalchemy import Engine, event, select +from oa_configurator import Dialect, ResolvedVectorStore, qualified +from sqlalchemy import Engine, event, inspect, select from omop_emb.backends.base_backend import ( EmbeddingBackend, @@ -92,8 +92,8 @@ def iter_stored_embeddings( record = self.model(model_name) if record is None: return - inspector = schema_inspect(self._engine, schema=self.schema) - if not inspector.has_table(record.storage_identifier): + inspector = inspect(self._engine) + if not inspector.has_table(record.storage_identifier, schema=self.schema): return schema = ( None @@ -136,8 +136,8 @@ def physical_indexes(self, model_name: str) -> tuple[str, ...]: expected_prefix = f"idx_{record.storage_identifier}_" return tuple( str(item["name"]) - for item in schema_inspect(self._engine, schema=self.schema).get_indexes( - record.storage_identifier, + for item in inspect(self._engine).get_indexes( + record.storage_identifier, schema=self.schema, ) if str(item["name"]).startswith(expected_prefix) ) @@ -146,7 +146,7 @@ def drop_index_sql(self, model_name: str) -> tuple[str, ...]: """Return reviewed index-removal statements without executing them.""" return tuple( - f"DROP INDEX IF EXISTS {qualified(self._engine, name, schema=self.schema)};" + f"DROP INDEX IF EXISTS {qualified(self._engine, name, physical_schema=self.schema)};" for name in self.physical_indexes(model_name) ) diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index b480e7c..3108115 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -5,8 +5,8 @@ from datetime import datetime, timezone from typing import Mapping, Optional -from oa_configurator import SCHEMA_TRANSLATE_MAP_KEY, ResolvedDatabase, schema_inspect -from sqlalchemy import Engine, select, update +from oa_configurator import SCHEMA_TRANSLATE_MAP_KEY, ResolvedDatabase +from sqlalchemy import Engine, inspect, select, update from sqlalchemy.orm import Session, sessionmaker from omop_emb.backends.index_config import ( @@ -59,9 +59,9 @@ def __init__( ) self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize - self._registry_available = schema_inspect( - self._embedding_engine, schema=resolve_registry_schema(self._embedding_engine) - ).has_table(ModelRegistry.__tablename__) + self._registry_available = inspect(self._embedding_engine).has_table( + ModelRegistry.__tablename__, schema=resolve_registry_schema(self._embedding_engine) + ) if initialize: ensure_registry_table(self._embedding_engine, resolved=resolved) self._registry_available = True diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index 0b9799f..c6b57b0 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -11,7 +11,6 @@ ensure_schema, guard_schema_provenance, qualified, - schema_inspect, supports_schemas, ) from sqlalchemy import ( @@ -22,6 +21,7 @@ JSON, String, func, + inspect, text, ) from sqlalchemy.orm import DeclarativeBase, mapped_column, validates, Mapped @@ -200,16 +200,12 @@ def resolve_registry_schema(bindable) -> str | 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 ----- - The registry table is shared across all resolved databases that share the same connection - (e.g. a single Postgres instance with CDM and Vector Store). Without specifying - `shared_as=MODEL_REGISTRY_SCHEMA`, each entry's own resolved.name would be tracked as a - separate identity yet the schema already exists in the database. A second configured - database with the same connection would raise a false-positive finding of "already populated". - + 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 ---------- @@ -232,9 +228,14 @@ def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = registry_schema = resolve_registry_schema(connection) guard = ( guard_schema_provenance( - connection, resolved, role=registry_schema, shared_as=MODEL_REGISTRY_SCHEMA + 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 registry_schema is not None + if resolved is not None and registry_schema is not None else nullcontext() ) with guard: @@ -254,7 +255,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 = schema_inspect(engine, schema=resolve_registry_schema(engine)).get_columns(ModelRegistry.__tablename__) + columns = inspect(engine).get_columns( + ModelRegistry.__tablename__, schema=resolve_registry_schema(engine) + ) provider_column = next( (column for column in columns if column["name"] == "provider_type"), None, @@ -275,7 +278,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: ) connection.execute( text( - f"ALTER TABLE {qualified(connection, ModelRegistry.__tablename__, schema=registry_schema)} " + f"ALTER TABLE {qualified(connection, ModelRegistry.__tablename__, physical_schema=registry_schema)} " "ALTER COLUMN provider_type TYPE VARCHAR " "USING provider_type::text" ) @@ -285,7 +288,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: # legacy partial table (provider_type only) doesn't have. connection.execute( text( - f"UPDATE {qualified(connection, ModelRegistry.__tablename__, schema=registry_schema)} " + 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/tests/test_pgvector.py b/tests/test_pgvector.py index aeacf47..43d8d09 100644 --- a/tests/test_pgvector.py +++ b/tests/test_pgvector.py @@ -13,7 +13,8 @@ "pgvector", reason="omop-emb[pgvector] not installed: skipping pgvector tests" ) -from oa_configurator import Role, schema_inspect +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 @@ -182,8 +183,8 @@ def test_table_and_index_lifecycle_stays_in_the_configured_schema( # 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 schema_inspect(pg_engine, schema="public").has_table( - record.storage_identifier + 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 @@ -191,8 +192,8 @@ def test_table_and_index_lifecycle_stays_in_the_configured_schema( 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 = schema_inspect(pg_engine, schema=schema).get_indexes( - record.storage_identifier + 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 @@ -222,7 +223,7 @@ def test_model_registry_lives_in_its_own_reserved_schema( 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-role schema + # 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"} @@ -232,4 +233,4 @@ def test_model_registry_lives_in_its_own_reserved_schema( assert len(public_registry.get_registered_models(model_name=MODEL_NAME)) == 1 # And genuinely never created under the storage schema's own name. - assert schema_inspect(pg_engine, schema=schema).has_table("model_registry") is False + assert sa.inspect(pg_engine).has_table("model_registry", schema=schema) is False diff --git a/tests/test_schema_provenance_guard.py b/tests/test_schema_provenance_guard.py index 85ce80c..b3efd85 100644 --- a/tests/test_schema_provenance_guard.py +++ b/tests/test_schema_provenance_guard.py @@ -25,6 +25,7 @@ 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 @@ -45,9 +46,9 @@ def _resolved(pg_db, *, database_name: str, schema: str): def _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) -> None: - """The registry's own provenance row is keyed by shared_as=MODEL_REGISTRY_SCHEMA + """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". + 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. """ @@ -55,11 +56,10 @@ def _establish_registry_baseline(pg_db, pg_engine, cleanup_after_test) -> None: with pg_engine.begin() as connection: record_schema_provenance( connection, - pg_db.resolved, - role=MODEL_REGISTRY_SCHEMA, - new_schema=MODEL_REGISTRY_SCHEMA, + database_name=MODEL_REGISTRY_SCHEMA, + schema_tag=REGISTRY_SCHEMA_KEY, + new_physical_schema=MODEL_REGISTRY_SCHEMA, reason="test setup: establish a known-correct baseline", - shared_as=MODEL_REGISTRY_SCHEMA, ) delete_rows_on_cleanup( cleanup_after_test, pg_engine, table, table.c.database_name == MODEL_REGISTRY_SCHEMA @@ -105,11 +105,10 @@ def test_registry_schema_guard_fires_on_genuine_registry_drift(pg_db, pg_engine, with pg_engine.begin() as connection: record_schema_provenance( connection, - resolved, - role=MODEL_REGISTRY_SCHEMA, - new_schema="a_previous_registry_schema_that_is_not_current", + 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", - shared_as=MODEL_REGISTRY_SCHEMA, ) engine = pg_engine.execution_options(schema_translate_map={"primary": "unrelated_primary_schema"}) @@ -120,7 +119,7 @@ def test_registry_schema_guard_fires_on_genuine_registry_drift(pg_db, pg_engine, 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-role drift coverage a prior rewrite replaced instead + 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) From 523142a8426113605e73b16ab046290f8e56b58e Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Wed, 23 Sep 2026 04:29:25 +0000 Subject: [PATCH 12/16] Follow-up from internal review --- src/omop_emb/backends/pgvector/pg_sql.py | 2 +- src/omop_emb/backends/read_only.py | 13 +++++++------ src/omop_emb/config.py | 2 ++ .../model_registry/model_registry_manager.py | 18 ++++++++++++------ .../model_registry/model_registry_orm.py | 6 +----- 5 files changed, 23 insertions(+), 18 deletions(-) diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index 63a0378..fb2f23a 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -78,7 +78,7 @@ def create_pg_embedding_table( physical_schema = schema_of(connection, schema_tag=Role.PRIMARY) ensure_schema(connection, physical_schema) guard = guard_schema_provenance_for( - connection, resolved, role=Role.PRIMARY, tables=[table_cls.__table__] # ty: ignore[invalid-argument-type] + 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] diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index 079306b..9c5b373 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 Dialect, ResolvedVectorStore, qualified +from oa_configurator import ( + Dialect, + ResolvedVectorStore, + qualified, + schema_if_supported +) from sqlalchemy import Engine, event, inspect, select from omop_emb.backends.base_backend import ( @@ -95,11 +100,7 @@ def iter_stored_embeddings( inspector = inspect(self._engine) if not inspector.has_table(record.storage_identifier, schema=self.schema): return - schema = ( - None - if self._engine.dialect.name == Dialect.SQLITE and self.schema == "main" - else self.schema - ) + schema = schema_if_supported(self.schema, self._engine) table = concept_metadata_table_descriptor( record.storage_identifier, schema=schema, diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index fc62a43..3e2a6d4 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -21,8 +21,10 @@ # 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(REGISTRY_SCHEMA_KEY, owner="omop_emb") class OmopEmbConfig(PackageConfigBase): diff --git a/src/omop_emb/model_registry/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index 3108115..8de1c67 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -51,12 +51,18 @@ def __init__( initialize: bool = True, resolved: ResolvedDatabase | None = None, ) -> None: - self._embedding_engine = embedding_engine.execution_options( - schema_translate_map={ - **(embedding_engine.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), - REGISTRY_SCHEMA_KEY: resolve_registry_schema(embedding_engine), - } - ) + # 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_schema(embedding_engine), + } + ) self._embedding_sessionmaker = sessionmaker(self._embedding_engine) self._read_only = not initialize self._registry_available = inspect(self._embedding_engine).has_table( diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index c6b57b0..f169e6e 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -30,16 +30,12 @@ from omop_emb.config import ( MODEL_REGISTRY_SCHEMA, + REGISTRY_SCHEMA_KEY, IndexType, MetricType, ) from omop_emb.backends.index_config import IndexConfig -# Schema name for the model registry table. Dialects with schema support -# store the registry table in a dedicated schema to allow schema-independent -# access to the registry table from any schema in the same database. -REGISTRY_SCHEMA_KEY = "registry" - class ModelRegistryBase(DeclarativeBase): """Dedicated declarative base for local model registry metadata.""" From 7b0216e4b219d15e3abc71ed801684fe8b3a16e1 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Wed, 23 Sep 2026 04:47:54 +0000 Subject: [PATCH 13/16] Disambiguate physical schema from schema tag --- src/omop_emb/backends/pgvector/pg_index_manager.py | 8 ++++---- src/omop_emb/backends/pgvector/pg_sql.py | 8 ++++---- src/omop_emb/config.py | 3 ++- 3 files changed, 10 insertions(+), 9 deletions(-) diff --git a/src/omop_emb/backends/pgvector/pg_index_manager.py b/src/omop_emb/backends/pgvector/pg_index_manager.py index 4ac8f0a..4f73455 100644 --- a/src/omop_emb/backends/pgvector/pg_index_manager.py +++ b/src/omop_emb/backends/pgvector/pg_index_manager.py @@ -17,7 +17,7 @@ import logging from typing import Generic, TypeVar -from oa_configurator import qualified, schema_of +from oa_configurator import qualified, physical_schema_of from sqlalchemy import Engine, inspect, text from omop_emb.config import IndexType, MetricType, VectorColumnType @@ -80,7 +80,7 @@ def has_index(self, metric_type: MetricType) -> bool: existing = { idx["name"] for idx in inspect(conn).get_indexes( - self._tablename, schema=schema_of(conn) + self._tablename, schema=physical_schema_of(conn) ) } return self._index_name(metric_type) in existing @@ -106,7 +106,7 @@ def drop_index(self, metric_type: MetricType) -> None: existed = self.has_index(metric_type) with self._engine.begin() as conn: conn.execute( - text(f"DROP INDEX IF EXISTS {qualified(conn, name, physical_schema=schema_of(conn))}") + text(f"DROP INDEX IF EXISTS {qualified(conn, name, physical_schema=physical_schema_of(conn))}") ) if existed: logger.info(f"Dropped pgvector index '{name}'.") @@ -196,7 +196,7 @@ 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=schema_of(self._engine)) + 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 {table_ref} " diff --git a/src/omop_emb/backends/pgvector/pg_sql.py b/src/omop_emb/backends/pgvector/pg_sql.py index fb2f23a..5131b83 100644 --- a/src/omop_emb/backends/pgvector/pg_sql.py +++ b/src/omop_emb/backends/pgvector/pg_sql.py @@ -22,7 +22,7 @@ ensure_schema, guard_schema_provenance_for, qualified, - schema_of, + 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 @@ -41,7 +41,7 @@ def table_exists(engine: Engine, table_name: str) -> bool: """Return ``True`` if *table_name* exists in the engine's configured schema.""" - return sa_inspect(engine).has_table(table_name, schema=schema_of(engine)) + return sa_inspect(engine).has_table(table_name, schema=physical_schema_of(engine)) def create_pg_embedding_table( engine: Engine, @@ -75,7 +75,7 @@ def create_pg_embedding_table( table_cls = pg_embedding_table_descriptor(model_record) with engine.begin() as connection: # Ensure that the schema exists before creating the table - physical_schema = schema_of(connection, schema_tag=Role.PRIMARY) + 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] @@ -96,7 +96,7 @@ 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 {qualified(conn, tablename, physical_schema=schema_of(conn))}") + text(f"DROP TABLE IF EXISTS {qualified(conn, tablename, physical_schema=physical_schema_of(conn))}") ) logger.info(f"Dropped embedding table '{tablename}'.") diff --git a/src/omop_emb/config.py b/src/omop_emb/config.py index 3e2a6d4..170435b 100644 --- a/src/omop_emb/config.py +++ b/src/omop_emb/config.py @@ -17,6 +17,7 @@ ResolvedVectorStore, VectorStoreConfig, register_reserved_schema, + register_reserved_schema_tag, ) # Guaranteed to be imported and registered if there is a config @@ -24,7 +25,7 @@ REGISTRY_SCHEMA_KEY: str = "registry" register_reserved_schema(MODEL_REGISTRY_SCHEMA, owner="omop_emb") -register_reserved_schema(REGISTRY_SCHEMA_KEY, owner="omop_emb") +register_reserved_schema_tag(REGISTRY_SCHEMA_KEY, owner="omop_emb") class OmopEmbConfig(PackageConfigBase): From 87d60109522f72afa1865809f70966c6b7965f39 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Wed, 23 Sep 2026 05:30:14 +0000 Subject: [PATCH 14/16] Correct naming and wiring --- src/omop_emb/backends/base_backend.py | 4 ++-- src/omop_emb/backends/read_only.py | 16 ++++++++-------- src/omop_emb/model_registry/__init__.py | 4 ++-- .../model_registry/model_registry_manager.py | 6 +++--- .../model_registry/model_registry_orm.py | 10 +++++----- tests/test_read_only_population.py | 6 +++--- 6 files changed, 23 insertions(+), 23 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index c0bbd66..70b740e 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -31,7 +31,7 @@ EmbeddingModelRecord, REGISTRY_SCHEMA_KEY, RegistryManager, - resolve_registry_schema, + resolve_registry_physical_schema, ) from omop_emb.utils.embedding_utils import ( EmbeddingConceptFilter, @@ -1055,7 +1055,7 @@ def resolve_backend( # 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_schema(database.connection.dialect_name) + registry_schema = resolve_registry_physical_schema(database.connection.dialect_name) registry_schema_translate_map = {REGISTRY_SCHEMA_KEY: registry_schema} if resolved_backend == BackendType.SQLITEVEC: diff --git a/src/omop_emb/backends/read_only.py b/src/omop_emb/backends/read_only.py index 9c5b373..dd431d4 100644 --- a/src/omop_emb/backends/read_only.py +++ b/src/omop_emb/backends/read_only.py @@ -49,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: @@ -98,12 +98,12 @@ 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 = schema_if_supported(self.schema, self._engine) + 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( @@ -138,7 +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) ) @@ -147,7 +147,7 @@ def drop_index_sql(self, model_name: str) -> tuple[str, ...]: """Return reviewed index-removal statements without executing them.""" return tuple( - f"DROP INDEX IF EXISTS {qualified(self._engine, name, physical_schema=self.schema)};" + f"DROP INDEX IF EXISTS {qualified(self._engine, name, physical_schema=self.physical_schema)};" for name in self.physical_indexes(model_name) ) @@ -187,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/model_registry/__init__.py b/src/omop_emb/model_registry/__init__.py index 231b0d4..34b37f3 100644 --- a/src/omop_emb/model_registry/__init__.py +++ b/src/omop_emb/model_registry/__init__.py @@ -4,7 +4,7 @@ REGISTRY_SCHEMA_KEY, ModelRegistry, ensure_registry_table, - resolve_registry_schema, + resolve_registry_physical_schema, ) __all__ = [ @@ -13,5 +13,5 @@ "ModelRegistry", "ensure_registry_table", "REGISTRY_SCHEMA_KEY", - "resolve_registry_schema", + "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 8de1c67..eddf51f 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -18,7 +18,7 @@ REGISTRY_SCHEMA_KEY, ModelRegistry, ensure_registry_table, - resolve_registry_schema, + resolve_registry_physical_schema, ) from omop_emb.model_registry.model_registry_types import EmbeddingModelRecord from omop_emb.utils.errors import ModelRegistrationConflictError @@ -60,13 +60,13 @@ def __init__( self._embedding_engine = embedding_engine.execution_options( schema_translate_map={ **existing_schema_translate_map, - REGISTRY_SCHEMA_KEY: resolve_registry_schema(embedding_engine), + 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(self._embedding_engine).has_table( - ModelRegistry.__tablename__, schema=resolve_registry_schema(self._embedding_engine) + ModelRegistry.__tablename__, schema=resolve_registry_physical_schema(self._embedding_engine) ) if initialize: ensure_registry_table(self._embedding_engine, resolved=resolved) diff --git a/src/omop_emb/model_registry/model_registry_orm.py b/src/omop_emb/model_registry/model_registry_orm.py index f169e6e..8ae9a43 100644 --- a/src/omop_emb/model_registry/model_registry_orm.py +++ b/src/omop_emb/model_registry/model_registry_orm.py @@ -189,7 +189,7 @@ def _validate_and_sync_index_config( return index_config.to_dict() -def resolve_registry_schema(bindable) -> str | None: +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 @@ -217,11 +217,11 @@ def ensure_registry_table(engine: Engine, *, resolved: ResolvedDatabase | None = connection = connection.execution_options( schema_translate_map={ **(connection.get_execution_options().get(SCHEMA_TRANSLATE_MAP_KEY) or {}), - REGISTRY_SCHEMA_KEY: resolve_registry_schema(connection), + REGISTRY_SCHEMA_KEY: resolve_registry_physical_schema(connection), } ) ensure_schema(connection, MODEL_REGISTRY_SCHEMA) - registry_schema = resolve_registry_schema(connection) + registry_schema = resolve_registry_physical_schema(connection) guard = ( guard_schema_provenance( connection, @@ -252,7 +252,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: can safely run it for both existing and newly-created registries. """ columns = inspect(engine).get_columns( - ModelRegistry.__tablename__, schema=resolve_registry_schema(engine) + ModelRegistry.__tablename__, schema=resolve_registry_physical_schema(engine) ) provider_column = next( (column for column in columns if column["name"] == "provider_type"), @@ -263,7 +263,7 @@ def _migrate_legacy_provider_type_column(engine: Engine) -> None: legacy_length = getattr(provider_column["type"], "length", None) with engine.begin() as connection: - registry_schema = resolve_registry_schema(connection) + 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 " diff --git a/tests/test_read_only_population.py b/tests/test_read_only_population.py index de40615..f066daf 100644 --- a/tests/test_read_only_population.py +++ b/tests/test_read_only_population.py @@ -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 @@ -50,7 +50,7 @@ def test_explicit_registry_initialization_is_visible_to_read_only_store() -> Non 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), From 41422236cd3a2eca576b7b4a763b2cad615d9af1 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Thu, 24 Sep 2026 00:16:13 +0000 Subject: [PATCH 15/16] Docstring rectifications --- src/omop_emb/backends/base_backend.py | 4 +++- src/omop_emb/backends/index_config.py | 4 +--- src/omop_emb/backends/pgvector/pg_backend.py | 6 +++++- src/omop_emb/interface.py | 3 ++- .../model_registry/model_registry_manager.py | 9 ++++++--- src/omop_emb/storage/embedding_bundle.py | 9 +++++---- src/omop_emb/utils/errors.py | 16 +++++++++++++--- 7 files changed, 35 insertions(+), 16 deletions(-) diff --git a/src/omop_emb/backends/base_backend.py b/src/omop_emb/backends/base_backend.py index 70b740e..f355577 100644 --- a/src/omop_emb/backends/base_backend.py +++ b/src/omop_emb/backends/base_backend.py @@ -294,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: 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 b3114f6..08acc99 100644 --- a/src/omop_emb/backends/pgvector/pg_backend.py +++ b/src/omop_emb/backends/pgvector/pg_backend.py @@ -173,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/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/model_registry_manager.py b/src/omop_emb/model_registry/model_registry_manager.py index eddf51f..9919e99 100644 --- a/src/omop_emb/model_registry/model_registry_manager.py +++ b/src/omop_emb/model_registry/model_registry_manager.py @@ -39,9 +39,10 @@ class RegistryManager: Notes ----- - The registry table is created in a schema named ``registry`` for dialects - supporting schema registration. Allows schema-independent access to the registry table - from any schema in the same database. + 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__( @@ -174,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/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/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): From d0fd1c7481803f55824be68dc7bbb08b54496b73 Mon Sep 17 00:00:00 2001 From: Nico Loesch Date: Thu, 24 Sep 2026 02:48:30 +0000 Subject: [PATCH 16/16] Updated Docs --- docs/usage/backend-selection.md | 2 +- docs/usage/installation.md | 2 +- docs/usage/interface-guide.md | 11 +++++++++-- 3 files changed, 11 insertions(+), 4 deletions(-) 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 11b794e..6f6771d 100644 --- a/docs/usage/interface-guide.md +++ b/docs/usage/interface-guide.md @@ -191,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"), @@ -272,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", ) @@ -284,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-...", )