From 60fde82218e6e8b36efb0c710bc23bf7bb9055a5 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Dec 2025 15:24:07 -0500 Subject: [PATCH] fix: search; protect agaainst failed vector create/delete --- ...a2b3c4d5e6_support_external_embeddings.py} | 37 ++- src/crud/document.py | 230 ++++++++++++++++-- src/crud/message.py | 16 +- src/crud/session.py | 68 ++++++ src/deriver/queue_manager.py | 81 +++++- src/models.py | 3 + src/utils/search.py | 4 +- src/vector_store/__init__.py | 14 +- src/vector_store/lancedb.py | 65 +++-- src/vector_store/turbopuffer.py | 6 +- tests/alembic/revisions/__init__.py | 4 +- ...a2b3c4d5e6_support_external_embeddings.py} | 0 tests/crud/test_document.py | 135 ++++++++++ 13 files changed, 594 insertions(+), 69 deletions(-) rename migrations/versions/{f1a2b3c4d5e6_make_embeddings_nullable.py => f1a2b3c4d5e6_support_external_embeddings.py} (65%) rename tests/alembic/revisions/{test_f1a2b3c4d5e6_make_embeddings_nullable.py => test_f1a2b3c4d5e6_support_external_embeddings.py} (100%) diff --git a/migrations/versions/f1a2b3c4d5e6_make_embeddings_nullable.py b/migrations/versions/f1a2b3c4d5e6_support_external_embeddings.py similarity index 65% rename from migrations/versions/f1a2b3c4d5e6_make_embeddings_nullable.py rename to migrations/versions/f1a2b3c4d5e6_support_external_embeddings.py index 5bdfeb14..f0870153 100644 --- a/migrations/versions/f1a2b3c4d5e6_make_embeddings_nullable.py +++ b/migrations/versions/f1a2b3c4d5e6_support_external_embeddings.py @@ -1,10 +1,12 @@ -"""add chunk_index to message_embeddings and make embeddings nullable +"""add chunk_index to message_embeddings, make embeddings nullable, add soft delete This migration: 1. Adds the chunk_index column to message_embeddings table for tracking chunked message embeddings in external vector stores (turbopuffer/lancedb). 2. Makes embedding columns nullable in both message_embeddings and documents tables since embeddings are now stored in external vector stores instead of PostgreSQL. +3. Adds deleted_at column to documents table for soft delete support, enabling + hybrid sync/soft delete pattern for vector store consistency. Revision ID: f1a2b3c4d5e6 Revises: baa22cad81e2 @@ -30,7 +32,7 @@ schema = get_schema() def upgrade() -> None: - """Add chunk_index column to message_embeddings and make embeddings nullable.""" + """Add chunk_index, make embeddings nullable, add deleted_at for soft delete.""" inspector = sa.inspect(op.get_bind()) # Add chunk_index column to message_embeddings if it doesn't exist @@ -68,11 +70,40 @@ def upgrade() -> None: schema=schema, ) + # Add deleted_at column to documents for soft delete support + # This enables hybrid sync/soft delete pattern: + # - Try to delete from vector store first + # - If successful, hard delete from DB + # - If vector delete fails, soft delete (set deleted_at) and let cleanup job handle it + if not column_exists("documents", "deleted_at", inspector): + op.add_column( + "documents", + sa.Column( + "deleted_at", + sa.DateTime(timezone=True), + nullable=True, + ), + schema=schema, + ) + # Create partial index for efficient cleanup queries (only index non-null values) + op.create_index( + "ix_documents_deleted_at", + "documents", + ["deleted_at"], + schema=schema, + postgresql_where=sa.text("deleted_at IS NOT NULL"), + ) + def downgrade() -> None: - """Remove chunk_index column and revert embedding columns to non-nullable.""" + """Remove chunk_index, deleted_at columns and revert embedding columns.""" inspector = sa.inspect(op.get_bind()) + # Remove deleted_at column and index from documents + if column_exists("documents", "deleted_at", inspector): + op.drop_index("ix_documents_deleted_at", table_name="documents", schema=schema) + op.drop_column("documents", "deleted_at", schema=schema) + # Revert documents.embedding back to nullable=True (it was originally nullable=True) op.alter_column( "documents", diff --git a/src/crud/document.py b/src/crud/document.py index 2fc02205..de072373 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -1,11 +1,14 @@ +import datetime from collections.abc import Sequence from logging import getLogger from typing import Any -from sqlalchemy import delete, select +from sqlalchemy import delete, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql import Select +from sqlalchemy.sql.functions import func +from tenacity import AsyncRetrying, stop_after_attempt, wait_exponential from src import models, schemas from src.config import settings @@ -15,7 +18,7 @@ from src.crud.session import get_session from src.embedding_client import embedding_client from src.exceptions import ResourceNotFoundException, ValidationException from src.utils.filter import apply_filter -from src.vector_store import VectorRecord, get_vector_store +from src.vector_store import VectorRecord, VectorStore, get_vector_store logger = getLogger(__name__) @@ -50,6 +53,7 @@ def get_all_documents( .where(models.Document.workspace_name == workspace_name) .where(models.Document.observer == observer) .where(models.Document.observed == observed) + .where(models.Document.deleted_at.is_(None)) # Exclude soft-deleted ) # Apply additional filters if provided @@ -87,8 +91,10 @@ def get_documents_with_filters( Returns: Select query for documents """ - stmt = select(models.Document).where( - models.Document.workspace_name == workspace_name + stmt = ( + select(models.Document) + .where(models.Document.workspace_name == workspace_name) + .where(models.Document.deleted_at.is_(None)) # Exclude soft-deleted ) # Apply additional filters if provided @@ -172,14 +178,17 @@ async def query_documents( document_ids = [result.id for result in vector_results] # Fetch documents from database - # No additional filtering needed since vector store already applied all supported filters stmt = ( select(models.Document) .where(models.Document.workspace_name == workspace_name) .where(models.Document.observer == observer) .where(models.Document.observed == observed) + .where(models.Document.deleted_at.is_(None)) .where(models.Document.id.in_(document_ids)) ) + # Re-apply all filters at the database layer to catch any constraints + # that aren't supported by the vector store metadata. + stmt = apply_filter(stmt, models.Document, filters) result = await db.execute(stmt) documents = {doc.id: doc for doc in result.scalars().all()} @@ -281,7 +290,20 @@ async def create_documents( }, ) ) - await vector_store.upsert_many(namespace, vector_records) + + # Retry vector upsert with exponential backoff + try: + async for attempt in AsyncRetrying( + stop=stop_after_attempt(3), + wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0), + reraise=True, + ): + with attempt: + await vector_store.upsert_many(namespace, vector_records) + except Exception as e: + # Final attempt failed - log but don't raise + # Documents exist in DB, vectors can be added manually later + logger.error(f"Failed to upsert vectors after retries: {e}") except IntegrityError as e: await db.rollback() @@ -302,7 +324,11 @@ async def delete_document( session_name: str | None = None, ) -> None: """ - Delete a single document by ID. + Delete a single document by ID using hybrid sync/soft delete pattern. + + Tries to delete from vector store first, then hard deletes from DB. + If vector store delete fails, soft deletes (sets deleted_at) and lets + cleanup job handle vector deletion later. Args: db: Database session @@ -315,25 +341,53 @@ async def delete_document( Raises: ResourceNotFoundException: If document not found or doesn't match criteria """ - stmt = delete(models.Document).where( + # Build base query conditions + conditions = [ models.Document.id == document_id, models.Document.workspace_name == workspace_name, models.Document.observer == observer, models.Document.observed == observed, - ) - - # If session is specified, ensure document belongs to that session + models.Document.deleted_at.is_(None), # Only delete non-deleted docs + ] if session_name is not None: - stmt = stmt.where(models.Document.session_name == session_name) + conditions.append(models.Document.session_name == session_name) - result = await db.execute(stmt) - await db.commit() + # Check document exists first + check_stmt = select(models.Document).where(*conditions) + result = await db.execute(check_stmt) + doc = result.scalar_one_or_none() - if result.rowcount == 0: + if doc is None: raise ResourceNotFoundException( f"Document {document_id} not found or does not belong to the specified collection/session" ) + # Try to delete from vector store first + vector_store = get_vector_store() + namespace = vector_store.get_document_namespace(workspace_name, observer, observed) + vector_deleted = False + + try: + await vector_store.delete_many(namespace, [document_id]) + vector_deleted = True + except Exception as e: + logger.warning(f"Failed to delete vector for document {document_id}: {e}") + + if vector_deleted: + # Happy path: hard delete from DB + delete_stmt = delete(models.Document).where(models.Document.id == document_id) + await db.execute(delete_stmt) + else: + # Fallback: soft delete, let cleanup job handle vector + update_stmt = ( + update(models.Document) + .where(models.Document.id == document_id) + .values(deleted_at=func.now()) + ) + await db.execute(update_stmt) + + await db.commit() + async def delete_document_by_id( db: AsyncSession, @@ -341,7 +395,11 @@ async def delete_document_by_id( document_id: str, ) -> None: """ - Delete a single document by ID and workspace. + Delete a single document by ID and workspace using hybrid sync/soft delete pattern. + + Tries to delete from vector store first, then hard deletes from DB. + If vector store delete fails, soft deletes (sets deleted_at) and lets + cleanup job handle vector deletion later. Args: db: Database session @@ -351,19 +409,48 @@ async def delete_document_by_id( Raises: ResourceNotFoundException: If document not found or doesn't belong to the workspace """ - stmt = delete(models.Document).where( + # Fetch document to get observer/observed for namespace + stmt = select(models.Document).where( models.Document.id == document_id, models.Document.workspace_name == workspace_name, + models.Document.deleted_at.is_(None), # Only delete non-deleted docs ) - result = await db.execute(stmt) - await db.commit() + doc = result.scalar_one_or_none() - if result.rowcount == 0: + if doc is None: raise ResourceNotFoundException( f"Document {document_id} not found or does not belong to workspace {workspace_name}" ) + # Try to delete from vector store first + vector_store = get_vector_store() + namespace = vector_store.get_document_namespace( + workspace_name, doc.observer, doc.observed + ) + vector_deleted = False + + try: + await vector_store.delete_many(namespace, [document_id]) + vector_deleted = True + except Exception as e: + logger.warning(f"Failed to delete vector for document {document_id}: {e}") + + if vector_deleted: + # Happy path: hard delete from DB + delete_stmt = delete(models.Document).where(models.Document.id == document_id) + await db.execute(delete_stmt) + else: + # Fallback: soft delete, let cleanup job handle vector + update_stmt = ( + update(models.Document) + .where(models.Document.id == document_id) + .values(deleted_at=func.now()) + ) + await db.execute(update_stmt) + + await db.commit() + async def create_observations( db: AsyncSession, @@ -479,7 +566,22 @@ async def create_observations( }, ) ) - await vector_store.upsert_many(namespace, vector_records) + + # Retry vector upsert with exponential backoff + try: + async for attempt in AsyncRetrying( + stop=stop_after_attempt(3), + wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0), + reraise=True, + ): + with attempt: + await vector_store.upsert_many(namespace, vector_records) + except Exception as e: + # Final attempt failed - log but don't raise + # Documents exist in DB, vectors can be added manually later + logger.error( + f"Failed to upsert vectors for {namespace} after retries: {e}" + ) except IntegrityError as e: await db.rollback() @@ -566,3 +668,89 @@ async def is_rejected_duplicate( f"[DUPLICATE DETECTION] Rejecting new in favor of existing. new='{doc.content}', existing='{existing_doc.content}'." ) return True + + +async def cleanup_soft_deleted_documents( + db: AsyncSession, + vector_store: VectorStore, + batch_size: int = 100, + older_than_minutes: int = 5, +) -> int: + """ + Clean up soft-deleted documents by deleting from vector store and hard deleting from DB. + + This function is designed to be called periodically (e.g., every 5 minutes) to reconcile + any documents that were soft-deleted when the vector store was unavailable. + + Steps: + 1. Find documents with deleted_at older than threshold + 2. Group by namespace (workspace/observer/observed) + 3. Delete from vector store (per namespace) + 4. Hard delete from DB only for documents where vector deletion succeeded + + If vector deletion fails for a namespace, those documents remain soft-deleted + and will be retried on the next cleanup run. + + Args: + db: Database session + vector_store: Vector store instance + batch_size: Maximum number of documents to process per call + older_than_minutes: Only process documents soft-deleted more than this many minutes ago + + Returns: + Count of documents cleaned up (only those where vector deletion succeeded) + """ + cutoff = datetime.datetime.now(datetime.UTC) - datetime.timedelta( + minutes=older_than_minutes + ) + + # Find soft-deleted documents ready for cleanup + # Use FOR UPDATE SKIP LOCKED to prevent multiple deriver instances from + # processing the same documents simultaneously + stmt = ( + select(models.Document) + .where(models.Document.deleted_at.is_not(None)) + .where(models.Document.deleted_at < cutoff) + .limit(batch_size) + .with_for_update(skip_locked=True) + ) + result = await db.execute(stmt) + documents = list(result.scalars().all()) + + if not documents: + return 0 + + # Group by namespace for batch vector deletion + by_namespace: dict[str, list[str]] = {} + for doc in documents: + namespace = vector_store.get_document_namespace( + doc.workspace_name, doc.observer, doc.observed + ) + by_namespace.setdefault(namespace, []).append(doc.id) + + # Delete from vector store (per namespace) and track successful deletions + successfully_deleted_ids: set[str] = set() + for namespace, ids in by_namespace.items(): + try: + await vector_store.delete_many(namespace, ids) + # Only add to successfully_deleted_ids if vector deletion succeeded + successfully_deleted_ids.update(ids) + except Exception as e: + # Log but continue - vectors may already be deleted or namespace may not exist + logger.warning(f"Failed to delete vectors from {namespace}: {e}") + + # Only hard delete documents where vector deletion succeeded + if successfully_deleted_ids: + await db.execute( + delete(models.Document).where( + models.Document.id.in_(successfully_deleted_ids) + ) + ) + await db.commit() + logger.debug( + f"Cleaned up {len(successfully_deleted_ids)} soft-deleted documents" + ) + return len(successfully_deleted_ids) + + # No documents were successfully deleted from vector store + return 0 diff --git a/src/crud/message.py b/src/crud/message.py index 6a05edc0..5a0aead7 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -4,6 +4,7 @@ from typing import Any from nanoid import generate as generate_nanoid from sqlalchemy import ColumnElement, Select, and_, func, select, text from sqlalchemy.ext.asyncio import AsyncSession +from tenacity import AsyncRetrying, stop_after_attempt, wait_exponential from src import models, schemas from src.config import settings @@ -179,9 +180,20 @@ async def create_messages( db.add_all(embedding_objects) await db.commit() - # Upsert vectors to external vector store + # Upsert vectors to external vector store with retry if vector_records: - await vector_store.upsert_many(namespace, vector_records) + try: + async for attempt in AsyncRetrying( + stop=stop_after_attempt(3), + wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0), + reraise=True, + ): + with attempt: + await vector_store.upsert_many(namespace, vector_records) + except Exception as e: + # Final attempt failed - log but don't raise + # MessageEmbedding records exist in DB, vectors can be added later + logger.error(f"Failed to upsert message vectors after retries: {e}") except Exception: logger.exception( diff --git a/src/crud/session.py b/src/crud/session.py index e447ba62..dfbd6cf6 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -18,6 +18,7 @@ from src.exceptions import ( ResourceNotFoundException, ) from src.utils.filter import apply_filter +from src.vector_store import get_vector_store from .peer import get_or_create_peers, get_peer from .workspace import get_or_create_workspace @@ -420,6 +421,36 @@ async def delete_session( ) ) + # Delete message vectors from vector store before deleting DB records + # Fetch all MessageEmbedding records to build vector IDs + embedding_result = await db.execute( + select( + models.MessageEmbedding.message_id, models.MessageEmbedding.chunk_index + ).where( + models.MessageEmbedding.session_name == session_name, + models.MessageEmbedding.workspace_name == workspace_name, + ) + ) + embeddings = embedding_result.all() + + if embeddings: + # Build vector IDs: {message_id}_{chunk_index} + vector_ids = [f"{e.message_id}_{e.chunk_index}" for e in embeddings] + + # Try to delete from vector store (best effort) + try: + vector_store = get_vector_store() + namespace = vector_store.get_message_namespace(workspace_name) + await vector_store.delete_many(namespace, vector_ids) + logger.debug( + f"Deleted {len(vector_ids)} message vectors for session {session_name}" + ) + except Exception as e: + # Log warning but continue - workspace deletion will clean up eventually + logger.warning( + f"Failed to delete message vectors for session {session_name}: {e}" + ) + # Delete MessageEmbedding entries in batches await _batch_delete_matching( db, @@ -431,6 +462,43 @@ async def delete_session( batch_size=5000, ) + # Delete document vectors from vector store before deleting DB records + # Fetch all Document records to get IDs and namespaces + doc_result = await db.execute( + select( + models.Document.id, + models.Document.observer, + models.Document.observed, + ).where( + models.Document.session_name == session_name, + models.Document.workspace_name == workspace_name, + ) + ) + documents = doc_result.all() + + if documents: + # Group document IDs by namespace (observer/observed) + docs_by_namespace: dict[str, list[str]] = {} + vector_store = get_vector_store() + for doc in documents: + namespace = vector_store.get_document_namespace( + workspace_name, doc.observer, doc.observed + ) + docs_by_namespace.setdefault(namespace, []).append(doc.id) + + # Try to delete from vector store (best effort, per namespace) + for namespace, doc_ids in docs_by_namespace.items(): + try: + await vector_store.delete_many(namespace, doc_ids) + logger.debug( + f"Deleted {len(doc_ids)} document vectors from {namespace}" + ) + except Exception as e: + # Log warning but continue - workspace deletion will clean up eventually + logger.warning( + f"Failed to delete document vectors from {namespace}: {e}" + ) + # Delete Document entries associated with this session in batches await _batch_delete_matching( db, diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 6a3c080d..f32690f8 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -49,6 +49,10 @@ class WorkerOwnership(NamedTuple): aqs_id: str # The ID of the ActiveQueueSession that the worker is processing +VECTOR_CLEANUP_INTERVAL_SECONDS = 300 # 5 minutes +QUEUE_CLEANUP_INTERVAL_SECONDS = 43200 # 12 hours + + class QueueManager: def __init__(self): self.shutdown_event: asyncio.Event = asyncio.Event() @@ -111,7 +115,7 @@ class QueueManager: ) logger.debug("Signal handlers registered") - # Start background maintenance loop + # Start background maintenance loop (handles both queue cleanup and vector cleanup) try: self._maintenance_task = asyncio.create_task(self._maintenance_loop()) except Exception: @@ -343,30 +347,85 @@ class QueueManager: await db.commit() async def _maintenance_loop(self) -> None: - """Run periodic maintenance tasks on the queue.""" + """ + Run periodic maintenance tasks. + + - Vector cleanup: every 5 minutes (clean up soft-deleted documents) + - Queue cleanup: every 12 hours (remove old processed/errored queue items) + """ + # Track when each task should next run + next_vector_cleanup = datetime.now(timezone.utc) + next_queue_cleanup = datetime.now(timezone.utc) + try: while not self.shutdown_event.is_set(): - try: - await self.cleanup_queue_items() - except Exception: - logger.exception("Error during maintenance cleanup") - if settings.SENTRY.ENABLED: - sentry_sdk.capture_exception() + now = datetime.now(timezone.utc) + + # Run vector cleanup if due + if now >= next_vector_cleanup: + try: + await self._run_vector_cleanup() + except Exception: + logger.exception("Error during vector cleanup") + if settings.SENTRY.ENABLED: + sentry_sdk.capture_exception() + next_vector_cleanup = now + timedelta( + seconds=VECTOR_CLEANUP_INTERVAL_SECONDS + ) + + # Run queue cleanup if due + if now >= next_queue_cleanup: + try: + await self.cleanup_queue_items() + except Exception: + logger.exception("Error during queue cleanup") + if settings.SENTRY.ENABLED: + sentry_sdk.capture_exception() + next_queue_cleanup = now + timedelta( + seconds=QUEUE_CLEANUP_INTERVAL_SECONDS + ) + + # Sleep until next task is due or shutdown + next_task_time = min(next_vector_cleanup, next_queue_cleanup) + sleep_seconds = max( + 0, (next_task_time - datetime.now(timezone.utc)).total_seconds() + ) - # Sleep until interval elapses or shutdown event is set try: await asyncio.wait_for( self.shutdown_event.wait(), - timeout=43200, # 12 hours + timeout=sleep_seconds + or 1, # At least 1 second to avoid busy loop ) break # Shutdown event set except asyncio.TimeoutError: - # Timeout means it's time for next cleanup + # Timeout means it's time for next task pass except asyncio.CancelledError: logger.debug("Maintenance loop cancelled") raise + async def _run_vector_cleanup(self) -> None: + """Run vector store cleanup for soft-deleted documents.""" + from src.crud.document import cleanup_soft_deleted_documents + from src.vector_store import get_vector_store + + async with tracked_db("vector_cleanup") as db: + vector_store = get_vector_store() + total_cleaned = 0 + + # Process in batches until no more soft-deleted documents + while True: + cleaned = await cleanup_soft_deleted_documents(db, vector_store) + total_cleaned += cleaned + if cleaned == 0: + break + + if total_cleaned > 0: + logger.info( + f"Vector cleanup: removed {total_cleaned} soft-deleted documents" + ) + async def _handle_processing_error( self, error: Exception, diff --git a/src/models.py b/src/models.py index 58739766..fea58113 100644 --- a/src/models.py +++ b/src/models.py @@ -389,6 +389,9 @@ class Document(Base): ForeignKey("workspaces.name"), nullable=False, index=True ) session_name: Mapped[str] = mapped_column(TEXT, index=True) + deleted_at: Mapped[datetime.datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True, index=True, default=None + ) collection = relationship("Collection", back_populates="documents") __table_args__ = ( diff --git a/src/utils/search.py b/src/utils/search.py index 015fd4d2..9b9e1bc7 100644 --- a/src/utils/search.py +++ b/src/utils/search.py @@ -130,11 +130,11 @@ async def _semantic_search( message_ids = list(seen_message_ids.keys()) - # Fetch messages from database by the IDs from vector search - # No additional filtering needed since vector store already applied all filters + # Fetch messages from database by the IDs from vector search and reapply filters semantic_query = select(models.Message).where( models.Message.public_id.in_(message_ids) ) + semantic_query = apply_filter(semantic_query, models.Message, filters) result = await db.execute(semantic_query) messages = {msg.public_id: msg for msg in result.scalars().all()} diff --git a/src/vector_store/__init__.py b/src/vector_store/__init__.py index 3f4cf2fb..1220e0c8 100644 --- a/src/vector_store/__init__.py +++ b/src/vector_store/__init__.py @@ -89,11 +89,9 @@ class VectorStore(ABC): Args: namespace: The namespace to store the vector in - id: Unique identifier for the vector - embedding: The embedding vector - metadata: Optional metadata to store with the vector + vector: VectorRecord containing id, embedding, and optional metadata """ - pass + ... @abstractmethod async def upsert_many( @@ -108,7 +106,7 @@ class VectorStore(ABC): namespace: The namespace to store the vectors in vectors: List of VectorRecord objects to upsert """ - pass + ... @abstractmethod async def query( @@ -133,7 +131,7 @@ class VectorStore(ABC): Returns: List of QueryResult objects, ordered by similarity (most similar first) """ - pass + ... @abstractmethod async def delete_many(self, namespace: str, ids: list[str]) -> None: @@ -144,7 +142,7 @@ class VectorStore(ABC): namespace: The namespace containing the vectors ids: List of vector identifiers to delete """ - pass + ... @abstractmethod async def delete_namespace(self, namespace: str) -> None: @@ -154,7 +152,7 @@ class VectorStore(ABC): Args: namespace: The namespace to delete """ - pass + ... # Singleton instance diff --git a/src/vector_store/lancedb.py b/src/vector_store/lancedb.py index 95ea5691..954dfe17 100644 --- a/src/vector_store/lancedb.py +++ b/src/vector_store/lancedb.py @@ -23,6 +23,8 @@ logger = logging.getLogger(__name__) # Additional metadata columns are added dynamically VECTOR_DIMENSION = 1536 +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false, reportUnknownParameterType=false + class LanceDBVectorStore(VectorStore): """ @@ -35,7 +37,7 @@ class LanceDBVectorStore(VectorStore): _db: AsyncConnection | None = None _db_path: str - def __init__(self): + def __init__(self) -> None: """Initialize the LanceDB vector store.""" super().__init__() self._db_path = settings.VECTOR_STORE.LANCEDB_PATH @@ -56,14 +58,14 @@ class LanceDBVectorStore(VectorStore): return None async def _get_or_create_table( - self, namespace: str, sample_data: list[dict[str, Any]] | None = None + self, + namespace: str, ) -> AsyncTable: """ Get existing table or create if not exists. Args: namespace: Table name (namespace) - sample_data: Optional sample data to infer schema from Returns: LanceDB async table @@ -73,18 +75,46 @@ class LanceDBVectorStore(VectorStore): if namespace in table_names: return await db.open_table(namespace) - # Create table with sample data if provided - if sample_data: - return await db.create_table(namespace, data=sample_data) - # Create empty table with base schema - schema = pa.schema( # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] - [ - pa.field("id", pa.string()), # pyright: ignore[reportUnknownMemberType] - pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)), # pyright: ignore[reportUnknownMemberType] + fields: list[pa.Field] = [ + pa.field("id", pa.string()), + pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)), + ] + 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 + + def _metadata_fields_for_namespace(self, namespace: str) -> list[pa.Field]: + """ + Infer standard metadata columns based on namespace structure. + + Namespaces: + - Documents: {prefix}.{workspace}.{observer}.{observed} + - Messages: {prefix}.{workspace}.messages + """ + parts = namespace.split(".") + if len(parts) < 3: + return [] + + if parts[-1] == "messages": + return [ + pa.field("message_id", pa.string(), nullable=True), + pa.field("session_name", pa.string(), nullable=True), + pa.field("peer_name", pa.string(), nullable=True), + pa.field("chunk_index", pa.int64(), nullable=True), ] - ) - return await db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType] + + if len(parts) == 4: + return [ + pa.field("workspace_name", pa.string(), nullable=True), + pa.field("observer", pa.string(), nullable=True), + pa.field("observed", pa.string(), nullable=True), + pa.field("session_name", pa.string(), nullable=True), + pa.field("level", pa.string(), nullable=True), + ] + + return [] def _row_to_dict(self, vector: VectorRecord) -> dict[str, Any]: """Convert a VectorRecord to a dict for LanceDB.""" @@ -94,7 +124,10 @@ class LanceDBVectorStore(VectorStore): } # Add metadata fields if vector.metadata: - row.update(vector.metadata) + reserved_keys = {"id", "vector", "_distance"} + for key in vector.metadata: + if key not in reserved_keys: + row[key] = vector.metadata[key] return row async def upsert( @@ -111,7 +144,7 @@ class LanceDBVectorStore(VectorStore): """ try: row = self._row_to_dict(vector) - table = await self._get_or_create_table(namespace, sample_data=[row]) + table = await self._get_or_create_table(namespace) # Use merge_insert for upsert behavior await ( @@ -145,7 +178,7 @@ class LanceDBVectorStore(VectorStore): try: rows = [self._row_to_dict(v) for v in vectors] - table = await self._get_or_create_table(namespace, sample_data=rows) + table = await self._get_or_create_table(namespace) # Use merge_insert for upsert behavior await ( diff --git a/src/vector_store/turbopuffer.py b/src/vector_store/turbopuffer.py index d6676612..7d745d97 100644 --- a/src/vector_store/turbopuffer.py +++ b/src/vector_store/turbopuffer.py @@ -70,9 +70,7 @@ class TurbopufferVectorStore(VectorStore): Args: namespace: The namespace to store the vector in - id: Unique identifier for the vector - embedding: The embedding vector - metadata: Optional metadata to store with the vector + vector: VectorRecord containing id, embedding, and optional metadata """ ns = self._get_namespace(namespace) attributes = vector.metadata or {} @@ -116,7 +114,7 @@ class TurbopufferVectorStore(VectorStore): { "id": v.id, "vector": v.embedding, - **v.metadata, + **(v.metadata or {}), } for v in vectors ] diff --git a/tests/alembic/revisions/__init__.py b/tests/alembic/revisions/__init__.py index cea9ff97..a995d2ca 100644 --- a/tests/alembic/revisions/__init__.py +++ b/tests/alembic/revisions/__init__.py @@ -22,7 +22,7 @@ from . import ( test_d429de0e5338_adopt_peer_paradigm, test_e9b705f9adf9_add_server_defaults_to_timestamp_, test_ec8f94139b02_codify_workspace_name_and_message_id_in_, - test_f1a2b3c4d5e6_make_embeddings_nullable, + test_f1a2b3c4d5e6_support_external_embeddings, ) __all__ = [ @@ -47,5 +47,5 @@ __all__ = [ "test_d429de0e5338_adopt_peer_paradigm", "test_e9b705f9adf9_add_server_defaults_to_timestamp_", "test_ec8f94139b02_codify_workspace_name_and_message_id_in_", - "test_f1a2b3c4d5e6_make_embeddings_nullable", + "test_f1a2b3c4d5e6_support_external_embeddings", ] diff --git a/tests/alembic/revisions/test_f1a2b3c4d5e6_make_embeddings_nullable.py b/tests/alembic/revisions/test_f1a2b3c4d5e6_support_external_embeddings.py similarity index 100% rename from tests/alembic/revisions/test_f1a2b3c4d5e6_make_embeddings_nullable.py rename to tests/alembic/revisions/test_f1a2b3c4d5e6_support_external_embeddings.py diff --git a/tests/crud/test_document.py b/tests/crud/test_document.py index 46533da1..c0599dac 100644 --- a/tests/crud/test_document.py +++ b/tests/crud/test_document.py @@ -1,3 +1,5 @@ +import datetime + import pytest from nanoid import generate as generate_nanoid from sqlalchemy import select @@ -139,6 +141,139 @@ class TestDocumentCRUD: assert len(results) == 2 + @pytest.mark.asyncio + async def test_query_documents_excludes_soft_deleted( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Query results should not include soft-deleted documents even if vectors remain""" + test_workspace, test_peer = sample_data + test_peer2, test_session, _ = await self._setup_test_data( + db_session, test_workspace, test_peer + ) + + # Create two documents and persist embeddings + doc_schemas = [ + schemas.DocumentCreate( + content="User likes pizza", + embedding=[0.9] * 1536, + session_name=test_session.name, + metadata=schemas.DocumentMetadata( + message_ids=[1], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + schemas.DocumentCreate( + content="User dislikes vegetables", + embedding=[0.1] * 1536, + session_name=test_session.name, + metadata=schemas.DocumentMetadata( + message_ids=[2], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + ] + await crud.create_documents( + db_session, + doc_schemas, + workspace_name=test_workspace.name, + observer=test_peer.name, + observed=test_peer2.name, + ) + + # Soft-delete one document without touching vectors + stmt = select(models.Document).where( + models.Document.workspace_name == test_workspace.name, + models.Document.observer == test_peer.name, + models.Document.observed == test_peer2.name, + ) + result = await db_session.execute(stmt) + docs = {doc.content: doc for doc in result.scalars().all()} + deleted_doc = docs["User likes pizza"] + kept_doc = docs["User dislikes vegetables"] + + deleted_doc.deleted_at = datetime.datetime.now(datetime.timezone.utc) + await db_session.commit() + + results = await crud.query_documents( + db_session, + workspace_name=test_workspace.name, + query="food preferences", + observer=test_peer.name, + observed=test_peer2.name, + top_k=10, + ) + + assert len(results) == 1 + assert results[0].id == kept_doc.id + + @pytest.mark.asyncio + async def test_query_documents_applies_additional_filters( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Filters beyond vector metadata should be enforced at the DB layer""" + test_workspace, test_peer = sample_data + test_peer2, test_session, _ = await self._setup_test_data( + db_session, test_workspace, test_peer + ) + + doc_schemas = [ + schemas.DocumentCreate( + content="Observation one", + embedding=[0.5] * 1536, + session_name=test_session.name, + times_derived=1, + metadata=schemas.DocumentMetadata( + message_ids=[1], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + schemas.DocumentCreate( + content="Observation two", + embedding=[0.5] * 1536, + session_name=test_session.name, + times_derived=2, + metadata=schemas.DocumentMetadata( + message_ids=[2], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + ] + await crud.create_documents( + db_session, + doc_schemas, + workspace_name=test_workspace.name, + observer=test_peer.name, + observed=test_peer2.name, + ) + + result = await db_session.execute( + select(models.Document).where( + models.Document.workspace_name == test_workspace.name, + models.Document.observer == test_peer.name, + models.Document.observed == test_peer2.name, + ) + ) + docs = result.scalars().all() + times_derived_map = {doc.times_derived: doc.id for doc in docs} + + results = await crud.query_documents( + db_session, + workspace_name=test_workspace.name, + query="any query", + observer=test_peer.name, + observed=test_peer2.name, + top_k=10, + filters={"times_derived": 2}, + embedding=[0.5] * 1536, + ) + + assert len(results) == 1 + assert results[0].id == times_derived_map[2] + @pytest.mark.asyncio async def test_delete_document_success( self,