fix(cost): inherit published cache rates from the backend model

Custom router ids often only store input/output, so completing one-sided
pricing must not bill Anthropic cache tokens at the normal input rate.
Unauthorized passthrough client rates stay stripped from litellm_params.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
liming 2026-08-27 18:08:06 +08:00
parent ff11623bb9
commit 2a08259f2a
3 changed files with 369 additions and 20 deletions

View file

@ -298,52 +298,159 @@ def extract_custom_cost_per_token(
)
def _published_model_info(
model: str | None,
custom_llm_provider: str | None,
) -> Mapping[str, object] | None:
if not model:
return None
try:
return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # get_model_info raises Exception for unmapped models
return None
def _rate_from_model_info(info: Mapping[str, object] | None, field: str) -> float | None:
if info is None:
return None
value: Final = info.get(field)
if value is None:
return None
return float(value)
def _published_token_rate(
model: str | None,
custom_llm_provider: str | None,
field: str,
) -> float:
if not model:
return 0.0
try:
info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return 0.0
value: Final = info.get(field)
) -> float | None:
return _rate_from_model_info(_published_model_info(model, custom_llm_provider), field)
def _unique_model_names(*names: str | None) -> tuple[str, ...]:
unique: list[str] = []
seen: set[str] = set()
for name in names:
if not isinstance(name, str) or not name or name in seen:
continue
seen.add(name)
unique.append(name)
if "/" in name:
tail: Final = name.split("/", 1)[1]
if tail and tail not in seen:
seen.add(tail)
unique.append(tail)
return tuple(unique)
def _cost_map_rate(key: str | None, field: str) -> float | None:
if not key:
return None
raw: Final = litellm.model_cost.get(key)
if not isinstance(raw, Mapping):
return None
value: Final = raw.get(field)
if value is None:
return 0.0
return None
return float(value)
def _declared_token_rate(
model: str | None,
custom_llm_provider: str | None,
field: str,
) -> float | None:
"""Return a price-map rate that was actually declared on the entry.
``get_model_info`` synthesizes ``input_cost_per_token`` / ``output_cost_per_token``
to 0 when they are missing. A custom ``router_model_id`` entry typically has
only those two fields; treating the zeros or missing cache keys as published
would skip the backend model that does have cache-specific rates.
"""
if not model:
return None
from_map: Final = _cost_map_rate(model, field)
if from_map is not None:
return from_map
if custom_llm_provider:
from_prefixed: Final = _cost_map_rate(f"{custom_llm_provider}/{model}", field)
if from_prefixed is not None:
return from_prefixed
info: Final = _published_model_info(model, custom_llm_provider)
if info is None:
return None
info_key: Final = info.get("key")
from_resolved: Final = _cost_map_rate(info_key if isinstance(info_key, str) else None, field)
if from_resolved is not None:
return from_resolved
if field in ("input_cost_per_token", "output_cost_per_token"):
return None
return _rate_from_model_info(info, field)
def _first_declared_token_rate(
models: Sequence[str | None],
custom_llm_provider: str | None,
field: str,
) -> float | None:
for candidate in _unique_model_names(*models):
rate: Final = _declared_token_rate(candidate, custom_llm_provider, field)
if rate is not None:
return rate
return None
def _complete_custom_cost_per_token(
rates: Mapping[str, float] | None,
*,
model: str | None,
custom_llm_provider: str | None,
fallback_models: Sequence[str | None] = (),
) -> CostPerToken | None:
"""Fill missing sides of a partial custom CostPerToken from declared price-map rates.
``model`` is often a custom ``router_model_id`` that only stores input/output.
``fallback_models`` should include the backend model so cache-specific rates
come from that published entry instead of the normal input rate.
"""
if rates is None:
return None
input_cost: Final = rates.get("input_cost_per_token")
output_cost: Final = rates.get("output_cost_per_token")
if input_cost is None and output_cost is None:
return None
lookup_models: Final = (model, *fallback_models)
resolved_input: Final = (
float(input_cost)
if input_cost is not None
else _published_token_rate(model, custom_llm_provider, "input_cost_per_token")
else (_first_declared_token_rate(lookup_models, custom_llm_provider, "input_cost_per_token") or 0.0)
)
resolved_output: Final = (
float(output_cost)
if output_cost is not None
else _published_token_rate(model, custom_llm_provider, "output_cost_per_token")
else (_first_declared_token_rate(lookup_models, custom_llm_provider, "output_cost_per_token") or 0.0)
)
cache_read: Final = rates.get("cache_read_input_token_cost")
cache_creation: Final = rates.get("cache_creation_input_token_cost")
published_cache_read: Final = _first_declared_token_rate(
lookup_models, custom_llm_provider, "cache_read_input_token_cost"
)
published_cache_creation: Final = _first_declared_token_rate(
lookup_models, custom_llm_provider, "cache_creation_input_token_cost"
)
completed: Final[CostPerToken] = {
"input_cost_per_token": resolved_input,
"output_cost_per_token": resolved_output,
"cache_read_input_token_cost": (float(cache_read) if cache_read is not None else resolved_input),
"cache_creation_input_token_cost": (float(cache_creation) if cache_creation is not None else resolved_input),
"cache_read_input_token_cost": (
float(cache_read)
if cache_read is not None
else (published_cache_read if published_cache_read is not None else resolved_input)
),
"cache_creation_input_token_cost": (
float(cache_creation)
if cache_creation is not None
else (published_cache_creation if published_cache_creation is not None else resolved_input)
),
}
return completed
@ -361,6 +468,29 @@ def _custom_cost_per_token_from_logging_obj(
return extract_custom_cost_per_token(nested)
def _backend_model_from_logging_obj(
litellm_logging_obj: LitellmLoggingObject | None,
) -> str | None:
if litellm_logging_obj is None:
return None
attr_params: Final = _litellm_params_as_mapping(getattr(litellm_logging_obj, "litellm_params", None))
if attr_params is not None:
attr_model: Final = attr_params.get("model")
if isinstance(attr_model, str) and attr_model:
return attr_model
details: Final = getattr(litellm_logging_obj, "model_call_details", None)
nested: Final = details.get("litellm_params") if isinstance(details, Mapping) else None
nested_params: Final = _litellm_params_as_mapping(nested)
if nested_params is not None:
nested_model: Final = nested_params.get("model")
if isinstance(nested_model, str) and nested_model:
return nested_model
logging_model: Final = getattr(litellm_logging_obj, "model", None)
if isinstance(logging_model, str) and logging_model:
return logging_model
return None
def _get_additional_costs(
model: str,
custom_llm_provider: str | None,
@ -1396,6 +1526,12 @@ def completion_cost(
_custom_cost_per_token_from_logging_obj(litellm_logging_obj),
model=selected_model,
custom_llm_provider=custom_llm_provider if isinstance(custom_llm_provider, str) else None,
fallback_models=(
model if isinstance(model, str) else None,
_get_response_model(completion_response),
base_model,
_backend_model_from_logging_obj(litellm_logging_obj),
),
)
)

View file

@ -682,11 +682,13 @@ def test_init_kwargs_filters_pricing_params(mock_request, mock_user_api_key_dict
assert parsed_body["temperature"] == 0.7
assert parsed_body["max_tokens"] == 100
# Verify pricing parameters are stored in litellm_params for internal use
# Unauthorized keys must not keep client rates in litellm_params; otherwise
# extract_custom_cost_per_token would bill from the request body (budget bypass).
# Authorized keys are covered by test_init_kwargs_keeps_client_pricing_when_key_allows_override.
litellm_params = result["litellm_params"]
assert litellm_params["input_cost_per_token"] == 0.00002
assert litellm_params["output_cost_per_token"] == 0.00002
# Note: Other pricing params are also stored but we test the key ones that caused the regression
assert "input_cost_per_token" not in litellm_params
assert "output_cost_per_token" not in litellm_params
assert extract_custom_cost_per_token(litellm_params) is None
def test_custom_pricing_used_in_cost_calculation():

View file

@ -1045,9 +1045,220 @@ def test_complete_custom_cost_per_token_defensive_branches(_local_model_cost_map
assert _complete_custom_cost_per_token(None, model="gpt-4o-mini", custom_llm_provider="openai") is None
assert _complete_custom_cost_per_token({}, model="gpt-4o-mini", custom_llm_provider="openai") is None
assert _custom_cost_per_token_from_logging_obj(None) is None
assert _published_token_rate(None, "openai", "input_cost_per_token") == 0.0
assert _published_token_rate("", "openai", "output_cost_per_token") == 0.0
assert _published_token_rate("gpt-4o-mini", "openai", "this_field_does_not_exist") == 0.0
assert _published_token_rate(None, "openai", "input_cost_per_token") is None
assert _published_token_rate("", "openai", "output_cost_per_token") is None
assert _published_token_rate("gpt-4o-mini", "openai", "this_field_does_not_exist") is None
assert _published_token_rate(
"litellm-unmapped-custom-priced-qwen", "anthropic", "input_cost_per_token"
) is None
def test_complete_output_only_keeps_published_cache_rates(_local_model_cost_map):
"""Output-only custom pricing must not bill cache at the normal input rate."""
model = "claude-sonnet-4-5-20250929"
published = litellm.get_model_info(model=model, custom_llm_provider="anthropic")
custom_output = 5e-06
assert published["cache_read_input_token_cost"] != published["input_cost_per_token"]
assert published["cache_creation_input_token_cost"] != published["input_cost_per_token"]
completed = _complete_custom_cost_per_token(
{"output_cost_per_token": custom_output},
model=model,
custom_llm_provider="anthropic",
)
assert completed is not None
assert completed["output_cost_per_token"] == custom_output
assert completed["input_cost_per_token"] == published["input_cost_per_token"]
assert completed["cache_read_input_token_cost"] == published["cache_read_input_token_cost"]
assert completed["cache_creation_input_token_cost"] == published["cache_creation_input_token_cost"]
def test_completion_cost_output_only_custom_rate_uses_published_cache_rates(
_local_model_cost_map,
):
import time
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
model = "claude-sonnet-4-5-20250929"
published = litellm.get_model_info(model=model, custom_llm_provider="anthropic")
custom_output = 5e-06
regular_prompt = 20
cache_read = 80
completion_tokens = 10
logging_obj = LiteLLMLoggingObj(
model=model,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-output-only-cache-rates",
function_id="test-fn",
)
logging_obj.update_environment_variables(
model=model,
user="",
optional_params={},
litellm_params={
"custom_llm_provider": "anthropic",
"output_cost_per_token": custom_output,
},
)
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
response = ModelResponse(
id="test-id",
model=model,
choices=[],
usage=Usage(
prompt_tokens=regular_prompt + cache_read,
completion_tokens=completion_tokens,
total_tokens=regular_prompt + cache_read + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cache_read),
),
)
cost = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="anthropic",
call_type="anthropic_messages",
litellm_logging_obj=logging_obj,
)
expected = (
regular_prompt * published["input_cost_per_token"]
+ cache_read * published["cache_read_input_token_cost"]
+ completion_tokens * custom_output
)
assert cost == pytest.approx(expected)
billed_cache_at_input = (
regular_prompt * published["input_cost_per_token"]
+ cache_read * published["input_cost_per_token"]
+ completion_tokens * custom_output
)
assert cost != pytest.approx(billed_cache_at_input)
def test_complete_output_only_router_id_uses_backend_cache_rates(_local_model_cost_map):
"""A custom router_model_id usually stores only input/output. Missing cache
rates must come from the backend Anthropic model, not the normal input rate.
"""
backend = "claude-sonnet-4-5-20250929"
router_id = "71ad2e1c-71db-4246-a558-d01480578941"
published = litellm.get_model_info(model=backend, custom_llm_provider="anthropic")
custom_output = 5e-06
litellm.register_model(
{
router_id: {
"input_cost_per_token": 1e-06,
"output_cost_per_token": custom_output,
"litellm_provider": "anthropic",
"mode": "chat",
}
},
persist_across_reloads=False,
)
assert litellm.model_cost[router_id].get("cache_read_input_token_cost") is None
assert published["cache_read_input_token_cost"] != published["input_cost_per_token"]
without_backend = _complete_custom_cost_per_token(
{"output_cost_per_token": custom_output},
model=f"anthropic/{router_id}",
custom_llm_provider="anthropic",
)
assert without_backend is not None
assert without_backend["cache_read_input_token_cost"] == without_backend["input_cost_per_token"]
completed = _complete_custom_cost_per_token(
{"output_cost_per_token": custom_output},
model=f"anthropic/{router_id}",
custom_llm_provider="anthropic",
fallback_models=(backend,),
)
assert completed is not None
assert completed["output_cost_per_token"] == custom_output
assert completed["input_cost_per_token"] == 1e-06
assert completed["cache_read_input_token_cost"] == published["cache_read_input_token_cost"]
assert completed["cache_creation_input_token_cost"] == published["cache_creation_input_token_cost"]
def test_completion_cost_router_id_uses_backend_cache_rates(_local_model_cost_map):
import time
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
backend = "claude-sonnet-4-5-20250929"
router_id = "test-router-custom-cache-uuid"
published = litellm.get_model_info(model=backend, custom_llm_provider="anthropic")
custom_input = 1e-06
custom_output = 5e-06
regular_prompt = 20
cache_read = 80
completion_tokens = 10
litellm.register_model(
{
router_id: {
"input_cost_per_token": custom_input,
"output_cost_per_token": custom_output,
"litellm_provider": "anthropic",
"mode": "chat",
}
},
persist_across_reloads=False,
)
logging_obj = LiteLLMLoggingObj(
model=backend,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-router-id-cache-rates",
function_id="test-fn",
)
logging_obj.update_environment_variables(
model=backend,
user="",
optional_params={},
litellm_params={
"model": backend,
"custom_llm_provider": "anthropic",
"input_cost_per_token": custom_input,
"output_cost_per_token": custom_output,
},
)
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
response = ModelResponse(
id="test-id",
model=backend,
choices=[],
usage=Usage(
prompt_tokens=regular_prompt + cache_read,
completion_tokens=completion_tokens,
total_tokens=regular_prompt + cache_read + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cache_read),
),
)
cost = completion_cost(
completion_response=response,
model=backend,
custom_llm_provider="anthropic",
call_type="anthropic_messages",
custom_pricing=True,
router_model_id=router_id,
litellm_logging_obj=logging_obj,
)
expected = (
regular_prompt * custom_input
+ cache_read * published["cache_read_input_token_cost"]
+ completion_tokens * custom_output
)
billed_cache_at_custom_input = (
regular_prompt * custom_input + cache_read * custom_input + completion_tokens * custom_output
)
assert cost == pytest.approx(expected)
assert cost != pytest.approx(billed_cache_at_custom_input)
def test_completion_cost_unknown_anthropic_model_uses_litellm_params_rates():