fix: steps toward deprecating MessageEmbedding table
This commit is contained in:
parent
8b356fe546
commit
6c12e6f413
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Reference in New Issue