mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
address review: extract _inject_cache_tokens, fix _empty, cover ResponseAPIUsage paths, add tests
This commit is contained in:
parent
db20c23adc
commit
57f6de7953
2 changed files with 98 additions and 15 deletions
|
|
@ -5172,36 +5172,66 @@ class StandardLoggingPayloadSetup:
|
|||
Like get_usage_from_response_obj but returns a plain dict, skipping
|
||||
the Pydantic Usage construction on the hot path.
|
||||
"""
|
||||
_empty: dict = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
_empty: dict = {
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
}
|
||||
if combined_usage_object is not None:
|
||||
d = combined_usage_object.model_dump()
|
||||
d["cache_read_input_tokens"] = combined_usage_object._cache_read_input_tokens
|
||||
d["cache_creation_input_tokens"] = combined_usage_object._cache_creation_input_tokens
|
||||
return d
|
||||
return StandardLoggingPayloadSetup._inject_cache_tokens(
|
||||
combined_usage_object.model_dump(), combined_usage_object
|
||||
)
|
||||
if not response_obj:
|
||||
return _empty
|
||||
_raw = response_obj.get("usage", None)
|
||||
if _raw is None:
|
||||
return _empty
|
||||
if isinstance(_raw, ResponseAPIUsage):
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_raw
|
||||
).model_dump()
|
||||
)
|
||||
return StandardLoggingPayloadSetup._inject_cache_tokens(
|
||||
usage.model_dump(), usage
|
||||
)
|
||||
if isinstance(_raw, dict):
|
||||
if ResponseAPILoggingUtils._is_response_api_usage(_raw):
|
||||
return (
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_raw
|
||||
).model_dump()
|
||||
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_raw
|
||||
)
|
||||
return StandardLoggingPayloadSetup._inject_cache_tokens(
|
||||
usage.model_dump(), usage
|
||||
)
|
||||
return _raw
|
||||
if isinstance(_raw, Usage):
|
||||
d = _raw.model_dump()
|
||||
d["cache_read_input_tokens"] = _raw._cache_read_input_tokens
|
||||
d["cache_creation_input_tokens"] = _raw._cache_creation_input_tokens
|
||||
return d
|
||||
return StandardLoggingPayloadSetup._inject_cache_tokens(
|
||||
_raw.model_dump(), _raw
|
||||
)
|
||||
return _empty
|
||||
|
||||
@staticmethod
|
||||
def _inject_cache_tokens(d: dict, usage: Usage) -> dict:
|
||||
"""Inject cache token counts into a usage dict produced by model_dump().
|
||||
|
||||
Pydantic PrivateAttr fields are excluded from model_dump(), so
|
||||
_cache_read_input_tokens and _cache_creation_input_tokens must be
|
||||
read directly from the Usage instance.
|
||||
|
||||
cache_read_input_tokens sources (in priority order):
|
||||
1. usage._cache_read_input_tokens — set by Anthropic and Bedrock
|
||||
2. prompt_tokens_details.cached_tokens — set by OpenAI, Gemini, DeepSeek
|
||||
|
||||
cache_creation_input_tokens: Anthropic and Bedrock only via private attr.
|
||||
"""
|
||||
ptd = d.get("prompt_tokens_details") or {}
|
||||
ptd_cached = (ptd.get("cached_tokens") or 0) if isinstance(ptd, dict) else 0
|
||||
d["cache_read_input_tokens"] = (
|
||||
getattr(usage, "_cache_read_input_tokens", 0) or ptd_cached
|
||||
)
|
||||
d["cache_creation_input_tokens"] = getattr(usage, "_cache_creation_input_tokens", 0) or 0
|
||||
return d
|
||||
|
||||
@staticmethod
|
||||
def get_model_cost_information(
|
||||
base_model: Optional[str],
|
||||
|
|
|
|||
|
|
@ -1101,3 +1101,56 @@ def test_merge_litellm_metadata_bedrock_passthrough_scenario():
|
|||
|
||||
# Verify total number of fields (9 user fields + 4 model fields = 13)
|
||||
assert len(result) == 13
|
||||
|
||||
|
||||
# --- cache token tests ---
|
||||
|
||||
def _make_usage(cache_read: int = 0, cache_creation: int = 0, ptd_cached: int = 0) -> Usage:
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper
|
||||
u = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
u._cache_read_input_tokens = cache_read
|
||||
u._cache_creation_input_tokens = cache_creation
|
||||
if ptd_cached:
|
||||
u.prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=ptd_cached)
|
||||
return u
|
||||
|
||||
|
||||
def test_get_usage_as_dict_cache_tokens_from_private_attrs():
|
||||
usage = _make_usage(cache_read=1000, cache_creation=500)
|
||||
d = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
response_obj=None, combined_usage_object=usage
|
||||
)
|
||||
assert d["cache_read_input_tokens"] == 1000
|
||||
assert d["cache_creation_input_tokens"] == 500
|
||||
|
||||
|
||||
def test_get_usage_as_dict_cache_read_falls_back_to_prompt_tokens_details():
|
||||
# OpenAI / Gemini / DeepSeek set ptd.cached_tokens but not the private attr
|
||||
usage = _make_usage(cache_read=0, cache_creation=0, ptd_cached=800)
|
||||
d = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
response_obj=None, combined_usage_object=usage
|
||||
)
|
||||
assert d["cache_read_input_tokens"] == 800
|
||||
assert d["cache_creation_input_tokens"] == 0
|
||||
|
||||
|
||||
def test_get_usage_as_dict_private_attr_wins_over_ptd():
|
||||
# When both are set, private attr takes priority
|
||||
usage = _make_usage(cache_read=600, cache_creation=200, ptd_cached=800)
|
||||
d = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
response_obj=None, combined_usage_object=usage
|
||||
)
|
||||
assert d["cache_read_input_tokens"] == 600
|
||||
|
||||
|
||||
def test_get_usage_as_dict_empty_when_no_usage():
|
||||
d = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj=None)
|
||||
assert d["cache_read_input_tokens"] == 0
|
||||
assert d["cache_creation_input_tokens"] == 0
|
||||
|
||||
|
||||
def test_get_usage_as_dict_from_response_obj():
|
||||
usage = _make_usage(cache_read=300, cache_creation=100)
|
||||
d = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj={"usage": usage})
|
||||
assert d["cache_read_input_tokens"] == 300
|
||||
assert d["cache_creation_input_tokens"] == 100
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue