From 472f8d035aff40e6443e0dbce0cf4f02a36acb9d Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 05:43:03 +0000 Subject: [PATCH] fix(cost): completion_window overrides service_tier only for providers that bill by window --- litellm/cost_calculator.py | 19 +++++- .../llms/openai_like/test_sail_provider.py | 62 ++++++++++++++++--- 2 files changed, 69 insertions(+), 12 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 0fb3e62f378..95a0ba003d2 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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) diff --git a/tests/test_litellm/llms/openai_like/test_sail_provider.py b/tests/test_litellm/llms/openai_like/test_sail_provider.py index 69b7f9cdf6f..008097a417f 100644 --- a/tests/test_litellm/llms/openai_like/test_sail_provider.py +++ b/tests/test_litellm/llms/openai_like/test_sail_provider.py @@ -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)