From 0878e27d22e2fd7c996ab9e0710d9567fe48605f Mon Sep 17 00:00:00 2001 From: unknown Date: Sun, 26 Jul 2026 12:34:40 +0530 Subject: [PATCH] fix(utils): keep token counts when merging batched embed responses Client.embed() sends every call through merge_embed_responses unless you pass batching=False or images, and that path rebuilds ApiMeta from scratch in merge_meta_field. The rebuild only sets api_version, four of the six ApiMetaBilledUnits fields, and warnings. So meta.tokens, meta.cached_tokens, billed_units.images and billed_units.image_tokens are discarded on the way out. The same request returns different metadata depending on whether batching is on: co.embed(texts=texts).meta.tokens # None co.embed(texts=texts, batching=False).meta.tokens # populated Anyone reading meta.tokens for usage accounting silently gets nothing. Sum the missing fields the same way the existing four are summed. cached_tokens sits directly on ApiMeta so it sums off the meta list rather than the billed_units list. tokens stays None when no input meta carried it, so a merged response with no token counts looks exactly like it does today and the existing equality assertions still hold. merge_meta_field names every field it copies by hand, which is why these four went missing when they were added to the models. One of the tests drives its assertions off the model fields instead of a fixed list, so the next field added to ApiMeta fails there rather than silently disappearing from merged responses. --- src/cohere/utils.py | 20 +++++++++++- tests/test_embed_utils.py | 65 +++++++++++++++++++++++++++++++++++++-- 2 files changed, 82 insertions(+), 3 deletions(-) diff --git a/src/cohere/utils.py b/src/cohere/utils.py index 1a23d4b0e..5f28f092d 100644 --- a/src/cohere/utils.py +++ b/src/cohere/utils.py @@ -9,7 +9,7 @@ from fastavro import parse_schema, reader, writer from . import EmbedResponse, EmbeddingsFloatsEmbedResponse, EmbeddingsByTypeEmbedResponse, ApiMeta, \ - EmbedByTypeResponseEmbeddings, ApiMetaBilledUnits, EmbedJob, CreateEmbedJobResponse, Dataset + EmbedByTypeResponseEmbeddings, ApiMetaBilledUnits, ApiMetaTokens, EmbedJob, CreateEmbedJobResponse, Dataset from .datasets import DatasetsCreateResponse, DatasetsGetResponse from .overrides import get_fields @@ -170,19 +170,37 @@ def sum_fields_if_not_none(obj: typing.Any, field: str) -> Optional[int]: def merge_meta_field(metas: typing.List[ApiMeta]) -> ApiMeta: api_version = metas[0].api_version if metas else None billed_units = [meta.billed_units for meta in metas] + images = sum_fields_if_not_none(billed_units, "images") input_tokens = sum_fields_if_not_none(billed_units, "input_tokens") + image_tokens = sum_fields_if_not_none(billed_units, "image_tokens") output_tokens = sum_fields_if_not_none(billed_units, "output_tokens") search_units = sum_fields_if_not_none(billed_units, "search_units") classifications = sum_fields_if_not_none(billed_units, "classifications") + + token_counts = [meta.tokens for meta in metas] + token_input = sum_fields_if_not_none(token_counts, "input_tokens") + token_output = sum_fields_if_not_none(token_counts, "output_tokens") + # Leave tokens unset rather than building an all-None ApiMetaTokens, so a + # merged response is indistinguishable from an unbatched one. + tokens = ApiMetaTokens( + input_tokens=token_input, + output_tokens=token_output, + ) if token_input is not None or token_output is not None else None + + cached_tokens = sum_fields_if_not_none(metas, "cached_tokens") warnings = {warning for meta in metas if meta.warnings for warning in meta.warnings} return ApiMeta( api_version=api_version, billed_units=ApiMetaBilledUnits( + images=images, input_tokens=input_tokens, + image_tokens=image_tokens, output_tokens=output_tokens, search_units=search_units, classifications=classifications ), + tokens=tokens, + cached_tokens=cached_tokens, warnings=list(warnings) ) diff --git a/tests/test_embed_utils.py b/tests/test_embed_utils.py index b522fc576..e5f2f849b 100644 --- a/tests/test_embed_utils.py +++ b/tests/test_embed_utils.py @@ -1,8 +1,9 @@ import unittest from cohere import EmbeddingsByTypeEmbedResponse, EmbedByTypeResponseEmbeddings, ApiMeta, ApiMetaBilledUnits, \ - ApiMetaApiVersion, EmbeddingsFloatsEmbedResponse -from cohere.utils import merge_embed_responses, sum_fields_if_not_none + ApiMetaApiVersion, ApiMetaTokens, EmbeddingsFloatsEmbedResponse +from cohere.overrides import get_fields +from cohere.utils import merge_embed_responses, merge_meta_field, sum_fields_if_not_none ebt_1 = EmbeddingsByTypeEmbedResponse( response_type="embeddings_by_type", @@ -205,6 +206,66 @@ 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_meta_field_keeps_tokens_and_image_units(self) -> None: + merged = merge_meta_field([ + ApiMeta( + api_version=ApiMetaApiVersion(version="1"), + billed_units=ApiMetaBilledUnits(input_tokens=1, images=1, image_tokens=10), + tokens=ApiMetaTokens(input_tokens=11, output_tokens=0), + cached_tokens=3, + ), + ApiMeta( + api_version=ApiMetaApiVersion(version="1"), + billed_units=ApiMetaBilledUnits(input_tokens=2, images=2, image_tokens=20), + tokens=ApiMetaTokens(input_tokens=22, output_tokens=0), + cached_tokens=4, + ), + ]) + + if merged.billed_units is None or merged.tokens is None: + raise Exception("this is just for mypy") + + self.assertEqual(merged.billed_units.input_tokens, 3) + self.assertEqual(merged.billed_units.images, 3) + self.assertEqual(merged.billed_units.image_tokens, 30) + self.assertEqual(merged.tokens.input_tokens, 33) + self.assertEqual(merged.tokens.output_tokens, 0) + self.assertEqual(merged.cached_tokens, 7) + + def test_merge_meta_field_leaves_tokens_unset_when_absent(self) -> None: + merged = merge_meta_field([ + ApiMeta(billed_units=ApiMetaBilledUnits(input_tokens=1)), + ApiMeta(billed_units=ApiMetaBilledUnits(input_tokens=2)), + ]) + + self.assertIsNone(merged.tokens) + self.assertIsNone(merged.cached_tokens) + + def test_merge_meta_field_sums_every_numeric_field_on_the_model(self) -> None: + # merge_meta_field lists the fields it copies by hand, so any field added + # to ApiMeta later is silently dropped from every merged response until + # someone remembers to update it. That is how images, image_tokens, + # tokens and cached_tokens went missing. Drive the assertion off the + # model itself so the next added field fails here instead of in the wild. + billed_fields = get_fields(ApiMetaBilledUnits()) + token_fields = get_fields(ApiMetaTokens()) + meta = ApiMeta( + billed_units=ApiMetaBilledUnits(**{field: 1 for field in billed_fields}), + tokens=ApiMetaTokens(**{field: 1 for field in token_fields}), + cached_tokens=1, + ) + + merged = merge_meta_field([meta, meta]) + + if merged.billed_units is None or merged.tokens is None: + raise Exception("this is just for mypy") + + for field in billed_fields: + self.assertEqual(getattr(merged.billed_units, field), 2, f"billed_units.{field} was dropped") + for field in token_fields: + self.assertEqual(getattr(merged.tokens, field), 2, f"tokens.{field} was dropped") + self.assertEqual(merged.cached_tokens, 2) + 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