mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
Fix tiered custom pricing for cache tokens
This commit is contained in:
parent
8bbe7edb71
commit
a47dae06ff
4 changed files with 151 additions and 1 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue