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