Skip to content

Commit 0878e27

Browse files
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.
1 parent 050f9c0 commit 0878e27

2 files changed

Lines changed: 82 additions & 3 deletions

File tree

src/cohere/utils.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
from fastavro import parse_schema, reader, writer
1010

1111
from . import EmbedResponse, EmbeddingsFloatsEmbedResponse, EmbeddingsByTypeEmbedResponse, ApiMeta, \
12-
EmbedByTypeResponseEmbeddings, ApiMetaBilledUnits, EmbedJob, CreateEmbedJobResponse, Dataset
12+
EmbedByTypeResponseEmbeddings, ApiMetaBilledUnits, ApiMetaTokens, EmbedJob, CreateEmbedJobResponse, Dataset
1313
from .datasets import DatasetsCreateResponse, DatasetsGetResponse
1414
from .overrides import get_fields
1515

@@ -170,19 +170,37 @@ def sum_fields_if_not_none(obj: typing.Any, field: str) -> Optional[int]:
170170
def merge_meta_field(metas: typing.List[ApiMeta]) -> ApiMeta:
171171
api_version = metas[0].api_version if metas else None
172172
billed_units = [meta.billed_units for meta in metas]
173+
images = sum_fields_if_not_none(billed_units, "images")
173174
input_tokens = sum_fields_if_not_none(billed_units, "input_tokens")
175+
image_tokens = sum_fields_if_not_none(billed_units, "image_tokens")
174176
output_tokens = sum_fields_if_not_none(billed_units, "output_tokens")
175177
search_units = sum_fields_if_not_none(billed_units, "search_units")
176178
classifications = sum_fields_if_not_none(billed_units, "classifications")
179+
180+
token_counts = [meta.tokens for meta in metas]
181+
token_input = sum_fields_if_not_none(token_counts, "input_tokens")
182+
token_output = sum_fields_if_not_none(token_counts, "output_tokens")
183+
# Leave tokens unset rather than building an all-None ApiMetaTokens, so a
184+
# merged response is indistinguishable from an unbatched one.
185+
tokens = ApiMetaTokens(
186+
input_tokens=token_input,
187+
output_tokens=token_output,
188+
) if token_input is not None or token_output is not None else None
189+
190+
cached_tokens = sum_fields_if_not_none(metas, "cached_tokens")
177191
warnings = {warning for meta in metas if meta.warnings for warning in meta.warnings}
178192
return ApiMeta(
179193
api_version=api_version,
180194
billed_units=ApiMetaBilledUnits(
195+
images=images,
181196
input_tokens=input_tokens,
197+
image_tokens=image_tokens,
182198
output_tokens=output_tokens,
183199
search_units=search_units,
184200
classifications=classifications
185201
),
202+
tokens=tokens,
203+
cached_tokens=cached_tokens,
186204
warnings=list(warnings)
187205
)
188206

tests/test_embed_utils.py

Lines changed: 63 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
11
import unittest
22

33
from cohere import EmbeddingsByTypeEmbedResponse, EmbedByTypeResponseEmbeddings, ApiMeta, ApiMetaBilledUnits, \
4-
ApiMetaApiVersion, EmbeddingsFloatsEmbedResponse
5-
from cohere.utils import merge_embed_responses, sum_fields_if_not_none
4+
ApiMetaApiVersion, ApiMetaTokens, EmbeddingsFloatsEmbedResponse
5+
from cohere.overrides import get_fields
6+
from cohere.utils import merge_embed_responses, merge_meta_field, sum_fields_if_not_none
67

78
ebt_1 = EmbeddingsByTypeEmbedResponse(
89
response_type="embeddings_by_type",
@@ -205,6 +206,66 @@ def test_merge_embeddings_by_type_with_none_field_in_later_response(self) -> Non
205206
result = merge_embed_responses([resp1, resp2])
206207
self.assertEqual(result.embeddings.float_, [[1.0, 2.0]]) # type: ignore
207208

209+
def test_merge_meta_field_keeps_tokens_and_image_units(self) -> None:
210+
merged = merge_meta_field([
211+
ApiMeta(
212+
api_version=ApiMetaApiVersion(version="1"),
213+
billed_units=ApiMetaBilledUnits(input_tokens=1, images=1, image_tokens=10),
214+
tokens=ApiMetaTokens(input_tokens=11, output_tokens=0),
215+
cached_tokens=3,
216+
),
217+
ApiMeta(
218+
api_version=ApiMetaApiVersion(version="1"),
219+
billed_units=ApiMetaBilledUnits(input_tokens=2, images=2, image_tokens=20),
220+
tokens=ApiMetaTokens(input_tokens=22, output_tokens=0),
221+
cached_tokens=4,
222+
),
223+
])
224+
225+
if merged.billed_units is None or merged.tokens is None:
226+
raise Exception("this is just for mypy")
227+
228+
self.assertEqual(merged.billed_units.input_tokens, 3)
229+
self.assertEqual(merged.billed_units.images, 3)
230+
self.assertEqual(merged.billed_units.image_tokens, 30)
231+
self.assertEqual(merged.tokens.input_tokens, 33)
232+
self.assertEqual(merged.tokens.output_tokens, 0)
233+
self.assertEqual(merged.cached_tokens, 7)
234+
235+
def test_merge_meta_field_leaves_tokens_unset_when_absent(self) -> None:
236+
merged = merge_meta_field([
237+
ApiMeta(billed_units=ApiMetaBilledUnits(input_tokens=1)),
238+
ApiMeta(billed_units=ApiMetaBilledUnits(input_tokens=2)),
239+
])
240+
241+
self.assertIsNone(merged.tokens)
242+
self.assertIsNone(merged.cached_tokens)
243+
244+
def test_merge_meta_field_sums_every_numeric_field_on_the_model(self) -> None:
245+
# merge_meta_field lists the fields it copies by hand, so any field added
246+
# to ApiMeta later is silently dropped from every merged response until
247+
# someone remembers to update it. That is how images, image_tokens,
248+
# tokens and cached_tokens went missing. Drive the assertion off the
249+
# model itself so the next added field fails here instead of in the wild.
250+
billed_fields = get_fields(ApiMetaBilledUnits())
251+
token_fields = get_fields(ApiMetaTokens())
252+
meta = ApiMeta(
253+
billed_units=ApiMetaBilledUnits(**{field: 1 for field in billed_fields}),
254+
tokens=ApiMetaTokens(**{field: 1 for field in token_fields}),
255+
cached_tokens=1,
256+
)
257+
258+
merged = merge_meta_field([meta, meta])
259+
260+
if merged.billed_units is None or merged.tokens is None:
261+
raise Exception("this is just for mypy")
262+
263+
for field in billed_fields:
264+
self.assertEqual(getattr(merged.billed_units, field), 2, f"billed_units.{field} was dropped")
265+
for field in token_fields:
266+
self.assertEqual(getattr(merged.tokens, field), 2, f"tokens.{field} was dropped")
267+
self.assertEqual(merged.cached_tokens, 2)
268+
208269
def test_sum_fields_if_not_none_with_none_entries(self) -> None:
209270
# billed_units list may contain None when ApiMeta.billed_units is unset;
210271
# sum_fields_if_not_none must skip None objects without raising AttributeError

0 commit comments

Comments
 (0)