fix(cost): keep one-sided custom rates and strip passthrough client pricing

Review feedback: allow a single configured token rate to keep the published
other side, avoid rebinding custom_cost_per_token, and drop untrusted
passthrough body prices before they can zero out spend.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
liming 2026-08-27 16:35:45 +08:00
parent 06a8444bdd
commit ff11623bb9
4 changed files with 468 additions and 51 deletions

View file

@ -2,8 +2,9 @@
## File for 'response_cost' calculation in Logging
import logging
import time
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from httpx import Response
@ -236,67 +237,127 @@ def _cost_per_token_custom_pricing_helper(
return None
def _litellm_params_as_mapping(litellm_params: object | None) -> Mapping[str, object] | None:
if litellm_params is None:
return None
if isinstance(litellm_params, Mapping):
return litellm_params
dump: Final = getattr(litellm_params, "model_dump", None)
if not callable(dump):
return None
dumped: Final = dump()
if not isinstance(dumped, Mapping):
return None
return dumped
def _model_info_from_params(params: Mapping[str, object], metadata_key: str) -> Mapping[str, object] | None:
metadata: Final = params.get(metadata_key)
if not isinstance(metadata, Mapping):
return None
return _litellm_params_as_mapping(metadata.get("model_info"))
def _custom_rates_from_mapping(source: Mapping[str, object] | None) -> Mapping[str, float] | None:
if source is None:
return None
input_cost: Final = source.get("input_cost_per_token")
output_cost: Final = source.get("output_cost_per_token")
if input_cost is None and output_cost is None:
return None
cache_read: Final = source.get("cache_read_input_token_cost")
cache_creation: Final = source.get("cache_creation_input_token_cost")
pairs: Final = (
("input_cost_per_token", input_cost),
("output_cost_per_token", output_cost),
("cache_read_input_token_cost", cache_read),
("cache_creation_input_token_cost", cache_creation),
)
return MappingProxyType({key: float(value) for key, value in pairs if value is not None})
def extract_custom_cost_per_token(
litellm_params: object | None,
) -> CostPerToken | None:
"""Return deployment token rates from litellm_params when both input and output are set.
) -> Mapping[str, float] | None:
"""Return deployment token rates from litellm_params when input and/or output is set.
Rates may sit on litellm_params itself (UI / model_list) or under
metadata.model_info / litellm_metadata.model_info (/v1/messages, /v1/responses).
One-sided rates are returned as-is; callers that need a complete CostPerToken
fill the missing side from the published price map.
Optional cache rates are copied when present so the custom-pricing helper can
apply them instead of falling back to the input rate.
"""
if litellm_params is None:
params: Final = _litellm_params_as_mapping(litellm_params)
if params is None:
return None
if not isinstance(litellm_params, dict):
dump = getattr(litellm_params, "model_dump", None)
if not callable(dump):
return None
dumped = dump()
if not isinstance(dumped, dict):
return None
litellm_params = dumped
return (
_custom_rates_from_mapping(params)
or _custom_rates_from_mapping(_model_info_from_params(params, "metadata"))
or _custom_rates_from_mapping(_model_info_from_params(params, "litellm_metadata"))
)
def _from_mapping(source: object) -> CostPerToken | None:
if not isinstance(source, dict):
return None
input_cost = source.get("input_cost_per_token")
output_cost = source.get("output_cost_per_token")
if input_cost is None or output_cost is None:
return None
result: CostPerToken = {
"input_cost_per_token": float(input_cost),
"output_cost_per_token": float(output_cost),
}
cache_read = source.get("cache_read_input_token_cost")
if cache_read is not None:
result["cache_read_input_token_cost"] = float(cache_read)
cache_creation = source.get("cache_creation_input_token_cost")
if cache_creation is not None:
result["cache_creation_input_token_cost"] = float(cache_creation)
return result
from_top = _from_mapping(litellm_params)
if from_top is not None:
return from_top
for metadata_key in ("metadata", "litellm_metadata"):
metadata = litellm_params.get(metadata_key) or {}
from_info = _from_mapping(metadata.get("model_info") if isinstance(metadata, dict) else None)
if from_info is not None:
return from_info
return None
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)
if value is None:
return 0.0
return float(value)
def _complete_custom_cost_per_token(
rates: Mapping[str, float] | None,
*,
model: str | None,
custom_llm_provider: str | None,
) -> CostPerToken | None:
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
resolved_input: Final = (
float(input_cost)
if input_cost is not None
else _published_token_rate(model, custom_llm_provider, "input_cost_per_token")
)
resolved_output: Final = (
float(output_cost)
if output_cost is not None
else _published_token_rate(model, custom_llm_provider, "output_cost_per_token")
)
cache_read: Final = rates.get("cache_read_input_token_cost")
cache_creation: Final = rates.get("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),
}
return completed
def _custom_cost_per_token_from_logging_obj(
litellm_logging_obj: LitellmLoggingObject | None,
) -> CostPerToken | None:
) -> Mapping[str, float] | None:
if litellm_logging_obj is None:
return None
extracted = extract_custom_cost_per_token(getattr(litellm_logging_obj, "litellm_params", None))
if extracted is not None:
return extracted
details = getattr(litellm_logging_obj, "model_call_details", None) or {}
nested = details.get("litellm_params") if isinstance(details, dict) else None
from_attr: Final = extract_custom_cost_per_token(getattr(litellm_logging_obj, "litellm_params", None))
if from_attr is not None:
return from_attr
details: Final = getattr(litellm_logging_obj, "model_call_details", None)
nested: Final = details.get("litellm_params") if isinstance(details, Mapping) else None
return extract_custom_cost_per_token(nested)
@ -1273,9 +1334,6 @@ def completion_cost(
- For un-mapped Replicate models, the cost is calculated based on the total time used for the request.
"""
try:
if custom_cost_per_token is None:
custom_cost_per_token = _custom_cost_per_token_from_logging_obj(litellm_logging_obj)
call_type = _infer_call_type(call_type, completion_response) or "completion"
if (
@ -1331,6 +1389,16 @@ def completion_cost(
if model is not None:
potential_model_names.append(model)
resolved_custom_cost_per_token: Final = (
custom_cost_per_token
if custom_cost_per_token is not None
else _complete_custom_cost_per_token(
_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,
)
)
for idx, model in enumerate(potential_model_names):
try:
if verbose_logger.isEnabledFor(logging.DEBUG):
@ -1680,7 +1748,7 @@ def completion_cost(
response_time_ms=total_time,
region_name=region_name,
custom_cost_per_second=custom_cost_per_second,
custom_cost_per_token=custom_cost_per_token,
custom_cost_per_token=resolved_custom_cost_per_token,
prompt_characters=prompt_characters,
completion_characters=completion_characters,
cache_creation_input_tokens=cache_creation_input_tokens,

View file

@ -77,7 +77,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_key_or_team_allows_client_pricing_override,
_strip_client_pricing_overrides,
)
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
@ -549,6 +553,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
from litellm.types.utils import all_litellm_params
_parsed_body = _parsed_body or {}
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
_strip_client_pricing_overrides(_parsed_body)
litellm_params_in_body: Final = {}
for k in all_litellm_params:

View file

@ -11,6 +11,7 @@ import httpx
import pytest
import litellm
from typing import AsyncGenerator
from litellm.cost_calculator import extract_custom_cost_per_token
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
from litellm.proxy.pass_through_endpoints.success_handler import (
@ -235,6 +236,113 @@ def test_init_kwargs_with_litellm_metadata(mock_request, mock_user_api_key_dict)
assert metadata["user_api_key"] == "test-key"
def _passthrough_logging_obj():
return LiteLLMLoggingObj(
model="test-model",
messages=[],
stream=False,
call_type="test-call-type",
start_time=datetime.now(),
litellm_call_id="test-call-id",
function_id="test-function-id",
)
def test_init_kwargs_strips_client_token_rates(mock_request, mock_user_api_key_dict):
"""Client-supplied 0 rates must not land in litellm_params (budget bypass)."""
request = mock_request()
parsed_body = {
"model": "claude-sonnet-4-5-20250929",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"messages": [{"role": "user", "content": "hi"}],
}
passthrough_payload = PassthroughStandardLoggingPayload(
url="https://test.com",
request_body={},
)
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=mock_user_api_key_dict,
passthrough_logging_payload=passthrough_payload,
_parsed_body=parsed_body,
litellm_call_id="test-call-id",
logging_obj=_passthrough_logging_obj(),
)
assert "input_cost_per_token" not in result["litellm_params"]
assert "output_cost_per_token" not in result["litellm_params"]
assert extract_custom_cost_per_token(result["litellm_params"]) is None
def test_init_kwargs_strips_client_model_info_pricing(
mock_request, mock_user_api_key_dict
):
request = mock_request()
parsed_body = {
"litellm_metadata": {
"tags": ["keep-me"],
"model_info": {
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
},
}
}
passthrough_payload = PassthroughStandardLoggingPayload(
url="https://test.com",
request_body={},
)
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=mock_user_api_key_dict,
passthrough_logging_payload=passthrough_payload,
_parsed_body=parsed_body,
litellm_call_id="test-call-id",
logging_obj=_passthrough_logging_obj(),
)
metadata = result["litellm_params"]["metadata"]
assert metadata["tags"] == ["keep-me"]
assert "model_info" not in metadata
def test_init_kwargs_keeps_client_pricing_when_key_allows_override(mock_request):
request = mock_request()
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="test-team",
end_user_id="test-user",
metadata={"allow_client_pricing_override": True},
)
parsed_body = {
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
}
passthrough_payload = PassthroughStandardLoggingPayload(
url="https://test.com",
request_body={},
)
result = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request,
user_api_key_dict=user_api_key_dict,
passthrough_logging_payload=passthrough_payload,
_parsed_body=parsed_body,
litellm_call_id="test-call-id",
logging_obj=_passthrough_logging_obj(),
)
assert result["litellm_params"]["input_cost_per_token"] == 0.0
assert result["litellm_params"]["output_cost_per_token"] == 0.0
assert extract_custom_cost_per_token(result["litellm_params"]) == {
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
}
def test_init_kwargs_with_tags_in_header(mock_request, mock_user_api_key_dict):
"""
Tags should be added to metadata if they exist in headers

View file

@ -11,6 +11,9 @@ import litellm
from litellm.cost_calculator import (
BaseTokenUsageProcessor,
RealtimeAPITokenUsageProcessor,
_complete_custom_cost_per_token,
_custom_cost_per_token_from_logging_obj,
_published_token_rate,
completion_cost,
cost_per_token,
extract_custom_cost_per_token,
@ -20,6 +23,7 @@ from litellm.cost_calculator import (
from litellm.types.llms.openai import OpenAIRealtimeStreamList
from litellm.types.utils import (
CacheCreationTokenDetails,
CustomPricingLiteLLMParams,
ModelInfo,
ModelResponse,
PromptTokensDetailsWrapper,
@ -956,7 +960,13 @@ def test_custom_pricing_cost_calc_uses_router_model_id_from_litellm_metadata():
def test_extract_custom_cost_per_token_from_litellm_params_and_model_info():
assert extract_custom_cost_per_token(None) is None
assert extract_custom_cost_per_token({"input_cost_per_token": 1.2e-05}) is None
assert extract_custom_cost_per_token({"custom_llm_provider": "anthropic"}) is None
assert extract_custom_cost_per_token({"input_cost_per_token": 1.2e-05}) == {
"input_cost_per_token": 1.2e-05,
}
assert extract_custom_cost_per_token({"output_cost_per_token": 3.6e-05}) == {
"output_cost_per_token": 3.6e-05,
}
assert extract_custom_cost_per_token(
{
"input_cost_per_token": 1.2e-05,
@ -968,6 +978,20 @@ def test_extract_custom_cost_per_token_from_litellm_params_and_model_info():
"output_cost_per_token": 3.6e-05,
"cache_read_input_token_cost": 1.2e-06,
}
assert extract_custom_cost_per_token(
{
"metadata": {
"model_info": {
"id": "deploy-meta",
"input_cost_per_token": 0.0002,
"output_cost_per_token": 0.0008,
},
},
}
) == {
"input_cost_per_token": 0.0002,
"output_cost_per_token": 0.0008,
}
assert extract_custom_cost_per_token(
{
"litellm_metadata": {
@ -984,6 +1008,48 @@ def test_extract_custom_cost_per_token_from_litellm_params_and_model_info():
}
def test_extract_custom_cost_per_token_from_pydantic_params():
both_sides = CustomPricingLiteLLMParams(
input_cost_per_token=1.2e-05,
output_cost_per_token=3.6e-05,
)
assert extract_custom_cost_per_token(both_sides) == {
"input_cost_per_token": 1.2e-05,
"output_cost_per_token": 3.6e-05,
}
input_only = CustomPricingLiteLLMParams(input_cost_per_token=1.2e-05)
assert extract_custom_cost_per_token(input_only) == {
"input_cost_per_token": 1.2e-05,
}
def test_extract_custom_cost_per_token_rejects_non_mapping_sources():
assert extract_custom_cost_per_token("not-params") is None
assert extract_custom_cost_per_token([1, 2]) is None
class _UncallableDump:
model_dump = "not-callable"
assert extract_custom_cost_per_token(_UncallableDump()) is None
class _NonDictDump:
def model_dump(self):
return ["not", "a", "mapping"]
assert extract_custom_cost_per_token(_NonDictDump()) is None
assert extract_custom_cost_per_token({"metadata": "not-a-dict"}) is None
assert extract_custom_cost_per_token({"metadata": {"model_info": "x"}}) is None
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
def test_completion_cost_unknown_anthropic_model_uses_litellm_params_rates():
"""Unknown anthropic models logged $0 on /v1/messages even when the
deployment set input/output rates in litellm_params.
@ -1111,6 +1177,175 @@ def test_anthropic_passthrough_unknown_model_spend_uses_litellm_params_rates():
assert kwargs["response_cost"] > 0
@pytest.mark.parametrize(
"declared",
[
{"input_cost_per_token": 1e-06},
{"output_cost_per_token": 5e-06},
],
ids=["input-only", "output-only"],
)
def test_completion_cost_one_sided_custom_rate_keeps_published_other_side(
_local_model_cost_map, declared
):
"""A deployment may configure only one direction.
The missing side must keep the published price-map rate, not 0.
"""
import time
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
model = "gpt-4o-mini"
published = litellm.get_model_info(model=model)
prompt_tokens = 100
completion_tokens = 20
input_cost = declared.get(
"input_cost_per_token", published["input_cost_per_token"]
)
output_cost = declared.get(
"output_cost_per_token", published["output_cost_per_token"]
)
logging_obj = LiteLLMLoggingObj(
model=model,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="completion",
start_time=time.time(),
litellm_call_id="test-one-sided-custom-pricing",
function_id="test-fn",
)
logging_obj.update_environment_variables(
model=model,
user="",
optional_params={},
litellm_params={"custom_llm_provider": "openai", **declared},
)
logging_obj.model_call_details["custom_llm_provider"] = "openai"
response = ModelResponse(
id="test-id",
model=model,
choices=[],
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
cost = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
litellm_logging_obj=logging_obj,
)
expected = prompt_tokens * input_cost + completion_tokens * output_cost
assert cost == pytest.approx(expected)
def test_completion_cost_one_sided_unknown_model_uses_zero_for_missing_side():
"""Unmapped models have no published other-side rate, so that side is 0."""
import time
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
unknown_model = "litellm-unmapped-custom-priced-qwen-onesided"
input_cost = 1.2e-05
prompt_tokens = 100
completion_tokens = 20
logging_obj = LiteLLMLoggingObj(
model=unknown_model,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-unmapped-one-sided",
function_id="test-fn",
)
logging_obj.update_environment_variables(
model=unknown_model,
user="",
optional_params={},
litellm_params={
"custom_llm_provider": "anthropic",
"input_cost_per_token": input_cost,
},
)
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
response = ModelResponse(
id="test-id",
model=unknown_model,
choices=[],
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
cost = completion_cost(
completion_response=response,
model=unknown_model,
custom_llm_provider="anthropic",
call_type="anthropic_messages",
custom_pricing=True,
litellm_logging_obj=logging_obj,
)
assert cost == pytest.approx(prompt_tokens * input_cost)
def test_completion_cost_reads_nested_litellm_params_from_model_call_details():
import time
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
unknown_model = "litellm-unmapped-nested-litellm-params"
input_cost = 1.2e-05
output_cost = 3.6e-05
prompt_tokens = 100
completion_tokens = 20
logging_obj = LiteLLMLoggingObj(
model=unknown_model,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-nested-litellm-params",
function_id="test-fn",
)
logging_obj.litellm_params = None
logging_obj.model_call_details["litellm_params"] = {
"custom_llm_provider": "anthropic",
"input_cost_per_token": input_cost,
"output_cost_per_token": output_cost,
}
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
response = ModelResponse(
id="test-id",
model=unknown_model,
choices=[],
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
cost = completion_cost(
completion_response=response,
model=unknown_model,
custom_llm_provider="anthropic",
call_type="anthropic_messages",
custom_pricing=True,
litellm_logging_obj=logging_obj,
)
expected = prompt_tokens * input_cost + completion_tokens * output_cost
assert cost == pytest.approx(expected)
def test_per_request_custom_pricing_with_router():
"""When custom pricing is passed as per-request kwargs (not in model_list),
_select_model_name_for_cost_calc should fall back to the model name