mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(bedrock): own the Converse batch usage shape in the provider layer
Shape detection and block normalization sat in the generic batch layer, which let batch and live parsing of the same wire format drift apart. Both now live on AmazonConverseConfig as is_converse_usage_shape and usage_from_batch_output, so batch_utils asks the provider adapter rather than knowing Bedrock's field names. Adds direct coverage for the shape predicate, the completion of an incomplete block, cache-count inflation, and the streaming usage event that shares the public transform. Drops the narrative banner from the batch tests.
This commit is contained in:
parent
7dbf2d57c5
commit
5fe7793a14
4 changed files with 81 additions and 40 deletions
|
|
@ -424,48 +424,17 @@ def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def _is_converse_usage_shape(usage_object: dict) -> bool:
|
||||
"""Converse-family models report camelCase token counts, not Anthropic's snake_case."""
|
||||
return "inputTokens" in usage_object or "outputTokens" in usage_object
|
||||
|
||||
|
||||
def _get_converse_batch_usage(usage_object: dict) -> Usage:
|
||||
"""Read a Converse-shaped usage block with the same transform the live Converse path uses,
|
||||
so a batch and an equivalent non-batch call agree on tokens."""
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
from litellm.types.llms.bedrock import ConverseTokenUsageBlock
|
||||
|
||||
input_tokens: Final = int(usage_object.get("inputTokens") or 0)
|
||||
output_tokens: Final = int(usage_object.get("outputTokens") or 0)
|
||||
cache_read: Final = int(
|
||||
usage_object.get("cacheReadInputTokens") or usage_object.get("cacheReadInputTokenCount") or 0
|
||||
)
|
||||
cache_write: Final = int(
|
||||
usage_object.get("cacheWriteInputTokens") or usage_object.get("cacheWriteInputTokenCount") or 0
|
||||
)
|
||||
return AmazonConverseConfig().transform_usage(
|
||||
ConverseTokenUsageBlock(
|
||||
inputTokens=input_tokens,
|
||||
outputTokens=output_tokens,
|
||||
totalTokens=int(usage_object.get("totalTokens") or input_tokens + output_tokens),
|
||||
cacheReadInputTokenCount=cache_read,
|
||||
cacheReadInputTokens=cache_read,
|
||||
cacheWriteInputTokenCount=cache_write,
|
||||
cacheWriteInputTokens=cache_write,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_provider: str = "openai") -> Usage:
|
||||
"""
|
||||
Get the tokens of a batch job from the response body
|
||||
"""
|
||||
if custom_llm_provider in ("anthropic", "bedrock"):
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
usage_object: Final = response_body.get("usage", None) or {}
|
||||
if custom_llm_provider == "bedrock" and _is_converse_usage_shape(usage_object):
|
||||
return _get_converse_batch_usage(usage_object)
|
||||
if custom_llm_provider == "bedrock" and AmazonConverseConfig.is_converse_usage_shape(usage_object):
|
||||
return AmazonConverseConfig().usage_from_batch_output(usage_object)
|
||||
anthropic_usage: Final = AnthropicConfig().calculate_usage(
|
||||
usage_object=usage_object,
|
||||
reasoning_content=None,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import copy
|
|||
import json
|
||||
import time
|
||||
import types
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
|
|
@ -1770,6 +1771,42 @@ class AmazonConverseConfig(BaseConfig):
|
|||
thinking_blocks_list.append(_redacted_block)
|
||||
return thinking_blocks_list
|
||||
|
||||
@staticmethod
|
||||
def is_converse_usage_shape(usage_object: Mapping[str, object]) -> bool:
|
||||
"""Converse-family models report camelCase token counts, not Anthropic's snake_case."""
|
||||
return "inputTokens" in usage_object or "outputTokens" in usage_object
|
||||
|
||||
@staticmethod
|
||||
def _usage_count(usage_object: Mapping[str, object], *keys: str) -> int:
|
||||
for key in keys:
|
||||
value = usage_object.get(key)
|
||||
if isinstance(value, (int, float)) and not isinstance(value, bool):
|
||||
return int(value)
|
||||
return 0
|
||||
|
||||
def usage_from_batch_output(self, usage_object: Mapping[str, object]) -> Usage:
|
||||
"""Read a Converse-shaped usage block out of a batch output line.
|
||||
|
||||
Batch output omits fields the live API always sends, so the block is
|
||||
completed before going through the same transform, keeping a batch and an
|
||||
equivalent non-batch call in agreement on tokens.
|
||||
"""
|
||||
input_tokens: Final = self._usage_count(usage_object, "inputTokens")
|
||||
output_tokens: Final = self._usage_count(usage_object, "outputTokens")
|
||||
cache_read: Final = self._usage_count(usage_object, "cacheReadInputTokens", "cacheReadInputTokenCount")
|
||||
cache_write: Final = self._usage_count(usage_object, "cacheWriteInputTokens", "cacheWriteInputTokenCount")
|
||||
return self.transform_usage(
|
||||
ConverseTokenUsageBlock(
|
||||
inputTokens=input_tokens,
|
||||
outputTokens=output_tokens,
|
||||
totalTokens=self._usage_count(usage_object, "totalTokens") or input_tokens + output_tokens,
|
||||
cacheReadInputTokenCount=cache_read,
|
||||
cacheReadInputTokens=cache_read,
|
||||
cacheWriteInputTokenCount=cache_write,
|
||||
cacheWriteInputTokens=cache_write,
|
||||
)
|
||||
)
|
||||
|
||||
def transform_usage(
|
||||
self,
|
||||
usage: ConverseTokenUsageBlock,
|
||||
|
|
|
|||
|
|
@ -1303,12 +1303,7 @@ async def test_output_file_content_bedrock_reads_with_deployment_aws_credentials
|
|||
|
||||
|
||||
# =========================================================================== #
|
||||
# Bedrock batch usage is parsed by the shape of the payload, not the provider
|
||||
#
|
||||
# Regression: every bedrock batch line went through the Anthropic usage parser,
|
||||
# which reads snake_case input_tokens/output_tokens. A Converse-family model
|
||||
# (Nova and friends) reports camelCase inputTokens/outputTokens, so usage read
|
||||
# 0/0/0 and the batch billed $0 despite real token consumption.
|
||||
# _get_batch_job_usage_from_response_body: bedrock usage shapes
|
||||
# =========================================================================== #
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -6003,3 +6003,43 @@ def test_adaptive_thinking_dropped_when_max_tokens_too_small_converse():
|
|||
)
|
||||
|
||||
assert "thinking" not in optional_params
|
||||
|
||||
|
||||
def test_is_converse_usage_shape_distinguishes_camel_case_from_anthropic():
|
||||
config = AmazonConverseConfig()
|
||||
assert config.is_converse_usage_shape({"inputTokens": 1, "outputTokens": 2}) is True
|
||||
assert config.is_converse_usage_shape({"outputTokens": 2}) is True
|
||||
assert config.is_converse_usage_shape({"input_tokens": 1, "output_tokens": 2}) is False
|
||||
assert config.is_converse_usage_shape({}) is False
|
||||
|
||||
|
||||
def test_usage_from_batch_output_completes_an_incomplete_block():
|
||||
"""Batch output omits totalTokens and the cache counts the live API always sends."""
|
||||
usage = AmazonConverseConfig().usage_from_batch_output({"inputTokens": 2202, "outputTokens": 540})
|
||||
assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (2202, 540, 2742)
|
||||
|
||||
|
||||
def test_usage_from_batch_output_inflates_input_by_cache_counts():
|
||||
usage = AmazonConverseConfig().usage_from_batch_output(
|
||||
{
|
||||
"inputTokens": 100,
|
||||
"outputTokens": 20,
|
||||
"totalTokens": 120,
|
||||
"cacheReadInputTokens": 800,
|
||||
"cacheWriteInputTokens": 200,
|
||||
}
|
||||
)
|
||||
assert usage.prompt_tokens == 1100
|
||||
assert usage.prompt_tokens_details.cached_tokens == 800
|
||||
assert usage.prompt_tokens_details.cache_creation_tokens == 200
|
||||
|
||||
|
||||
def test_streaming_usage_chunk_is_transformed():
|
||||
"""The streaming decoder's usage event feeds the same public transform."""
|
||||
from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
|
||||
|
||||
decoder = AWSEventStreamDecoder(model="us.amazon.nova-lite-v1:0")
|
||||
chunk = decoder.converse_chunk_parser({"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}})
|
||||
assert chunk.usage.prompt_tokens == 11
|
||||
assert chunk.usage.completion_tokens == 4
|
||||
assert chunk.usage.total_tokens == 15
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue