fix(cost): completion_window overrides service_tier only for providers that bill by window

This commit is contained in:
shrey kharbanda 2026-09-24 05:43:03 +00:00
parent 27be70891d
commit 472f8d035a
2 changed files with 69 additions and 12 deletions

View file

@ -77,6 +77,7 @@ from litellm.llms.openai.cost_calculation import (
from litellm.llms.openai.cost_calculation import (
cost_per_token as openai_cost_per_token,
)
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.llms.perplexity.cost_calculator import (
cost_per_token as perplexity_cost_per_token,
)
@ -976,6 +977,15 @@ def _completion_window_value(metadata: object) -> str | None:
return None
def _provider_bills_by_completion_window(custom_llm_provider: str | None) -> bool:
"""True only for JSON-configured providers that translate ``service_tier`` into a
provider ``metadata.completion_window`` on the wire (currently Sail)."""
if custom_llm_provider is None:
return False
provider: Final = JSONProviderRegistry.get(custom_llm_provider)
return provider is not None and provider.special_handling.get("service_tier_as_completion_window") is True
def _service_tier_from_completion_window(optional_params: dict[str, object]) -> str | None:
"""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)."""
@ -1393,11 +1403,16 @@ def completion_cost(
)
rerank_billed_units: RerankBilledUnits | None = None
# Providers that bill by completion window: an explicit window on the request wins
# over service_tier, matching what the provider actually sees on the wire
if optional_params is not None and _provider_bills_by_completion_window(custom_llm_provider):
window_tier: Final = _service_tier_from_completion_window(optional_params)
if window_tier is not None:
service_tier = window_tier
# Extract service_tier from optional_params if not provided directly
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(optional_params)
service_tier = _normalize_service_tier(service_tier)

View file

@ -432,9 +432,7 @@ class TestSailServiceTierAsCompletionWindow:
"service_tier,expected_window",
[("flex", "flex"), ("balanced", "balanced"), ("priority", "asap")],
)
def test_responses_api_service_tier_maps_to_completion_window(
self, service_tier: str, expected_window: str
):
def test_responses_api_service_tier_maps_to_completion_window(self, service_tier: str, expected_window: str):
body = _sail_responses_body({"service_tier": service_tier})
assert "service_tier" not in body
assert body["metadata"]["completion_window"] == expected_window
@ -524,9 +522,9 @@ class TestSailTierPricing:
def test_completion_window_in_optional_params_bills_at_tier_rates(self, optional_params: dict):
rates = litellm.model_cost[MODEL]
prompt_tokens, completion_tokens = 1000, 200
window = (
optional_params.get("extra_body", {}).get("metadata") or optional_params["metadata"]
)["completion_window"]
window = (optional_params.get("extra_body", {}).get("metadata") or optional_params["metadata"])[
"completion_window"
]
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
@ -541,6 +539,53 @@ class TestSailTierPricing:
)
assert cost == pytest.approx(expected)
@pytest.mark.parametrize(
"service_tier,window",
[("flex", "balanced"), ("balanced", "flex")],
ids=["flex_tier_balanced_window", "balanced_tier_flex_window"],
)
def test_completion_window_overrides_service_tier_on_sail(self, service_tier: str, window: str):
rates = litellm.model_cost[MODEL]
prompt_tokens, completion_tokens = 1000, 200
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
model=MODEL,
custom_llm_provider="sail",
optional_params={
"service_tier": service_tier,
"extra_body": {"metadata": {"completion_window": window}},
},
)
expected = (
prompt_tokens * rates[f"input_cost_per_token_{window}"]
+ completion_tokens * rates[f"output_cost_per_token_{window}"]
)
assert cost == pytest.approx(expected)
@pytest.mark.parametrize(
"optional_params",
[
{"extra_body": {"metadata": {"completion_window": "flex"}}},
{"metadata": {"completion_window": "flex"}},
],
ids=["extra_body_flex_window", "metadata_flex_window"],
)
def test_completion_window_ignored_for_non_sail_provider(self, optional_params: dict):
rates = litellm.model_cost["azure/gpt-5.4"]
prompt_tokens, completion_tokens = 1000, 200
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
model="azure/gpt-5.4",
custom_llm_provider="azure",
optional_params=optional_params,
)
expected = prompt_tokens * rates["input_cost_per_token"] + completion_tokens * rates["output_cost_per_token"]
assert cost == pytest.approx(expected)
def test_completion_window_asap_bills_at_base_rates(self):
rates = litellm.model_cost[MODEL]
prompt_tokens, completion_tokens = 1000, 200
@ -552,10 +597,7 @@ class TestSailTierPricing:
optional_params={"extra_body": {"metadata": {"completion_window": "asap"}}},
)
expected = (
prompt_tokens * rates["input_cost_per_token"]
+ completion_tokens * rates["output_cost_per_token"]
)
expected = prompt_tokens * rates["input_cost_per_token"] + completion_tokens * rates["output_cost_per_token"]
assert cost == pytest.approx(expected)