From e04c9a12b978ef5182208aa0730af000cb6ee4ec Mon Sep 17 00:00:00 2001 From: liming Date: Thu, 27 Aug 2026 19:32:50 +0800 Subject: [PATCH] fix(cost): keep custom token-rate helpers inside the basedpyright budget Narrow rate conversion and stop importing private strip helpers across modules so the lint job no longer exceeds the per-rule basedpyright ceiling. Co-authored-by: Cursor --- litellm/cost_calculator.py | 25 +++++++++++-------- litellm/proxy/litellm_pre_call_utils.py | 9 +++++++ .../pass_through_endpoints.py | 6 ++--- 3 files changed, 25 insertions(+), 15 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 3a4483bd6f4..1f596893ef3 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -258,6 +258,14 @@ def _model_info_from_params(params: Mapping[str, object], metadata_key: str) -> return _litellm_params_as_mapping(metadata.get("model_info")) +def _as_token_rate(value: object) -> float | None: + if isinstance(value, bool) or value is None: + return None + if isinstance(value, (int, float)): + return float(value) + return None + + def _custom_rates_from_mapping(source: Mapping[str, object] | None) -> Mapping[str, float] | None: if source is None: return None @@ -273,7 +281,7 @@ def _custom_rates_from_mapping(source: Mapping[str, object] | None) -> Mapping[s ("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}) + return MappingProxyType({key: rate for key, value in pairs if (rate := _as_token_rate(value)) is not None}) def extract_custom_cost_per_token( @@ -305,18 +313,16 @@ def _published_model_info( if not model: return None try: - return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + info: Final = 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 + return MappingProxyType({str(key): value for key, value in info.items()}) 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) + return _as_token_rate(info.get(field)) def _published_token_rate( @@ -345,10 +351,7 @@ def _cost_map_rate(key: str | None, field: str) -> float | 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 None - return float(value) + return _as_token_rate(raw.get(field)) def _declared_token_rate( @@ -390,7 +393,7 @@ def _first_declared_token_rate( field: str, ) -> float | None: for candidate in _unique_model_names(*models): - rate: Final = _declared_token_rate(candidate, custom_llm_provider, field) + rate = _declared_token_rate(candidate, custom_llm_provider, field) if rate is not None: return rate return None diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 064b53e07b7..10db60185bc 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -534,6 +534,15 @@ def _strip_client_pricing_overrides(data: dict[str, Any]) -> None: ) +def strip_unauthorized_client_pricing( + data: dict[str, Any], # mutable-ok: in-place strip of the caller request body + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """Drop client pricing overrides unless the key or team allows them.""" + if not _key_or_team_allows_client_pricing_override(user_api_key_dict): + _strip_client_pricing_overrides(data) + + def _get_metadata_variable_name(request: Request) -> str: """ Helper to return what the "metadata" field should be called in the request data diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e5939a81816..0133fb17b84 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -79,8 +79,7 @@ from litellm.proxy.common_utils.sse_keepalive import ( ) from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, - _key_or_team_allows_client_pricing_override, - _strip_client_pricing_overrides, + strip_unauthorized_client_pricing, ) from litellm.proxy.utils import normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository @@ -553,8 +552,7 @@ 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) + strip_unauthorized_client_pricing(_parsed_body, user_api_key_dict) litellm_params_in_body: Final = {} for k in all_litellm_params: