diff --git a/src/reconciler/sync_vectors.py b/src/reconciler/sync_vectors.py index 5edd6612..032069b9 100644 --- a/src/reconciler/sync_vectors.py +++ b/src/reconciler/sync_vectors.py @@ -26,7 +26,7 @@ logger = logging.getLogger(__name__) # Constants RECONCILIATION_BATCH_SIZE = 50 RECONCILIATION_TIME_BUDGET_SECONDS = 240 # Leave headroom for other maintenance work -MAX_SYNC_ATTEMPTS = 5 # After this many failures, mark as permanently_failed +MAX_SYNC_ATTEMPTS = 5 # After this many failures, mark as failed @dataclass @@ -89,11 +89,10 @@ async def _get_message_embeddings_needing_sync( batch_size: int = RECONCILIATION_BATCH_SIZE, ) -> list[models.MessageEmbedding]: """ - Get message embeddings that need to be synced to the vector store. + Get pending message embeddings that need to be synced to the vector store. - Selects models.MessageEmbedding records where sync_state is "pending", - regardless of whether an embedding vector exists in the database. - Records missing embeddings will be re-embedded during reconciliation. + Returns only pending embeddings (with full data including embedding vectors). + The batch_size limits the number of embeddings returned. Uses FOR UPDATE SKIP LOCKED to prevent concurrent processing and orders by last_sync_at (nulls first) to prioritize never-synced records. @@ -102,9 +101,7 @@ async def _get_message_embeddings_needing_sync( """ stmt = ( select(models.MessageEmbedding) - .where( - models.MessageEmbedding.sync_state == "pending" # Only pending items - ) + .where(models.MessageEmbedding.sync_state == "pending") .order_by(models.MessageEmbedding.last_sync_at.asc().nullsfirst()) .limit(batch_size) .with_for_update(skip_locked=True) @@ -209,11 +206,9 @@ async def _sync_documents( .where(models.Document.id == doc_id) .values(embedding=embedding) ) - except Exception as e: - logger.warning( - "Failed to re-embed %s documents for reconciliation: %s", - len(missing_docs), - e, + except Exception: + logger.exception( + "Failed to re-embed %s documents for reconciliation", len(missing_docs) ) missing_after_embed: list[models.Document] = [] @@ -286,9 +281,10 @@ async def _sync_documents( ) synced_count += len(docs_with_vectors) - except Exception as e: - logger.warning( - f"Failed to sync documents to external vector store {namespace}: {e}" + except Exception: + logger.exception( + "Failed to sync documents to external vector store %s", + namespace, ) # Increment attempts and mark as failed if we've hit max attempts await _bump_document_sync_attempts(db, docs_with_vectors) @@ -303,7 +299,12 @@ async def _sync_message_embeddings( external_vector_store: VectorStore, ) -> tuple[int, int]: """ - Sync a batch of message embeddings to the external vector store. + Sync a batch of pending message embeddings to the external vector store. + + Args: + db: Database session + embeddings: List of pending MessageEmbedding records to sync + external_vector_store: External vector store to sync to Returns (synced_count, failed_count). """ @@ -319,7 +320,7 @@ async def _sync_message_embeddings( settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED ) - # Re-embed missing payloads so reconciliation can heal non-pgvector writes + # Re-embed embeddings that are missing their vector payload missing_embs = [emb for emb in embeddings if emb.embedding is None] reembedded_by_id: dict[int, list[float]] = {} @@ -336,13 +337,13 @@ async def _sync_message_embeddings( .where(models.MessageEmbedding.id == emb.id) .values(embedding=new_embedding) ) - except Exception as e: - logger.warning( - "Failed to re-embed %s message embeddings for reconciliation: %s", + except Exception: + logger.exception( + "Failed to re-embed %s message embeddings for reconciliation", len(missing_embs), - e, ) + # Track embeddings that still don't have an embedding after re-embed attempt missing_after_embed: list[models.MessageEmbedding] = [] for emb in embeddings: if emb.embedding is None and emb.id not in reembedded_by_id: @@ -351,7 +352,31 @@ async def _sync_message_embeddings( await _bump_message_embedding_sync_attempts(db, missing_after_embed) failed_count += len(missing_after_embed) - # Group by namespace (workspace) + # Compute chunk position for each embedding within its parent message. + # Messages can be split into multiple embedding chunks; we need to track + # which chunk position (0, 1, 2, ...) each MessageEmbedding represents. + # Fetch sibling embedding IDs (lightweight query) to compute correct positions. + message_ids = list({emb.message_id for emb in embeddings}) + sibling_stmt = ( + select(models.MessageEmbedding.id, models.MessageEmbedding.message_id) + .where(models.MessageEmbedding.message_id.in_(message_ids)) + .order_by(models.MessageEmbedding.message_id, models.MessageEmbedding.id) + ) + sibling_result = await db.execute(sibling_stmt) + sibling_rows = sibling_result.all() + + # Build position mapping from sibling IDs + embeddings_by_message_id: dict[str, list[int]] = {} + for emb_id, msg_id in sibling_rows: + embeddings_by_message_id.setdefault(msg_id, []).append(emb_id) + + chunk_position_by_emb_id: dict[int, int] = {} + for emb_ids in embeddings_by_message_id.values(): + # IDs are already sorted by the query + for position, emb_id in enumerate(emb_ids): + chunk_position_by_emb_id[emb_id] = position + + # Group embeddings by namespace (workspace) by_namespace: dict[str, list[models.MessageEmbedding]] = {} for emb in embeddings: namespace = external_vector_store.get_vector_namespace( @@ -359,20 +384,6 @@ async def _sync_message_embeddings( ) by_namespace.setdefault(namespace, []).append(emb) - # Compute chunk position for each embedding within its parent message. - # Messages can be split into multiple embedding chunks; we need to track - # which chunk position (0, 1, 2, ...) each MessageEmbedding represents. - embeddings_by_message_id: dict[str, list[models.MessageEmbedding]] = {} - for emb in embeddings: - embeddings_by_message_id.setdefault(emb.message_id, []).append(emb) - - # Sort each message's embeddings by id and build position mapping - chunk_position_by_emb_id: dict[int, int] = {} - for msg_embeddings in embeddings_by_message_id.values(): - msg_embeddings.sort(key=lambda e: e.id) - for position, msg_emb in enumerate(msg_embeddings): - chunk_position_by_emb_id[msg_emb.id] = position - # Sync each namespace batch for namespace, embs in by_namespace.items(): embs_with_vectors: list[models.MessageEmbedding] = [] @@ -424,9 +435,10 @@ async def _sync_message_embeddings( ) synced_count += len(embs_with_vectors) - except Exception as e: - logger.warning( - f"Failed to sync message embeddings to external vector store {namespace}: {e}" + except Exception: + logger.exception( + "Failed to sync message embeddings to external vector store %s", + namespace, ) # Increment attempts and mark as failed if we've hit max attempts await _bump_message_embedding_sync_attempts(db, embs_with_vectors) @@ -519,11 +531,10 @@ async def run_vector_reconciliation_cycle() -> ReconciliationMetrics: synced, failed = await _sync_message_embeddings( db, embs, external_vector_store ) - except Exception as e: - logger.warning( - "Message embedding reconciliation failed for %s embeddings: %s", + except Exception: + logger.exception( + "Message embedding reconciliation failed for %s embeddings", len(embs), - e, ) await _bump_message_embedding_sync_attempts(db, embs) synced = 0 diff --git a/src/vector_store/lancedb.py b/src/vector_store/lancedb.py index 333b9b6b..3c35e2c8 100644 --- a/src/vector_store/lancedb.py +++ b/src/vector_store/lancedb.py @@ -6,6 +6,7 @@ This module provides a LanceDB-based implementation of the VectorStore interface import asyncio import logging +import re from collections.abc import Sequence from typing import Any, cast @@ -20,6 +21,9 @@ from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult logger = logging.getLogger(__name__) +# Pattern for valid SQL identifiers (alphanumeric + underscore, not starting with digit) +_VALID_IDENTIFIER_PATTERN = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$") + # Schema for LanceDB tables # id: string, vector: fixed_size_list of float32 (1536 dimensions for OpenAI embeddings) # Additional metadata columns are added dynamically @@ -84,6 +88,8 @@ class LanceDBVectorStore(VectorStore): return await db.open_table(namespace) # Create empty table with base schema + # Handle race condition: another worker may have created the table + # between our check and create_table call fields: list[pa.Field] = [ pa.field("id", pa.string()), pa.field( @@ -92,8 +98,12 @@ class LanceDBVectorStore(VectorStore): ] fields.extend(self._metadata_fields_for_namespace(namespace)) schema = pa.schema(fields) - table = await db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType] - return table + try: + table = await db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType] + return table + except Exception: + # Table may have been created by another worker, try to open it + return await db.open_table(namespace) def _metadata_fields_for_namespace(self, namespace: str) -> list[pa.Field]: """ @@ -263,12 +273,19 @@ class LanceDBVectorStore(VectorStore): Returns: SQL WHERE clause string or None if no filters + + Raises: + ValueError: If a filter key is not a valid SQL identifier """ if not filters: return None conditions: list[str] = [] for key, value in filters.items(): + # Validate key is a safe SQL identifier to prevent injection + if not _VALID_IDENTIFIER_PATTERN.match(key): + raise ValueError(f"Invalid filter key: {key!r}") + # Check if value is a dict with "in" operator if isinstance(value, dict) and "in" in value: # IN clause for list membership diff --git a/tests/conftest.py b/tests/conftest.py index 71f7fffd..3dc4d7db 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -482,7 +482,7 @@ def mock_vector_store(): patch("src.crud.session.get_external_vector_store", return_value=mock_vs), patch("src.crud.message.get_external_vector_store", return_value=mock_vs), patch( - "src.deriver.vector_reconciliation.get_external_vector_store", + "src.reconciler.sync_vectors.get_external_vector_store", return_value=mock_vs, ), patch("src.utils.search.get_external_vector_store", return_value=mock_vs), diff --git a/tests/deriver/test_vector_reconciliation.py b/tests/deriver/test_vector_reconciliation.py index bd3ff28d..4ee93e1e 100644 --- a/tests/deriver/test_vector_reconciliation.py +++ b/tests/deriver/test_vector_reconciliation.py @@ -425,9 +425,7 @@ class TestReEmbedding: await db_session.refresh(doc) # Mock embedding client - with patch( - "src.deriver.vector_reconciliation.embedding_client" - ) as mock_embed_client: + with patch("src.reconciler.sync_vectors.embedding_client") as mock_embed_client: mock_embed_client.simple_batch_embed = AsyncMock( return_value=[[float(i)] * 1536 for i in range(3)] ) @@ -501,9 +499,7 @@ class TestReEmbedding: batch_call_count += 1 return [[1.0] * 1536 for _ in contents] - with patch( - "src.deriver.vector_reconciliation.embedding_client" - ) as mock_embed_client: + with patch("src.reconciler.sync_vectors.embedding_client") as mock_embed_client: mock_embed_client.simple_batch_embed = track_batch_embed # Mock vector store @@ -785,13 +781,13 @@ class TestEndToEndReconciliation: # This would be an integration test with the full cycle # For now, we verify the function signature and return type with ( - patch("src.deriver.vector_reconciliation.tracked_db") as mock_tracked_db, - patch("src.deriver.vector_reconciliation.get_external_vector_store"), + patch("src.reconciler.sync_vectors.tracked_db") as mock_tracked_db, + patch("src.reconciler.sync_vectors.get_external_vector_store"), patch( - "src.deriver.vector_reconciliation._get_documents_needing_sync" + "src.reconciler.sync_vectors._get_documents_needing_sync" ) as mock_get_docs, patch( - "src.deriver.vector_reconciliation._get_message_embeddings_needing_sync" + "src.reconciler.sync_vectors._get_message_embeddings_needing_sync" ) as mock_get_embs, patch("src.crud.document.cleanup_soft_deleted_documents") as mock_cleanup, ):