fix(embeddings): guard embed_batch count mismatch in OpenAI and Azure OpenAI (#5966)
This commit is contained in:
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user