From 9806f6455f14a824d742df03299b9391a7643645 Mon Sep 17 00:00:00 2001 From: Oliver Slapinski Date: Sat, 15 Aug 2026 19:20:28 -0400 Subject: [PATCH] fix(utils): preserve embedding types across batches Signed-off-by: Oliver Slapinski --- src/cohere/utils.py | 9 +++++++-- tests/test_embed_utils.py | 14 ++++++++++++++ 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/src/cohere/utils.py b/src/cohere/utils.py index 1a23d4b0e..b1e1c408c 100644 --- a/src/cohere/utils.py +++ b/src/cohere/utils.py @@ -224,8 +224,13 @@ def merge_embed_responses(responses: typing.List[EmbedResponse]) -> EmbedRespons for response in embeddings_type ] - # only get set keys from the pydantic model (i.e. exclude fields that are set to 'None') - fields = [x for x in get_fields(embeddings_type[0].embeddings) if getattr(embeddings_type[0].embeddings, x) is not None] + # Include a type when any batch returned it. Looking only at the first + # response can silently discard valid embeddings from later batches. + fields = [ + field + for field in get_fields(embeddings_type[0].embeddings) + if any(getattr(response.embeddings, field) is not None for response in embeddings_type) + ] merged_dicts = { field: [ diff --git a/tests/test_embed_utils.py b/tests/test_embed_utils.py index b522fc576..552b233ed 100644 --- a/tests/test_embed_utils.py +++ b/tests/test_embed_utils.py @@ -205,6 +205,20 @@ def test_merge_embeddings_by_type_with_none_field_in_later_response(self) -> Non result = merge_embed_responses([resp1, resp2]) self.assertEqual(result.embeddings.float_, [[1.0, 2.0]]) # type: ignore + def test_merge_embeddings_by_type_keeps_field_from_later_response(self) -> None: + resp1 = EmbeddingsByTypeEmbedResponse( + response_type="embeddings_by_type", id="1", + embeddings=EmbedByTypeResponseEmbeddings(float_=[[1.0, 2.0]])) + resp2 = EmbeddingsByTypeEmbedResponse( + response_type="embeddings_by_type", id="2", + embeddings=EmbedByTypeResponseEmbeddings( + float_=[[3.0, 4.0]], int8=[[3, 4]])) + + result = merge_embed_responses([resp1, resp2]) + + self.assertEqual(result.embeddings.float_, [[1.0, 2.0], [3.0, 4.0]]) # type: ignore + self.assertEqual(result.embeddings.int8, [[3, 4]]) # type: ignore + def test_sum_fields_if_not_none_with_none_entries(self) -> None: # billed_units list may contain None when ApiMeta.billed_units is unset; # sum_fields_if_not_none must skip None objects without raising AttributeError