From 7c155f99a2e3afe664acbfeebaa15b0636b339d2 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Tue, 13 Jan 2026 12:55:27 -0500 Subject: [PATCH] fix: coderabbit comments --- ...19a52b73c60_support_external_embeddings.py | 29 +++++++++----- src/crud/document.py | 12 +++--- src/crud/message.py | 29 ++++++++------ src/deriver/vector_reconciliation.py | 40 ++++++++++--------- src/exceptions.py | 8 ++++ src/vector_store/lancedb.py | 12 ++++-- tests/integration/test_message_embeddings.py | 17 -------- tests/routes/test_conclusions.py | 6 +-- tests/utils/test_agent_tools.py | 3 +- 9 files changed, 86 insertions(+), 70 deletions(-) diff --git a/migrations/versions/119a52b73c60_support_external_embeddings.py b/migrations/versions/119a52b73c60_support_external_embeddings.py index e3cc12af..31b1f09f 100644 --- a/migrations/versions/119a52b73c60_support_external_embeddings.py +++ b/migrations/versions/119a52b73c60_support_external_embeddings.py @@ -79,7 +79,9 @@ def upgrade() -> None: "sync_state", sa.TEXT(), nullable=False, - server_default="pending", # Existing records need reconciliation + server_default=sa.text( + "'pending'" + ), # Existing records need reconciliation ), schema=schema, ) @@ -108,7 +110,7 @@ def upgrade() -> None: "sync_attempts", sa.Integer(), nullable=False, - server_default="0", + server_default=sa.text("0"), ), schema=schema, ) @@ -131,7 +133,9 @@ def upgrade() -> None: "sync_state", sa.TEXT(), nullable=False, - server_default="pending", # Existing records need reconciliation + server_default=sa.text( + "'pending'" + ), # Existing records need reconciliation ), schema=schema, ) @@ -160,7 +164,7 @@ def upgrade() -> None: "sync_attempts", sa.Integer(), nullable=False, - server_default="0", + server_default=sa.text("0"), ), schema=schema, ) @@ -182,13 +186,18 @@ def downgrade() -> None: """Remove deleted_at columns and revert embedding columns.""" inspector = sa.inspect(op.get_bind()) - if column_exists("message_embeddings", "sync_state", inspector): - # Drop composite index first + # Drop message_embeddings indexes if they exist + if index_exists( + "message_embeddings", "ix_message_embeddings_sync_state_last_sync_at", inspector + ): op.drop_index( "ix_message_embeddings_sync_state_last_sync_at", table_name="message_embeddings", schema=schema, ) + if index_exists( + "message_embeddings", "ix_message_embeddings_sync_state", inspector + ): op.drop_index( "ix_message_embeddings_sync_state", table_name="message_embeddings", @@ -204,13 +213,14 @@ def downgrade() -> None: if column_exists("message_embeddings", "last_sync_at", inspector): op.drop_column("message_embeddings", "last_sync_at", schema=schema) - if column_exists("documents", "sync_state", inspector): - # Drop composite index first + # Drop documents indexes if they exist + if index_exists("documents", "ix_documents_sync_state_last_sync_at", inspector): op.drop_index( "ix_documents_sync_state_last_sync_at", table_name="documents", schema=schema, ) + if index_exists("documents", "ix_documents_sync_state", inspector): op.drop_index("ix_documents_sync_state", table_name="documents", schema=schema) if column_exists("documents", "sync_state", inspector): @@ -224,8 +234,9 @@ def downgrade() -> None: op.drop_column("documents", "last_sync_at", schema=schema) # Remove deleted_at column and index from documents - if column_exists("documents", "deleted_at", inspector): + if index_exists("documents", "ix_documents_deleted_at", inspector): op.drop_index("ix_documents_deleted_at", table_name="documents", schema=schema) + if column_exists("documents", "deleted_at", inspector): op.drop_column("documents", "deleted_at", schema=schema) # NOTE: This downgrade does NOT restore the NOT NULL constraint on embedding columns diff --git a/src/crud/document.py b/src/crud/document.py index ac452378..aa8697e1 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -467,9 +467,9 @@ async def create_documents( ) await db.commit() - except Exception as e: + except Exception: # Failed after retries - increment sync_attempts for reconciliation - logger.error(f"Failed to upsert vectors after retries: {e}") + logger.exception("Failed to upsert vectors after retries") await db.execute( update(models.Document) .where(models.Document.id.in_(doc_ids)) @@ -752,10 +752,10 @@ async def create_observations( ) await db.commit() - except Exception as e: + except Exception: # Failed after retries - increment sync_attempts for reconciliation - logger.error( - f"Failed to upsert vectors for {namespace} after retries: {e}" + logger.exception( + f"Failed to upsert vectors for {namespace} after retries" ) await db.execute( update(models.Document) @@ -925,6 +925,8 @@ async def cleanup_soft_deleted_documents( return len(successfully_deleted_ids) # No documents were successfully deleted from vector store + # Release FOR UPDATE locks by rolling back the transaction + await db.rollback() return 0 diff --git a/src/crud/message.py b/src/crud/message.py index 581ff431..040374e0 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -221,9 +221,11 @@ async def create_messages( # Create MessageEmbedding entries embedding_objects: list[models.MessageEmbedding] = [] + # Maps emb index -> (chunk_position, embedding vector) + pending_embedding_data: dict[int, tuple[int, list[float]]] = {} for message_obj in message_objects: embeddings = embedding_dict.get(message_obj.public_id, []) - for chunk_idx, embedding in enumerate(embeddings): + for chunk_position, embedding in enumerate(embeddings): embedding_obj = models.MessageEmbedding( content=message_obj.content, message_id=message_obj.public_id, @@ -233,13 +235,14 @@ async def create_messages( sync_state="pending", embedding=embedding if store_embeddings_in_postgres else None, ) - embedding_obj._chunk_index = chunk_idx - embedding_obj._pending_embedding = embedding + emb_idx = len(embedding_objects) + pending_embedding_data[emb_idx] = (chunk_position, embedding) embedding_objects.append(embedding_obj) - # Add MessageEmbedding rows to database only if storing in postgres + # Always create MessageEmbedding rows so reconciliation can track sync state + # even when embeddings aren't stored in postgres embedding_ids: list[int] = [] - if embedding_objects and store_embeddings_in_postgres: + if embedding_objects: db.add_all(embedding_objects) await db.flush() embedding_ids = [emb.id for emb in embedding_objects] @@ -265,15 +268,15 @@ async def create_messages( "message", workspace_name ) - # Build vector records with {message_id}_{chunk_index} as vector ID + # Build vector records with {message_id}_{chunk_position} as vector ID vector_records: list[VectorRecord] = [] - for emb in embedding_objects: - vector_id = f"{emb.message_id}_{emb._chunk_index}" - embedding_data = list(emb._pending_embedding) + for emb_idx, emb in enumerate(embedding_objects): + chunk_position, embedding = pending_embedding_data[emb_idx] + vector_id = f"{emb.message_id}_{chunk_position}" vector_records.append( VectorRecord( id=vector_id, - embedding=embedding_data, + embedding=list(embedding), metadata={ "message_id": emb.message_id, "session_name": emb.session_name, @@ -301,10 +304,10 @@ async def create_messages( ) await db.commit() - except Exception as e: + except Exception: # Failed after retries - increment sync_attempts for reconciliation - logger.error( - f"Failed to upsert message vectors after retries: {e}" + logger.exception( + "Failed to upsert message vectors after retries" ) if embedding_ids: await db.execute( diff --git a/src/deriver/vector_reconciliation.py b/src/deriver/vector_reconciliation.py index fb7cec8c..5edd6612 100644 --- a/src/deriver/vector_reconciliation.py +++ b/src/deriver/vector_reconciliation.py @@ -91,12 +91,14 @@ async def _get_message_embeddings_needing_sync( """ Get message embeddings that need to be synced to the vector store. - Finds embeddings where: - - has an embedding stored in the database - - sync_state is "pending" (never synced or retry needed) - - Note: "synced" = done forever, "failed" = permanent failure (manual intervention) + 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. - Uses FOR UPDATE SKIP LOCKED to prevent concurrent processing. + Uses FOR UPDATE SKIP LOCKED to prevent concurrent processing and + orders by last_sync_at (nulls first) to prioritize never-synced records. + + Note: "synced" = done forever, "failed" = permanent failure (manual intervention) """ stmt = ( select(models.MessageEmbedding) @@ -357,17 +359,19 @@ async def _sync_message_embeddings( ) 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]] = {} + # 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: - message_chunks.setdefault(emb.message_id, []).append(emb) + embeddings_by_message_id.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 + # 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(): @@ -384,8 +388,8 @@ 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}" + # Use {message_id}_{chunk_position} as vector ID (consistent with creation) + vector_id = f"{emb.message_id}_{chunk_position_by_emb_id[emb.id]}" vector_records.append( VectorRecord( @@ -543,8 +547,8 @@ async def run_vector_reconciliation_cycle() -> ReconciliationMetrics: did_work = True if not did_work: - print("No work done, breaking") + logger.debug("No work done, breaking reconciliation loop") break - print("Vector reconciliation cycle completed") + logger.info("Vector reconciliation cycle completed") return metrics diff --git a/src/exceptions.py b/src/exceptions.py index 3a445527..129324b6 100644 --- a/src/exceptions.py +++ b/src/exceptions.py @@ -125,6 +125,14 @@ class SpecialistExecutionError(HonchoException): detail = "Specialist execution failed" +@final +class VectorStoreError(HonchoException): + """Exception raised when a vector store operation fails.""" + + status_code = 500 + detail = "Vector store operation failed" + + class LLMError(Exception): """Exception raised when an LLM call fails. diff --git a/src/vector_store/lancedb.py b/src/vector_store/lancedb.py index 5176b2f6..333b9b6b 100644 --- a/src/vector_store/lancedb.py +++ b/src/vector_store/lancedb.py @@ -14,6 +14,7 @@ import pyarrow as pa from lancedb import AsyncConnection, AsyncTable from src.config import settings +from src.exceptions import VectorStoreError from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult @@ -167,11 +168,13 @@ class LanceDBVectorStore(VectorStore): logger.debug(f"Upserted {len(vectors)} vectors to namespace {namespace}") return VectorUpsertResult(ok=True) - except Exception: + except Exception as e: logger.exception( f"Failed to upsert {len(vectors)} vectors to namespace {namespace}" ) - raise + raise VectorStoreError( + f"Failed to upsert {len(vectors)} vectors to namespace {namespace}" + ) from e async def query( self, @@ -341,7 +344,8 @@ class LanceDBVectorStore(VectorStore): async def close(self) -> None: """Close the LanceDB connection and release resources.""" if self._db is not None: - # LanceDB AsyncConnection doesn't have an explicit close method, - # but we clear the reference to allow garbage collection + # AsyncConnection provides an explicit close() method (synchronous) + # which we invoke to ensure proper cleanup of resources + self._db.close() self._db = None logger.debug("LanceDB connection closed") diff --git a/tests/integration/test_message_embeddings.py b/tests/integration/test_message_embeddings.py index 691d12c1..de544ee0 100644 --- a/tests/integration/test_message_embeddings.py +++ b/tests/integration/test_message_embeddings.py @@ -12,26 +12,12 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from src import models -from src.config import settings from src.crud import create_messages from src.models import Peer, Workspace from src.schemas import MessageCreate from src.utils.search import search -def _stores_embeddings_in_postgres() -> bool: - """Check if current config stores MessageEmbedding rows in postgres.""" - return settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED - - -# Skip tests that depend on MessageEmbedding rows when using external store in migrated mode -requires_postgres_embeddings = pytest.mark.skipif( - not _stores_embeddings_in_postgres(), - reason="MessageEmbedding rows not created when TYPE != 'pgvector' and MIGRATED=true", -) - - -@requires_postgres_embeddings @pytest.mark.asyncio async def test_message_embedding_created_when_setting_enabled( db_session: AsyncSession, @@ -137,7 +123,6 @@ async def test_message_embedding_not_created_when_setting_disabled( assert embedding_record is None -@requires_postgres_embeddings @pytest.mark.asyncio async def test_multiple_message_embeddings_created_when_setting_enabled( db_session: AsyncSession, @@ -197,7 +182,6 @@ async def test_multiple_message_embeddings_created_when_setting_enabled( assert embedding_record.peer_name == test_peer.name -@requires_postgres_embeddings @pytest.mark.asyncio async def test_semantic_search_when_embeddings_enabled( db_session: AsyncSession, @@ -273,7 +257,6 @@ async def test_semantic_search_when_embeddings_enabled( assert created_message.public_id in found_message_ids -@requires_postgres_embeddings @pytest.mark.asyncio async def test_message_chunking_creates_multiple_embeddings( db_session: AsyncSession, diff --git a/tests/routes/test_conclusions.py b/tests/routes/test_conclusions.py index 4cf000cd..245a657b 100644 --- a/tests/routes/test_conclusions.py +++ b/tests/routes/test_conclusions.py @@ -346,11 +346,11 @@ class TestConclusionRoutes: db_session.add(test_session) await db_session.commit() - # Create test observations via API (this populates the vector store) + # Create test conclusions via API (this populates the vector store) _create_response = client.post( - f"/v2/workspaces/{test_workspace.name}/observations", + f"/v2/workspaces/{test_workspace.name}/conclusions", json={ - "observations": [ + "conclusions": [ { "content": "User loves pizza and pasta", "observer_id": test_peer.name, diff --git a/tests/utils/test_agent_tools.py b/tests/utils/test_agent_tools.py index b54098e0..7f135bcb 100644 --- a/tests/utils/test_agent_tools.py +++ b/tests/utils/test_agent_tools.py @@ -11,6 +11,7 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from src import crud, models +from src.config import settings from src.utils.agent_tools import ( ToolContext, _handle_create_observations, # pyright: ignore[reportPrivateUsage] @@ -318,7 +319,7 @@ class TestSearchMemory: ): """Returns observations matching semantic query.""" # Force pgvector queries since test documents are created directly in postgres - monkeypatch.setattr("src.config.settings.VECTOR_STORE.MIGRATED", False) + monkeypatch.setattr(settings.VECTOR_STORE, "MIGRATED", False) ctx = make_tool_context()