diff --git a/src/crud/message.py b/src/crud/message.py index 31f7e6e8..cfa9e0a6 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -154,74 +154,85 @@ async def create_messages( for message_obj in message_objects: embeddings = embedding_dict.get(message_obj.public_id, []) - for embedding in embeddings: - # Create MessageEmbedding record - embedding_obj = models.MessageEmbedding( - content=message_obj.content, - message_id=message_obj.public_id, - workspace_name=workspace_name, - session_name=session_name, - peer_name=message_obj.peer_name, - sync_state="pending", - ) + for chunk_idx, embedding in enumerate(embeddings): if pgvector_in_use: - # pgvector in use: write embedding to ORM (postgres) - embedding_obj.embedding = embedding + # Create MessageEmbedding record for pgvector storage + embedding_obj = models.MessageEmbedding( + content=message_obj.content, + message_id=message_obj.public_id, + workspace_name=workspace_name, + session_name=session_name, + peer_name=message_obj.peer_name, + sync_state="pending", + embedding=embedding, + ) + # Track chunk index in-memory (not persisted to avoid HNSW recomputation) + embedding_obj._chunk_index = chunk_idx + embedding_objects.append(embedding_obj) else: - # store in memory for vector store upsert only + # pgvector not in use: don't create MessageEmbedding, track in-memory only + # Create a minimal object to hold metadata for vector store upsert + embedding_obj = models.MessageEmbedding( + content=message_obj.content, + message_id=message_obj.public_id, + workspace_name=workspace_name, + session_name=session_name, + peer_name=message_obj.peer_name, + sync_state="pending", + ) embedding_obj._pending_embedding = embedding - embedding_objects.append(embedding_obj) + embedding_obj._chunk_index = chunk_idx + embedding_objects.append(embedding_obj) - # Add all embedding metadata objects to the session - if embedding_objects: + # Add MessageEmbedding rows to database only if pgvector in use + if embedding_objects and pgvector_in_use: db.add_all(embedding_objects) await db.flush() # Track embedding IDs for sync state updates embedding_ids = [emb.id for emb in embedding_objects] + else: + embedding_ids = [] - # Build vector records - source depends on whether pgvector is in use - vector_records: list[VectorRecord] = [] - for emb in embedding_objects: - if pgvector_in_use: - # pgvector in use: embedding is on ORM object (numpy array) - if emb.embedding is not None: - vector_records.append( - VectorRecord( - id=str(emb.id), - embedding=[float(x) for x in emb.embedding], - metadata={ - "message_id": emb.message_id, - "session_name": emb.session_name, - "peer_name": emb.peer_name, - }, - ) - ) - else: - # pgvector not in use: embedding is in _pending_embedding - if ( - hasattr(emb, "_pending_embedding") - and emb._pending_embedding is not None - ): - vector_records.append( - VectorRecord( - id=str(emb.id), - embedding=list(emb._pending_embedding), - metadata={ - "message_id": emb.message_id, - "session_name": emb.session_name, - "peer_name": emb.peer_name, - }, - ) - ) - await db.commit() + # Build vector records with {message_id}_{chunk_index} as vector ID + vector_records: list[VectorRecord] = [] + for emb in embedding_objects: + # Always use {message_id}_{chunk_index} as vector ID (all stores) + vector_id = f"{emb.message_id}_{emb._chunk_index}" - # Upsert to vector store with retry and update sync state - if vector_records: - try: - result = await upsert_with_retry( - vector_store, namespace, vector_records - ) + # Get embedding from appropriate source + if pgvector_in_use and emb.embedding is not None: + embedding_data = [float(x) for x in emb.embedding] + elif ( + hasattr(emb, "_pending_embedding") + and emb._pending_embedding is not None + ): + embedding_data = list(emb._pending_embedding) + else: + continue + + vector_records.append( + VectorRecord( + id=vector_id, + embedding=embedding_data, + metadata={ + "message_id": emb.message_id, + "session_name": emb.session_name, + "peer_name": emb.peer_name, + }, + ) + ) + + await db.commit() + + # Upsert to vector store with retry and update sync state + if vector_records: + try: + result = await upsert_with_retry( + vector_store, namespace, vector_records + ) + # Only update MessageEmbedding sync state if pgvector is in use + if pgvector_in_use and embedding_ids: if result is not None and result.secondary_ok is False: # Partial success: primary has data but secondary doesn't logger.warning( @@ -251,11 +262,11 @@ async def create_messages( ) await db.commit() - except Exception as e: - # Total failure: primary write failed after retries - logger.error( - f"Failed to upsert message vectors after retries: {e}" - ) + except Exception as e: + # Total failure: primary write failed after retries + logger.error(f"Failed to upsert message vectors after retries: {e}") + # Only update MessageEmbedding sync state if pgvector is in use + if pgvector_in_use and embedding_ids: await db.execute( update(models.MessageEmbedding) .where(models.MessageEmbedding.id.in_(embedding_ids)) diff --git a/src/crud/session.py b/src/crud/session.py index f9f9246a..30f5d505 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -422,18 +422,28 @@ async def delete_session( ) # Delete message vectors from vector store before deleting DB records - # Fetch all MessageEmbedding records to build vector IDs + # Fetch all MessageEmbedding records to build vector IDs with {message_id}_{chunk_index} embedding_result = await db.execute( - select(models.MessageEmbedding.id).where( + select(models.MessageEmbedding).where( models.MessageEmbedding.session_name == session_name, models.MessageEmbedding.workspace_name == workspace_name, ) ) - embeddings = embedding_result.all() + embeddings = list(embedding_result.scalars().all()) vector_store = get_vector_store() if embeddings: - vector_ids = [str(e.id) for e in embeddings] + # Compute chunk_index for each embedding based on message_id ordering + message_chunks: dict[str, list[models.MessageEmbedding]] = {} + for emb in embeddings: + message_chunks.setdefault(emb.message_id, []).append(emb) + + # Sort each message's chunks by id and build vector IDs + vector_ids: list[str] = [] + for chunks in message_chunks.values(): + chunks.sort(key=lambda e: e.id) + for chunk_idx, chunk in enumerate(chunks): + vector_ids.append(f"{chunk.message_id}_{chunk_idx}") # Try to delete from vector store (best effort) try: diff --git a/src/deriver/vector_reconciliation.py b/src/deriver/vector_reconciliation.py index 7f1e7042..1dea0aa1 100644 --- a/src/deriver/vector_reconciliation.py +++ b/src/deriver/vector_reconciliation.py @@ -362,11 +362,23 @@ async def _sync_message_embeddings( namespace = vector_store.get_vector_namespace("message", emb.workspace_name) by_namespace.setdefault(namespace, []).append(emb) + # Compute chunk_index for each embedding based on message_id ordering + # Group embeddings by message_id and assign chunk_index + message_chunks: dict[str, list[models.MessageEmbedding]] = {} + for emb in embeddings: + message_chunks.setdefault(emb.message_id, []).append(emb) + + # Sort each message's chunks by id and assign chunk_index + for chunks in message_chunks.values(): + chunks.sort(key=lambda e: e.id) + for chunk_idx, chunk in enumerate(chunks): + chunk._chunk_index = chunk_idx + # Sync each namespace batch for namespace, embs in by_namespace.items(): embs_with_vectors: list[models.MessageEmbedding] = [] try: - # Build vector records + # Build vector records with {message_id}_{chunk_index} format vector_records: list[VectorRecord] = [] for emb in embs: embedding = ( @@ -377,9 +389,12 @@ async def _sync_message_embeddings( if embedding is None: continue + # Use {message_id}_{chunk_index} as vector ID (consistent with creation) + vector_id = f"{emb.message_id}_{emb._chunk_index}" + vector_records.append( VectorRecord( - id=str(emb.id), + id=vector_id, embedding=[float(x) for x in embedding], metadata={ "message_id": emb.message_id,