diff --git a/scripts/generate_message_embeddings.py b/scripts/generate_message_embeddings.py index fd1b6374..47d088e8 100644 --- a/scripts/generate_message_embeddings.py +++ b/scripts/generate_message_embeddings.py @@ -21,7 +21,6 @@ import sys project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) sys.path.insert(0, project_root) -import tiktoken # noqa: E402 from sqlalchemy import select # noqa: E402 from sqlalchemy.ext.asyncio import AsyncSession # noqa: E402 @@ -90,20 +89,16 @@ async def create_embeddings_for_messages( if not messages: return 0 - # Initialize tiktoken encoding (same as used in MessageCreate schema) - encoding = tiktoken.get_encoding("o200k_base") - - # Prepare data for batch embedding with proper token encoding id_resource_dict = { - message.public_id: ( - message.content, - encoding.encode(message.content), # Properly encode the content - ) + message.public_id: message.content for message in messages + if message.content and message.content.strip() } # Generate embeddings - embedding_dict = await embedding_client.batch_embed(id_resource_dict) + embedding_dict = ( + await embedding_client.batch_embed(id_resource_dict) if id_resource_dict else {} + ) # Create MessageEmbedding objects embedding_objects: list[models.MessageEmbedding] = [] diff --git a/src/crud/message.py b/src/crud/message.py index 3334c6e8..4a177d39 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -262,18 +262,16 @@ async def create_messages( await db.commit() try: if settings.EMBED_MESSAGES: - encoded_message_lookup = { - msg.public_id: orig_msg.encoded_message - for msg, orig_msg in zip(message_objects, messages, strict=True) - } id_resource_dict = { - message.public_id: ( - message.content, - encoded_message_lookup[message.public_id], - ) + message.public_id: message.content for message in message_objects + if message.content and message.content.strip() } - embedding_dict = await embedding_client.batch_embed(id_resource_dict) + embedding_dict = ( + await embedding_client.batch_embed(id_resource_dict) + if id_resource_dict + else {} + ) external_vector_store = get_external_vector_store() diff --git a/src/embedding_client.py b/src/embedding_client.py index d6c8b46d..833980a0 100644 --- a/src/embedding_client.py +++ b/src/embedding_client.py @@ -65,7 +65,10 @@ class _EmbeddingClient: self.max_embedding_tokens = max_input_tokens self.max_batch_size = 2048 # OpenAI batch limit - self.encoding: tiktoken.Encoding = tiktoken.get_encoding("o200k_base") + try: + self.encoding: tiktoken.Encoding = tiktoken.encoding_for_model(self.model) + except KeyError: + self.encoding = tiktoken.get_encoding("o200k_base") self.max_embedding_tokens_per_request: int = max_tokens_per_request @property @@ -156,13 +159,13 @@ class _EmbeddingClient: return embeddings async def batch_embed( - self, id_resource_dict: dict[str, tuple[str, list[int]]] + self, id_resource_dict: dict[str, str] ) -> dict[str, list[list[float]]]: """ Embed multiple texts, chunking long ones and batching API calls. Args: - id_resource_dict: Maps text IDs to (text, encoded_tokens) tuples + id_resource_dict: Maps text IDs to text content Returns: Maps text IDs to lists of embedding vectors (one per chunk) @@ -185,27 +188,29 @@ class _EmbeddingClient: return self._accumulate_embeddings(batch_results) def _prepare_chunks( - self, id_resource_dict: dict[str, tuple[str, list[int]]] + self, id_resource_dict: dict[str, str] ) -> dict[str, list[tuple[str, int]]]: """ Chunk texts that exceed token limits. Args: - id_resource_dict: Maps text IDs to (text, encoded_tokens) tuples + id_resource_dict: Maps text IDs to text content. We tokenize with + the embedding client's own encoding so token IDs match the + decoder vocabulary used by the target embedding API. Returns: Maps text IDs to lists of (chunk_text, token_count) tuples """ - return { - text_id: ( - _chunk_text_with_tokens( - text, encoded_tokens, self.max_embedding_tokens, self.encoding + out: dict[str, list[tuple[str, int]]] = {} + for text_id, text in id_resource_dict.items(): + tokens = self.encoding.encode(text) + if len(tokens) > self.max_embedding_tokens: + out[text_id] = _chunk_text_with_tokens( + text, tokens, self.max_embedding_tokens, self.encoding ) - if len(encoded_tokens) > self.max_embedding_tokens - else [(text, len(encoded_tokens))] - ) - for text_id, (text, encoded_tokens) in id_resource_dict.items() - } + else: + out[text_id] = [(text, len(tokens))] + return out def _create_batches( self, text_chunks: dict[str, list[tuple[str, int]]] @@ -440,7 +445,7 @@ class EmbeddingClient: return await self._get_client().simple_batch_embed(texts) async def batch_embed( - self, id_resource_dict: dict[str, tuple[str, list[int]]] + self, id_resource_dict: dict[str, str] ) -> dict[str, list[list[float]]]: """Embed multiple texts, chunking long ones and batching API calls.""" return await self._get_client().batch_embed(id_resource_dict) diff --git a/tests/conftest.py b/tests/conftest.py index 80389d34..b7d64778 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -481,11 +481,11 @@ def mock_openai_embeddings(request: pytest.FixtureRequest): # Mock the batch_embed method to return content-dependent embeddings async def mock_batch_embed_func( - id_resource_dict: dict[str, tuple[str, list[int]]], + id_resource_dict: dict[str, str], ) -> dict[str, list[list[float]]]: return { - text_id: [_content_to_embedding(resource[0])] - for text_id, resource in id_resource_dict.items() + text_id: [_content_to_embedding(content)] + for text_id, content in id_resource_dict.items() } mock_batch_embed.side_effect = mock_batch_embed_func diff --git a/tests/integration/test_message_embeddings.py b/tests/integration/test_message_embeddings.py index 091e8cf3..6ca0904b 100644 --- a/tests/integration/test_message_embeddings.py +++ b/tests/integration/test_message_embeddings.py @@ -77,6 +77,68 @@ async def test_message_embedding_created_when_setting_enabled( assert embedding_record.peer_name == test_peer.name +@pytest.mark.asyncio +async def test_blank_messages_are_not_sent_for_embedding( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, + mock_openai_embeddings: dict[str, Any], +): + """Blank messages should be persisted but excluded from embedding batches.""" + monkeypatch.setattr("src.config.settings.EMBED_MESSAGES", True) + + test_workspace, test_peer = sample_data + + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + blank_content = " " + nonblank_content = "This message should be embedded" + messages = [ + MessageCreate( + content=blank_content, + peer_id=test_peer.name, + metadata={"test": "blank_embedding"}, + ), + MessageCreate( + content=nonblank_content, + peer_id=test_peer.name, + metadata={"test": "blank_embedding"}, + ), + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert [message.content for message in created_messages] == [ + blank_content, + nonblank_content, + ] + + mock_openai_embeddings["batch_embed"].assert_awaited_once() + batch_arg = mock_openai_embeddings["batch_embed"].await_args.args[0] + assert batch_arg == {created_messages[1].public_id: nonblank_content} + + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id.in_( + [message.public_id for message in created_messages] + ) + ) + result = await db_session.execute(stmt) + embedding_records = list(result.scalars().all()) + + assert len(embedding_records) == 1 + assert embedding_records[0].message_id == created_messages[1].public_id + assert embedding_records[0].content == nonblank_content + + @pytest.mark.asyncio async def test_message_embedding_not_created_when_setting_disabled( db_session: AsyncSession, @@ -492,7 +554,7 @@ async def test_message_chunking_creates_multiple_embeddings( test_message_content = "This is a very long message that should be chunked into multiple pieces because it exceeds the token limit that we set for testing purposes. This message contains many words and should definitely be split into multiple chunks." def mock_batch_embed_chunked( - id_resource_dict: dict[str, tuple[str, list[int]]], + id_resource_dict: dict[str, str], ) -> dict[str, list[list[float]]]: return { text_id: [[0.1] * 1536, [0.2] * 1536, [0.3] * 1536] # 3 chunks per message