mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(cost): completion_window overrides service_tier only for providers that bill by window
This commit is contained in:
parent
27be70891d
commit
472f8d035a
2 changed files with 69 additions and 12 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue