fix: steps toward deprecating MessageEmbedding table

This commit is contained in:
Rajat Ahuja 2026-01-12 14:15:42 -05:00
parent 8b356fe546
commit 6c12e6f413
3 changed files with 104 additions and 68 deletions

View File

@ -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))

View File

@ -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:

View File

@ -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,