fix(embeddings): guard embed_batch count mismatch in OpenAI and Azure OpenAI (#5966)

This commit is contained in:
Bartok
2026-07-01 07:23:28 -06:00
committed by GitHub
parent bc05fd9623
commit 152d1e66f7
8 changed files with 102 additions and 0 deletions
+5
View File
@@ -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;
}
}
+5
View File
@@ -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;
}
}
@@ -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,
@@ -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",
+5
View File
@@ -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
+5
View File
@@ -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
@@ -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"])
@@ -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"])