fix: migrate mental_models.embedding dimension alongside memory_units (#526)
ensure_embedding_dimension() now also checks and migrates mental_models.embedding, fixing silent failures when changing embedding model dimensions. Extracted shared per-table logic into _migrate_table_embedding_dimension() to avoid duplication. Adds test coverage for the mental_models dimension migration path. Fixes #523
This commit is contained in:
parent
fa3501d448
commit
7accac94b2
2 changed files with 203 additions and 119 deletions
|
|
@ -25,7 +25,7 @@ from pathlib import Path
|
|||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from alembic.script.revision import ResolutionError
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy import Connection, create_engine, text
|
||||
|
||||
from .utils import mask_network_location
|
||||
|
||||
|
|
@ -471,6 +471,125 @@ def check_migration_status(
|
|||
return None, None
|
||||
|
||||
|
||||
def _migrate_table_embedding_dimension(
|
||||
conn: Connection,
|
||||
schema_name: str,
|
||||
table_name: str,
|
||||
required_dimension: int,
|
||||
vector_ext: str,
|
||||
) -> None:
|
||||
"""
|
||||
Migrate the embedding column of a single table to the required dimension.
|
||||
|
||||
- If dimensions match: no action needed
|
||||
- If dimensions differ and table is empty: ALTER COLUMN to new dimension
|
||||
- If dimensions differ and table has data: raise error with migration guidance
|
||||
"""
|
||||
current_dim = conn.execute(
|
||||
text("""
|
||||
SELECT atttypmod
|
||||
FROM pg_attribute a
|
||||
JOIN pg_class c ON a.attrelid = c.oid
|
||||
JOIN pg_namespace n ON c.relnamespace = n.oid
|
||||
WHERE n.nspname = :schema
|
||||
AND c.relname = :table
|
||||
AND a.attname = 'embedding'
|
||||
"""),
|
||||
{"schema": schema_name, "table": table_name},
|
||||
).scalar()
|
||||
|
||||
if current_dim is None:
|
||||
logger.debug(f"No embedding column found on {table_name}, skipping")
|
||||
return
|
||||
|
||||
if current_dim == required_dimension:
|
||||
logger.debug(f"Embedding dimension OK for {table_name}: {current_dim}")
|
||||
return
|
||||
|
||||
logger.info(
|
||||
f"Embedding dimension mismatch on {table_name}: database has {current_dim}, model requires {required_dimension}"
|
||||
)
|
||||
|
||||
row_count = conn.execute(
|
||||
text(f"SELECT COUNT(*) FROM {schema_name}.{table_name} WHERE embedding IS NOT NULL")
|
||||
).scalar()
|
||||
|
||||
if row_count > 0:
|
||||
raise RuntimeError(
|
||||
f"Cannot change embedding dimension from {current_dim} to {required_dimension}: "
|
||||
f"{table_name} table contains {row_count} rows with embeddings. "
|
||||
f"To change dimensions, you must either:\n"
|
||||
f" 1. Re-embed all data: DELETE FROM {schema_name}.{table_name}; then restart\n"
|
||||
f" 2. Use a model with {current_dim}-dimensional embeddings"
|
||||
)
|
||||
|
||||
logger.info(f"Altering {table_name}.embedding column dimension from {current_dim} to {required_dimension}")
|
||||
|
||||
# Drop existing vector index (works for both HNSW and vchordrq)
|
||||
conn.execute(
|
||||
text(f"""
|
||||
DO $$
|
||||
DECLARE idx_name TEXT;
|
||||
BEGIN
|
||||
FOR idx_name IN
|
||||
SELECT indexname FROM pg_indexes
|
||||
WHERE schemaname = '{schema_name}'
|
||||
AND tablename = '{table_name}'
|
||||
AND (indexdef LIKE '%hnsw%' OR indexdef LIKE '%vchordrq%' OR indexdef LIKE '%diskann%')
|
||||
AND indexdef LIKE '%embedding%'
|
||||
LOOP
|
||||
EXECUTE 'DROP INDEX IF EXISTS {schema_name}.' || idx_name;
|
||||
END LOOP;
|
||||
END $$;
|
||||
""")
|
||||
)
|
||||
|
||||
conn.execute(
|
||||
text(f"ALTER TABLE {schema_name}.{table_name} ALTER COLUMN embedding TYPE vector({required_dimension})")
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# Recreate index with appropriate type based on detected extension
|
||||
if vector_ext == "pgvectorscale":
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_{table_name}_embedding_diskann
|
||||
ON {schema_name}.{table_name}
|
||||
USING diskann (embedding vector_cosine_ops)
|
||||
WITH (num_neighbors = 50)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created DiskANN index on {table_name} for {required_dimension}-dimensional embeddings")
|
||||
elif vector_ext == "vchord":
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_{table_name}_embedding_vchordrq
|
||||
ON {schema_name}.{table_name}
|
||||
USING vchordrq (embedding vector_l2_ops)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created vchordrq index on {table_name} for {required_dimension}-dimensional embeddings")
|
||||
else: # pgvector
|
||||
if required_dimension > 2000:
|
||||
raise RuntimeError(
|
||||
f"Embedding dimension {required_dimension} exceeds pgvector HNSW index limit of 2000. "
|
||||
f"Use an embedding model with <= 2000 dimensions, or switch to a vector extension "
|
||||
f"that supports higher dimensions (e.g., pgvectorscale/DiskANN)."
|
||||
)
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_{table_name}_embedding_hnsw
|
||||
ON {schema_name}.{table_name}
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
WITH (m = 16, ef_construction = 64)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created HNSW index on {table_name} for {required_dimension}-dimensional embeddings")
|
||||
conn.commit()
|
||||
|
||||
logger.info(f"Successfully changed {table_name}.embedding dimension to {required_dimension}")
|
||||
|
||||
|
||||
def ensure_embedding_dimension(
|
||||
database_url: str,
|
||||
required_dimension: int,
|
||||
|
|
@ -478,10 +597,9 @@ def ensure_embedding_dimension(
|
|||
vector_extension: str = "pgvector",
|
||||
) -> None:
|
||||
"""
|
||||
Ensure the embedding column dimension matches the model's dimension.
|
||||
Ensure the embedding column dimension matches the model's dimension for all tables.
|
||||
|
||||
This function checks the current vector column dimension in the database
|
||||
and adjusts it if necessary:
|
||||
Checks and adjusts memory_units.embedding and mental_models.embedding:
|
||||
- If dimensions match: no action needed
|
||||
- If dimensions differ and table is empty: ALTER COLUMN to new dimension
|
||||
- If dimensions differ and table has data: raise error with migration guidance
|
||||
|
|
@ -499,7 +617,7 @@ def ensure_embedding_dimension(
|
|||
|
||||
engine = create_engine(database_url)
|
||||
with engine.connect() as conn:
|
||||
# Check if memory_units table exists
|
||||
# Check if memory_units table exists (proxy for schema being initialized)
|
||||
table_exists = conn.execute(
|
||||
text("""
|
||||
SELECT EXISTS (
|
||||
|
|
@ -518,117 +636,8 @@ def ensure_embedding_dimension(
|
|||
vector_ext = _detect_vector_extension(conn, vector_extension)
|
||||
logger.info(f"Using vector extension: {vector_ext}")
|
||||
|
||||
# Get current column dimension from pg_attribute
|
||||
# pgvector stores dimension in atttypmod
|
||||
current_dim = conn.execute(
|
||||
text("""
|
||||
SELECT atttypmod
|
||||
FROM pg_attribute a
|
||||
JOIN pg_class c ON a.attrelid = c.oid
|
||||
JOIN pg_namespace n ON c.relnamespace = n.oid
|
||||
WHERE n.nspname = :schema
|
||||
AND c.relname = 'memory_units'
|
||||
AND a.attname = 'embedding'
|
||||
"""),
|
||||
{"schema": schema_name},
|
||||
).scalar()
|
||||
|
||||
if current_dim is None:
|
||||
logger.warning("Could not determine current embedding dimension, skipping check")
|
||||
return
|
||||
|
||||
# pgvector stores dimension directly in atttypmod (no offset like other types)
|
||||
current_dimension = current_dim
|
||||
|
||||
if current_dimension == required_dimension:
|
||||
logger.debug(f"Embedding dimension OK: {current_dimension}")
|
||||
return
|
||||
|
||||
logger.info(
|
||||
f"Embedding dimension mismatch: database has {current_dimension}, model requires {required_dimension}"
|
||||
)
|
||||
|
||||
# Check if table has data
|
||||
row_count = conn.execute(
|
||||
text(f"SELECT COUNT(*) FROM {schema_name}.memory_units WHERE embedding IS NOT NULL")
|
||||
).scalar()
|
||||
|
||||
if row_count > 0:
|
||||
raise RuntimeError(
|
||||
f"Cannot change embedding dimension from {current_dimension} to {required_dimension}: "
|
||||
f"memory_units table contains {row_count} rows with embeddings. "
|
||||
f"To change dimensions, you must either:\n"
|
||||
f" 1. Re-embed all data: DELETE FROM {schema_name}.memory_units; then restart\n"
|
||||
f" 2. Use a model with {current_dimension}-dimensional embeddings"
|
||||
)
|
||||
|
||||
# Table is empty, safe to alter column
|
||||
logger.info(f"Altering embedding column dimension from {current_dimension} to {required_dimension}")
|
||||
|
||||
# Drop existing vector index (works for both HNSW and vchordrq)
|
||||
conn.execute(
|
||||
text(f"""
|
||||
DO $$
|
||||
DECLARE idx_name TEXT;
|
||||
BEGIN
|
||||
FOR idx_name IN
|
||||
SELECT indexname FROM pg_indexes
|
||||
WHERE schemaname = '{schema_name}'
|
||||
AND tablename = 'memory_units'
|
||||
AND (indexdef LIKE '%hnsw%' OR indexdef LIKE '%vchordrq%')
|
||||
AND indexdef LIKE '%embedding%'
|
||||
LOOP
|
||||
EXECUTE 'DROP INDEX IF EXISTS {schema_name}.' || idx_name;
|
||||
END LOOP;
|
||||
END $$;
|
||||
""")
|
||||
)
|
||||
|
||||
# Alter the column type
|
||||
conn.execute(
|
||||
text(f"ALTER TABLE {schema_name}.memory_units ALTER COLUMN embedding TYPE vector({required_dimension})")
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# Recreate index with appropriate type based on detected extension
|
||||
if vector_ext == "pgvectorscale":
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_embedding_diskann
|
||||
ON {schema_name}.memory_units
|
||||
USING diskann (embedding vector_cosine_ops)
|
||||
WITH (num_neighbors = 50)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created DiskANN index for {required_dimension}-dimensional embeddings")
|
||||
elif vector_ext == "vchord":
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_embedding_vchordrq
|
||||
ON {schema_name}.memory_units
|
||||
USING vchordrq (embedding vector_l2_ops)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created vchordrq index for {required_dimension}-dimensional embeddings")
|
||||
else: # pgvector
|
||||
if required_dimension > 2000:
|
||||
raise RuntimeError(
|
||||
f"Embedding dimension {required_dimension} exceeds pgvector HNSW index limit of 2000. "
|
||||
f"Use an embedding model with <= 2000 dimensions, or switch to a vector extension "
|
||||
f"that supports higher dimensions (e.g., pgvectorscale/DiskANN)."
|
||||
)
|
||||
conn.execute(
|
||||
text(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_embedding_hnsw
|
||||
ON {schema_name}.memory_units
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
WITH (m = 16, ef_construction = 64)
|
||||
""")
|
||||
)
|
||||
logger.info(f"Created HNSW index for {required_dimension}-dimensional embeddings")
|
||||
conn.commit()
|
||||
|
||||
logger.info(f"Successfully changed embedding dimension to {required_dimension}")
|
||||
_migrate_table_embedding_dimension(conn, schema_name, "memory_units", required_dimension, vector_ext)
|
||||
_migrate_table_embedding_dimension(conn, schema_name, "mental_models", required_dimension, vector_ext)
|
||||
|
||||
|
||||
def ensure_vector_extension(
|
||||
|
|
|
|||
|
|
@ -75,7 +75,7 @@ def drop_schema(db_url: str, schema_name: str):
|
|||
conn.commit()
|
||||
|
||||
|
||||
def get_column_dimension(db_url: str, schema: str = "public") -> int | None:
|
||||
def get_column_dimension(db_url: str, schema: str = "public", table: str = "memory_units") -> int | None:
|
||||
"""Get the current embedding column dimension from the database."""
|
||||
engine = create_engine(db_url)
|
||||
with engine.connect() as conn:
|
||||
|
|
@ -86,10 +86,10 @@ def get_column_dimension(db_url: str, schema: str = "public") -> int | None:
|
|||
JOIN pg_class c ON a.attrelid = c.oid
|
||||
JOIN pg_namespace n ON c.relnamespace = n.oid
|
||||
WHERE n.nspname = :schema
|
||||
AND c.relname = 'memory_units'
|
||||
AND c.relname = :table
|
||||
AND a.attname = 'embedding'
|
||||
"""),
|
||||
{"schema": schema},
|
||||
{"schema": schema, "table": table},
|
||||
).scalar()
|
||||
return result
|
||||
|
||||
|
|
@ -125,6 +125,38 @@ def clear_embeddings(db_url: str, schema: str):
|
|||
conn.commit()
|
||||
|
||||
|
||||
def insert_test_mental_model_embedding(db_url: str, schema: str, dimension: int):
|
||||
"""Insert a test mental model row with a dummy embedding."""
|
||||
engine = create_engine(db_url)
|
||||
embedding = [0.1] * dimension
|
||||
embedding_str = "[" + ",".join(str(x) for x in embedding) + "]"
|
||||
|
||||
with engine.connect() as conn:
|
||||
# Ensure test bank exists
|
||||
conn.execute(
|
||||
text(f"""
|
||||
INSERT INTO {schema}.banks (bank_id, name)
|
||||
VALUES ('test-bank-mm', 'Test Bank')
|
||||
ON CONFLICT (bank_id) DO NOTHING
|
||||
""")
|
||||
)
|
||||
conn.execute(
|
||||
text(f"""
|
||||
INSERT INTO {schema}.mental_models (bank_id, name, source_query, content, embedding)
|
||||
VALUES ('test-bank-mm', 'test model', 'test query', 'test content', '{embedding_str}'::vector)
|
||||
""")
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
|
||||
def clear_mental_model_embeddings(db_url: str, schema: str):
|
||||
"""Clear all rows from mental_models."""
|
||||
engine = create_engine(db_url)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(text(f"DELETE FROM {schema}.mental_models"))
|
||||
conn.commit()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Embedding Dimension Tests (Local Embeddings)
|
||||
# =============================================================================
|
||||
|
|
@ -199,6 +231,49 @@ class TestEmbeddingDimension:
|
|||
# Cleanup
|
||||
clear_embeddings(db_url, schema)
|
||||
|
||||
def test_mental_models_dimension_matches_no_change(self, dimension_test_schema):
|
||||
"""When mental_models dimension matches, no changes should be made."""
|
||||
db_url, schema = dimension_test_schema
|
||||
|
||||
initial_dim = get_column_dimension(db_url, schema, table="mental_models")
|
||||
assert initial_dim == 384, f"Expected 384, got {initial_dim}"
|
||||
|
||||
ensure_embedding_dimension(db_url, 384, schema=schema)
|
||||
|
||||
assert get_column_dimension(db_url, schema, table="mental_models") == 384
|
||||
|
||||
def test_mental_models_dimension_change_empty_table(self, dimension_test_schema):
|
||||
"""When mental_models is empty, dimension can be changed."""
|
||||
db_url, schema = dimension_test_schema
|
||||
|
||||
clear_mental_model_embeddings(db_url, schema)
|
||||
|
||||
ensure_embedding_dimension(db_url, 768, schema=schema)
|
||||
|
||||
assert get_column_dimension(db_url, schema, table="mental_models") == 768
|
||||
|
||||
# Change back for other tests
|
||||
ensure_embedding_dimension(db_url, 384, schema=schema)
|
||||
assert get_column_dimension(db_url, schema, table="mental_models") == 384
|
||||
|
||||
def test_mental_models_dimension_change_blocked_with_data(self, dimension_test_schema):
|
||||
"""When mental_models has data, dimension change should be blocked."""
|
||||
db_url, schema = dimension_test_schema
|
||||
|
||||
clear_mental_model_embeddings(db_url, schema)
|
||||
insert_test_mental_model_embedding(db_url, schema, 384)
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
ensure_embedding_dimension(db_url, 768, schema=schema)
|
||||
|
||||
assert "Cannot change embedding dimension" in str(exc_info.value)
|
||||
assert "mental_models" in str(exc_info.value)
|
||||
|
||||
assert get_column_dimension(db_url, schema, table="mental_models") == 384
|
||||
|
||||
# Cleanup
|
||||
clear_mental_model_embeddings(db_url, schema)
|
||||
|
||||
def test_local_embeddings_dimension_detection(self, embeddings):
|
||||
"""Test that LocalSTEmbeddings correctly detects dimension."""
|
||||
# Initialize embeddings if not already done
|
||||
|
|
|
|||
Loading…
Reference in a new issue