From 152d1e66f703f67dd3feb1ff6bedf5fc4ca7d462 Mon Sep 17 00:00:00 2001 From: Bartok Date: Wed, 1 Jul 2026 07:23:28 -0600 Subject: [PATCH] fix(embeddings): guard embed_batch count mismatch in OpenAI and Azure OpenAI (#5966) --- mem0-ts/src/oss/src/embeddings/azure.ts | 5 ++++ mem0-ts/src/oss/src/embeddings/openai.ts | 5 ++++ mem0-ts/src/oss/tests/azure-embedder.test.ts | 13 +++++++++ mem0-ts/src/oss/tests/openai-embedder.test.ts | 15 +++++++++++ mem0/embeddings/azure_openai.py | 5 ++++ mem0/embeddings/openai.py | 5 ++++ .../test_azure_openai_embeddings.py | 27 +++++++++++++++++++ tests/embeddings/test_openai_embeddings.py | 27 +++++++++++++++++++ 8 files changed, 102 insertions(+) diff --git a/mem0-ts/src/oss/src/embeddings/azure.ts b/mem0-ts/src/oss/src/embeddings/azure.ts index 49cbe5413..48c164d32 100644 --- a/mem0-ts/src/oss/src/embeddings/azure.ts +++ b/mem0-ts/src/oss/src/embeddings/azure.ts @@ -52,6 +52,11 @@ export class AzureOpenAIEmbedder implements Embedder { .map((item) => item.embedding), ); } + if (allEmbeddings.length !== texts.length) { + throw new Error( + `Azure OpenAI embedBatch() returned ${allEmbeddings.length} embeddings for ${texts.length} texts using model '${this.model}'`, + ); + } return allEmbeddings; } } diff --git a/mem0-ts/src/oss/src/embeddings/openai.ts b/mem0-ts/src/oss/src/embeddings/openai.ts index dd389886f..4fe90beee 100644 --- a/mem0-ts/src/oss/src/embeddings/openai.ts +++ b/mem0-ts/src/oss/src/embeddings/openai.ts @@ -47,6 +47,11 @@ export class OpenAIEmbedder implements Embedder { .map((item) => item.embedding), ); } + if (allEmbeddings.length !== texts.length) { + throw new Error( + `OpenAI embedBatch() returned ${allEmbeddings.length} embeddings for ${texts.length} texts using model '${this.model}'`, + ); + } return allEmbeddings; } } diff --git a/mem0-ts/src/oss/tests/azure-embedder.test.ts b/mem0-ts/src/oss/tests/azure-embedder.test.ts index 5b6f46ac3..c26971db8 100644 --- a/mem0-ts/src/oss/tests/azure-embedder.test.ts +++ b/mem0-ts/src/oss/tests/azure-embedder.test.ts @@ -134,6 +134,19 @@ describe("AzureOpenAIEmbedder (unit)", () => { expect(result).toEqual(batch); }); + it("embedBatch() throws when provider returns fewer embeddings than texts", async () => { + // Provider returns only 1 embedding for 2 texts (short-batch response). + mockEmbeddingsCreate.mockResolvedValue({ + data: [{ index: 0, embedding: mockEmbedding }], + }); + + const embedder = new AzureOpenAIEmbedder(baseConfig); + + await expect(embedder.embedBatch(["text1", "text2"])).rejects.toThrow( + /returned 1 embeddings for 2 texts/, + ); + }); + it("uses custom model when provided", async () => { const embedder = new AzureOpenAIEmbedder({ ...baseConfig, diff --git a/mem0-ts/src/oss/tests/openai-embedder.test.ts b/mem0-ts/src/oss/tests/openai-embedder.test.ts index e5733e4ea..af52516db 100644 --- a/mem0-ts/src/oss/tests/openai-embedder.test.ts +++ b/mem0-ts/src/oss/tests/openai-embedder.test.ts @@ -142,6 +142,21 @@ describe("OpenAIEmbedder (unit)", () => { expect(result).toEqual(batch); }); + it("embedBatch() throws when provider returns fewer embeddings than texts", async () => { + // Provider returns only 1 embedding for 2 texts (short-batch response). + mockEmbeddingsCreate.mockResolvedValue({ + data: [{ index: 0, embedding: mockEmbedding }], + }); + + const embedder = new OpenAIEmbedder({ + apiKey: "test-key", + }); + + await expect(embedder.embedBatch(["text1", "text2"])).rejects.toThrow( + /returned 1 embeddings for 2 texts/, + ); + }); + it("uses custom model when provided", async () => { const embedder = new OpenAIEmbedder({ apiKey: "test-key", diff --git a/mem0/embeddings/azure_openai.py b/mem0/embeddings/azure_openai.py index 95670a8a1..e07fb042e 100644 --- a/mem0/embeddings/azure_openai.py +++ b/mem0/embeddings/azure_openai.py @@ -69,4 +69,9 @@ class AzureOpenAIEmbedding(EmbeddingBase): model=self.config.model, ) all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index)) + if len(all_embeddings) != len(texts): + raise ValueError( + f"Azure OpenAI embed_batch() returned {len(all_embeddings)} embeddings for {len(texts)} texts" + f" using model '{self.config.model}'" + ) return all_embeddings diff --git a/mem0/embeddings/openai.py b/mem0/embeddings/openai.py index cede4f1cb..20f1cdf24 100644 --- a/mem0/embeddings/openai.py +++ b/mem0/embeddings/openai.py @@ -73,4 +73,9 @@ class OpenAIEmbedding(EmbeddingBase): kwargs["dimensions"] = self.config.embedding_dims response = self.client.embeddings.create(**kwargs) all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index)) + if len(all_embeddings) != len(texts): + raise ValueError( + f"OpenAI embed_batch() returned {len(all_embeddings)} embeddings for {len(texts)} texts" + f" using model '{self.config.model}'" + ) return all_embeddings diff --git a/tests/embeddings/test_azure_openai_embeddings.py b/tests/embeddings/test_azure_openai_embeddings.py index c45e76c53..08b9a3fbe 100644 --- a/tests/embeddings/test_azure_openai_embeddings.py +++ b/tests/embeddings/test_azure_openai_embeddings.py @@ -164,3 +164,30 @@ def test_init_with_placeholder_api_key(monkeypatch, base_embedder_config): http_client=None, default_headers=None, ) + + +def test_embed_batch_returns_all_embeddings(mock_openai_client): + config = BaseEmbedderConfig(model="text-embedding-ada-002") + embedder = AzureOpenAIEmbedding(config) + mock_response = Mock() + mock_response.data = [ + Mock(index=0, embedding=[0.1, 0.2]), + Mock(index=1, embedding=[0.3, 0.4]), + ] + mock_openai_client.embeddings.create.return_value = mock_response + + result = embedder.embed_batch(["first text", "second text"]) + + assert result == [[0.1, 0.2], [0.3, 0.4]] + + +def test_embed_batch_count_mismatch_raises(mock_openai_client): + config = BaseEmbedderConfig(model="text-embedding-ada-002") + embedder = AzureOpenAIEmbedding(config) + # Provider returns fewer embeddings than inputs (partial/dropped batch). + mock_response = Mock() + mock_response.data = [Mock(index=0, embedding=[0.1, 0.2])] + mock_openai_client.embeddings.create.return_value = mock_response + + with pytest.raises(ValueError, match="returned 1 embeddings for 2 texts"): + embedder.embed_batch(["first text", "second text"]) diff --git a/tests/embeddings/test_openai_embeddings.py b/tests/embeddings/test_openai_embeddings.py index 4291d1a91..4d1088b50 100644 --- a/tests/embeddings/test_openai_embeddings.py +++ b/tests/embeddings/test_openai_embeddings.py @@ -122,3 +122,30 @@ def test_embed_passes_dimensions_only_when_explicit(mock_openai_client): mock_openai_client.embeddings.create.assert_called_once_with( input=["truncate me"], model="text-embedding-3-small", dimensions=256, encoding_format="float" ) + + +def test_embed_batch_returns_all_embeddings(mock_openai_client): + config = BaseEmbedderConfig() + embedder = OpenAIEmbedding(config) + mock_response = Mock() + mock_response.data = [ + Mock(index=0, embedding=[0.1, 0.2]), + Mock(index=1, embedding=[0.3, 0.4]), + ] + mock_openai_client.embeddings.create.return_value = mock_response + + result = embedder.embed_batch(["first text", "second text"]) + + assert result == [[0.1, 0.2], [0.3, 0.4]] + + +def test_embed_batch_count_mismatch_raises(mock_openai_client): + config = BaseEmbedderConfig() + embedder = OpenAIEmbedding(config) + # Provider returns fewer embeddings than inputs (partial/dropped batch). + mock_response = Mock() + mock_response.data = [Mock(index=0, embedding=[0.1, 0.2])] + mock_openai_client.embeddings.create.return_value = mock_response + + with pytest.raises(ValueError, match="returned 1 embeddings for 2 texts"): + embedder.embed_batch(["first text", "second text"])