Fix tiered custom pricing for cache tokens

This commit is contained in:
its-amann 2026-06-11 19:40:14 +05:30 • committed by Aman Agnihotri
parent 8bbe7edb71
commit a47dae06ff
4 changed files with 151 additions and 1 deletions

View file

@ -5884,12 +5884,14 @@ class StandardLoggingPayloadSetup:
custom_llm_provider: str | None,
init_response_obj: Any | BaseModel | dict,
api_base: str | None = None,
router_model_id: str | None = None,
) -> StandardLoggingModelInformation:
model_cost_name: Final = _select_model_name_for_cost_calc(
model=base_model if custom_pricing else None,
completion_response=init_response_obj,
base_model=base_model,
custom_pricing=custom_pricing,
router_model_id=router_model_id,
)
if model_cost_name is None:
model_cost_information = StandardLoggingModelInformation(model_map_key="", model_map_value=None)
@ -6557,6 +6559,7 @@ def get_standard_logging_object_payload(
custom_llm_provider=kwargs.get("custom_llm_provider"),
init_response_obj=init_response_obj,
api_base=litellm_params.get("api_base"),
router_model_id=_model_id or None,
)
error_information, error_str = StandardLoggingPayloadSetup.get_error_information_for_logging_payload(

View file

@ -609,6 +609,13 @@ def _apply_off_peak_to_base_costs(
)
def _get_effective_prompt_tokens_for_tiered_pricing(usage: Usage) -> float:
prompt_tokens = float(getattr(usage, "prompt_tokens", 0) or 0)
cache_read_tokens = float(getattr(usage, "cache_read_input_tokens", 0) or 0)
cache_creation_tokens = float(getattr(usage, "cache_creation_input_tokens", 0) or 0)
return prompt_tokens + cache_read_tokens + cache_creation_tokens
def _get_token_base_cost(
model_info: ModelInfo,
usage: Usage,
@ -671,6 +678,8 @@ def _get_token_base_cost(
if k.startswith("input_cost_per_token_above_") and not k.endswith(_NON_STANDARD_THRESHOLD_SUFFIXES)
]
effective_prompt_tokens = _get_effective_prompt_tokens_for_tiered_pricing(usage)
# Only sort the threshold keys (typically 1-2 keys instead of 66+)
threshold: float | None = None
for key in sorted(threshold_keys, key=_parse_above_token_threshold, reverse=True):
@ -680,7 +689,7 @@ def _get_token_base_cost(
# Handle both formats: _above_128k_tokens and _above_128_tokens
threshold_str = key.split("_above_")[1].split("_tokens")[0]
threshold = _parse_above_token_threshold(key)
if _prompt_exceeds_threshold(usage.prompt_tokens, threshold, threshold_is_inclusive):
if _prompt_exceeds_threshold(effective_prompt_tokens, threshold, threshold_is_inclusive):
# Prefer a service_tier-specific above-threshold key when available,
# e.g. input_cost_per_token_priority_above_200k_tokens for Gemini
# ON_DEMAND_PRIORITY. Falls back to the standard key automatically

View file

@ -398,6 +398,86 @@ def test_standard_logging_payload_uses_deployment_when_no_base_model():
assert payload["model_map_information"]["model_map_key"] != "invoke_test_claude"
def test_standard_logging_model_map_uses_router_model_info_for_custom_pricing():
from litellm.litellm_core_utils.litellm_logging import (
Logging,
get_standard_logging_object_payload,
)
model_id = "tiered-custom-pricing-logging-test"
model_info = {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 0.000005,
"output_cost_per_token": 0.000025,
"cache_creation_input_token_cost": 0.00000625,
"cache_read_input_token_cost": 0.0000005,
"input_cost_per_token_above_200k_tokens": 0.00001,
"output_cost_per_token_above_200k_tokens": 0.0000375,
"cache_creation_input_token_cost_above_200k_tokens": 0.0000125,
"cache_read_input_token_cost_above_200k_tokens": 0.000001,
}
litellm.register_model(model_cost={model_id: model_info})
try:
logging_obj = Logging(
model="vertex_ai/claude-opus-4-6",
messages=[{"role": "user", "content": "Hello"}],
stream=False,
call_type="completion",
start_time=datetime.now(),
litellm_call_id="test-tiered-pricing",
function_id="test-function",
)
payload = get_standard_logging_object_payload(
kwargs={
"model": "vertex_ai/claude-opus-4-6",
"messages": [{"role": "user", "content": "Hello"}],
"response_cost": 1.0,
"custom_llm_provider": "vertex_ai",
"litellm_params": {
"api_base": "",
"metadata": {
"model_info": {
"id": model_id,
**model_info,
},
},
},
},
init_response_obj={
"id": "chatcmpl-test",
"model": "vertex_ai/claude-opus-4-6",
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 585659,
"cache_creation_input_tokens": 585659,
},
"choices": [],
},
start_time=datetime.now(),
end_time=datetime.now(),
logging_obj=logging_obj,
status="success",
)
assert payload is not None
model_map_value = payload["model_map_information"]["model_map_value"]
assert model_map_value is not None
assert (
model_map_value["cache_creation_input_token_cost_above_200k_tokens"]
== model_info["cache_creation_input_token_cost_above_200k_tokens"]
)
assert (
model_map_value["input_cost_per_token_above_200k_tokens"]
== model_info["input_cost_per_token_above_200k_tokens"]
)
finally:
litellm.model_cost.pop(model_id, None)
litellm.model_cost.pop(f"vertex_ai/{model_id}", None)
def test_get_hidden_params():
"""Test get_hidden_params with different inputs"""
# Test with None

View file

@ -2536,6 +2536,64 @@ def test_custom_pricing_applies_cache_creation_input_cost_via_cache_write_tokens
# ---------------------------------------------------------------------------
def test_custom_pricing_above_200k_cache_creation_tokens_uses_tiered_rate():
"""
Custom deployment pricing should apply above-200k cache write rates when
Anthropic-style usage reports cache creation tokens outside prompt_tokens.
"""
model_id = "tiered-custom-pricing-deploy-test"
model_info = {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 0.000005,
"output_cost_per_token": 0.000025,
"cache_creation_input_token_cost": 0.00000625,
"cache_read_input_token_cost": 0.0000005,
"input_cost_per_token_above_200k_tokens": 0.00001,
"output_cost_per_token_above_200k_tokens": 0.0000375,
"cache_creation_input_token_cost_above_200k_tokens": 0.0000125,
"cache_read_input_token_cost_above_200k_tokens": 0.000001,
}
cache_creation_tokens = 585659
litellm.register_model(model_cost={model_id: model_info})
try:
response = ModelResponse(
id="test-id",
created=1234567890,
model="vertex_ai/claude-opus-4-6",
object="chat.completion",
choices=[],
usage=Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=cache_creation_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=0,
),
)
cost = completion_cost(
completion_response=response,
model="vertex_ai/claude-opus-4-6",
custom_llm_provider="vertex_ai",
custom_pricing=True,
router_model_id=model_id,
)
assert cost == pytest.approx(
cache_creation_tokens
* model_info["cache_creation_input_token_cost_above_200k_tokens"]
)
assert cost != pytest.approx(
cache_creation_tokens * model_info["cache_creation_input_token_cost"]
)
finally:
litellm.model_cost.pop(model_id, None)
litellm.model_cost.pop(f"vertex_ai/{model_id}", None)
def test_extract_cache_read_tokens_anthropic_top_level():
from litellm.proxy.spend_tracking.savings import extract_cache_read_tokens as _extract_cache_read_tokens