fix(cost): satisfy type-discipline gates and sync generated types for balanced tier

This commit is contained in:
shrey kharbanda 2026-09-24 05:24:14 +00:00
parent a8c8fdc64e
commit 27be70891d
5 changed files with 58 additions and 31 deletions

View file

@ -286,6 +286,9 @@ pub struct ModelInfo {
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
/// Balanced service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_balanced: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_batches: Option<f64>,
/// Flex service-tier rate for the same-named base field.
@ -380,6 +383,9 @@ pub struct ModelInfo {
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_512k_tokens: Option<f64>,
/// Balanced service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_balanced: Option<f64>,
/// USD per prompt token via the provider's batch API.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_batches: Option<f64>,
@ -501,6 +507,9 @@ pub struct ModelInfo {
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_512k_tokens: Option<f64>,
/// Balanced service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_balanced: Option<f64>,
/// USD per generated token via the provider's batch API.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_batches: Option<f64>,

View file

@ -970,7 +970,7 @@ def _completion_window_value(metadata: object) -> str | None:
pricing, so it returns None."""
if not isinstance(metadata, dict):
return None
window: Final = cast(dict[str, object], metadata).get("completion_window")
window: Final[object] = metadata.get("completion_window")
if isinstance(window, str) and window in (ServiceTier.FLEX.value, ServiceTier.BALANCED.value):
return window
return None
@ -980,9 +980,8 @@ def _service_tier_from_completion_window(optional_params: dict[str, object]) ->
"""Read ``metadata.completion_window`` from ``extra_body`` or a top-level ``metadata``
param (the two shapes callers use to pick a provider completion window directly)."""
extra_body: Final = optional_params.get("extra_body")
return _completion_window_value(
cast(dict[str, object], extra_body).get("metadata") if isinstance(extra_body, dict) else None
) or _completion_window_value(optional_params.get("metadata"))
extra_metadata: Final[object] = extra_body.get("metadata") if isinstance(extra_body, dict) else None
return _completion_window_value(extra_metadata) or _completion_window_value(optional_params.get("metadata"))
def get_usage_object(
@ -1398,7 +1397,7 @@ def completion_cost(
if service_tier is None and optional_params is not None:
service_tier = _normalize_service_tier(optional_params.get("service_tier"))
if service_tier is None:
service_tier = _service_tier_from_completion_window(cast(dict[str, object], optional_params))
service_tier = _service_tier_from_completion_window(optional_params)
service_tier = _normalize_service_tier(service_tier)

View file

@ -4,7 +4,7 @@ Dynamic configuration class generator for JSON-based providers.
from collections.abc import Coroutine, Mapping
from types import MappingProxyType
from typing import Any, Final, Literal, cast, overload
from typing import Any, Final, Literal, overload
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -33,13 +33,17 @@ def _apply_service_tier_as_completion_window(body: Mapping[str, object]) -> dict
window: Final[str | None] = (
_SERVICE_TIER_TO_COMPLETION_WINDOW.get(service_tier.lower()) if isinstance(service_tier, str) else None
)
metadata: Final[dict[str, object]] = (
cast(dict[str, object], body["metadata"]) if isinstance(body.get("metadata"), dict) else {}
)
new_body: Final[dict[str, object]] = {key: value for key, value in body.items() if key != "service_tier"}
raw_metadata: Final = body.get("metadata")
metadata: Final[Mapping[str, object]] = raw_metadata if isinstance(raw_metadata, dict) else MappingProxyType({})
new_body: Final[dict[str, object]] = { # mutable-ok: transform_request returns a plain dict
key: value for key, value in body.items() if key != "service_tier"
}
if window is None or "completion_window" in metadata:
return new_body
return {**new_body, "metadata": {**metadata, "completion_window": window}}
return { # mutable-ok: transform_request returns a plain dict
**new_body,
"metadata": {**metadata, "completion_window": window}, # mutable-ok: transform_request returns a plain dict
}
def _service_tier_as_completion_window_enabled(provider: SimpleProviderConfig) -> bool:
@ -126,15 +130,12 @@ def create_config_class(provider: SimpleProviderConfig):
litellm_params: dict[str, object],
headers: dict[str, object],
) -> dict[str, object]:
body: Final[dict[str, object]] = cast(
dict[str, object],
super().transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
),
body: Final = super().transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
if _service_tier_as_completion_window_enabled(provider):
return _apply_service_tier_as_completion_window(body)
@ -169,7 +170,7 @@ def create_config_class(provider: SimpleProviderConfig):
param for param in (*base_params, *extra_params) if param not in excluded_params
)
return list(supported_params)
return list(supported_params) # mutable-ok: get_supported_openai_params contract returns a list
def map_openai_params(
self,
@ -282,15 +283,12 @@ def create_responses_config_class(provider: SimpleProviderConfig):
) -> dict[str, object]:
if provider.special_handling.get("force_store_false"):
response_api_optional_request_params["store"] = False
body: Final[dict[str, object]] = cast(
dict[str, object],
super().transform_responses_api_request(
model=model,
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
),
body: Final = super().transform_responses_api_request(
model=model,
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
if _service_tier_as_completion_window_enabled(provider):
return _apply_service_tier_as_completion_window(body)

View file

@ -5,6 +5,7 @@ JSON-based provider configuration loader for OpenAI-compatible providers.
import json
from collections.abc import Mapping
from pathlib import Path
from types import MappingProxyType
from typing import Final
from litellm._logging import verbose_logger
@ -21,7 +22,7 @@ class SimpleProviderConfig:
self.base_class = data.get("base_class", "openai_gpt")
self.param_mappings = data.get("param_mappings", {})
self.constraints = data.get("constraints", {})
self.special_handling: Mapping[str, object] = data.get("special_handling", {})
self.special_handling: Mapping[str, object] = data.get("special_handling") or MappingProxyType({})
self.supported_endpoints = data.get("supported_endpoints", [])
self.unsupported_params: Final = tuple(data.get("unsupported_params", ()))

View file

@ -31323,6 +31323,8 @@ export interface components {
cache_creation_input_token_cost_above_272k_tokens_flex?: number | null;
/** Cache Creation Input Token Cost Above 272K Tokens Priority */
cache_creation_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Creation Input Token Cost Balanced */
cache_creation_input_token_cost_balanced?: number | null;
/** Cache Creation Input Token Cost Batches */
cache_creation_input_token_cost_batches?: number | null;
/** Cache Creation Input Token Cost Flex */
@ -31351,6 +31353,8 @@ export interface components {
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 512K Tokens */
cache_read_input_token_cost_above_512k_tokens?: number | null;
/** Cache Read Input Token Cost Balanced */
cache_read_input_token_cost_balanced?: number | null;
/** Cache Read Input Token Cost Batches */
cache_read_input_token_cost_batches?: number | null;
/** Cache Read Input Token Cost Flex */
@ -31429,6 +31433,8 @@ export interface components {
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Above 512K Tokens */
input_cost_per_token_above_512k_tokens?: number | null;
/** Input Cost Per Token Balanced */
input_cost_per_token_balanced?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -31516,6 +31522,8 @@ export interface components {
output_cost_per_pixel?: number | null;
/** Output Cost Per Reasoning Token */
output_cost_per_reasoning_token?: number | null;
/** Output Cost Per Reasoning Token Balanced */
output_cost_per_reasoning_token_balanced?: number | null;
/** Output Cost Per Reasoning Token Flex */
output_cost_per_reasoning_token_flex?: number | null;
/** Output Cost Per Reasoning Token Priority */
@ -31552,6 +31560,8 @@ export interface components {
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Above 512K Tokens */
output_cost_per_token_above_512k_tokens?: number | null;
/** Output Cost Per Token Balanced */
output_cost_per_token_balanced?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */
@ -42218,6 +42228,8 @@ export interface components {
cache_creation_input_token_cost_above_272k_tokens_flex?: number | null;
/** Cache Creation Input Token Cost Above 272K Tokens Priority */
cache_creation_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Creation Input Token Cost Balanced */
cache_creation_input_token_cost_balanced?: number | null;
/** Cache Creation Input Token Cost Batches */
cache_creation_input_token_cost_batches?: number | null;
/** Cache Creation Input Token Cost Flex */
@ -42246,6 +42258,8 @@ export interface components {
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 512K Tokens */
cache_read_input_token_cost_above_512k_tokens?: number | null;
/** Cache Read Input Token Cost Balanced */
cache_read_input_token_cost_balanced?: number | null;
/** Cache Read Input Token Cost Batches */
cache_read_input_token_cost_batches?: number | null;
/** Cache Read Input Token Cost Flex */
@ -42324,6 +42338,8 @@ export interface components {
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Above 512K Tokens */
input_cost_per_token_above_512k_tokens?: number | null;
/** Input Cost Per Token Balanced */
input_cost_per_token_balanced?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -42411,6 +42427,8 @@ export interface components {
output_cost_per_pixel?: number | null;
/** Output Cost Per Reasoning Token */
output_cost_per_reasoning_token?: number | null;
/** Output Cost Per Reasoning Token Balanced */
output_cost_per_reasoning_token_balanced?: number | null;
/** Output Cost Per Reasoning Token Flex */
output_cost_per_reasoning_token_flex?: number | null;
/** Output Cost Per Reasoning Token Priority */
@ -42447,6 +42465,8 @@ export interface components {
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Above 512K Tokens */
output_cost_per_token_above_512k_tokens?: number | null;
/** Output Cost Per Token Balanced */
output_cost_per_token_balanced?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */