test(sail): completion_window transform, tier pricing, registry consistency

This commit is contained in:
shrey kharbanda 2026-09-24 04:57:26 +00:00
parent 4756fd6964
commit 7f58a7cd23
3 changed files with 297 additions and 0 deletions

View file

@ -3789,3 +3789,21 @@ def test_azure_gpt_6_foundry_price_sheet(_local_model_cost_map, model_base):
assert azure_ai_info[field] == base
assert azure_us_info[field] == pytest.approx(1.1 * base)
assert azure_eu_info[field] == pytest.approx(1.2 * base)
@pytest.mark.parametrize(
"base_key,tier,expected",
[
("input_cost_per_token", "balanced", "input_cost_per_token_balanced"),
("output_cost_per_token", "balanced", "output_cost_per_token_balanced"),
("cache_read_input_token_cost", "balanced", "cache_read_input_token_cost_balanced"),
("input_cost_per_token", "flex", "input_cost_per_token_flex"),
("input_cost_per_token", "BALANCED", "input_cost_per_token_balanced"),
("input_cost_per_token", None, "input_cost_per_token"),
("input_cost_per_token", "auto", "input_cost_per_token"),
],
)
def test_get_service_tier_cost_key_balanced(base_key, tier, expected):
from litellm.litellm_core_utils.llm_cost_calc.utils import _get_service_tier_cost_key
assert _get_service_tier_cost_key(base_key, tier) == expected

View file

@ -19,6 +19,52 @@ sys.path.insert(0, workspace_path)
import litellm
def _providers_json() -> dict:
import json
from pathlib import Path
return json.loads(
(Path(litellm.__file__).parent / "llms" / "openai_like" / "providers.json").read_text()
)
class TestProvidersJsonConsistency:
def test_every_slug_is_an_llm_provider_enum_member(self):
from litellm import LlmProviders
known_non_enum_slugs = {
"abliteration",
"aihubmix",
"crusoe",
"empiriolabs",
"gmi",
"llamagate",
"sarvam",
"veniceai",
}
enum_slugs = {provider.value for provider in LlmProviders}
unknown = sorted(set(_providers_json()) - enum_slugs - known_non_enum_slugs)
assert unknown == [], f"providers.json slugs missing from LlmProviders: {unknown}"
assert known_non_enum_slugs - enum_slugs == known_non_enum_slugs
def test_every_chat_completions_slug_resolves_via_get_llm_provider(self):
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
unresolved = []
for slug, config in _providers_json().items():
if "/v1/chat/completions" not in config.get("supported_endpoints", []):
continue
try:
_, provider, _, _ = get_llm_provider(
model=f"{slug}/x", custom_llm_provider=None, api_base=None, api_key=None
)
if provider != slug:
unresolved.append(f"{slug}: resolved to {provider}")
except Exception as exc:
unresolved.append(f"{slug}: {exc}")
assert unresolved == []
class TestJSONProviderLoader:
"""Test JSON provider loading and configuration"""

View file

@ -10,6 +10,12 @@ import respx
import litellm
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.openai_like.dynamic_config import (
create_config_class,
create_responses_config_class,
)
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
SAIL_BASE_URL = "https://api.sailresearch.com/v1"
@ -349,3 +355,230 @@ class TestSailCostTracking:
+ completion_tokens * rates["output_cost_per_token"]
)
assert cost == pytest.approx(expected)
_MESSAGES = [{"role": "user", "content": "hi"}]
def _sail_chat_body(optional_params: dict) -> dict:
config = create_config_class(JSONProviderRegistry.get("sail"))()
return config.transform_request(
model="zai-org/GLM-5.3",
messages=_MESSAGES,
optional_params=dict(optional_params),
litellm_params={},
headers={},
)
def _sail_responses_body(optional_params: dict) -> dict:
config = create_responses_config_class(JSONProviderRegistry.get("sail"))()
return config.transform_responses_api_request(
model="zai-org/GLM-5.3",
input="hi",
response_api_optional_request_params=dict(optional_params),
litellm_params=GenericLiteLLMParams(),
headers={},
)
class TestSailServiceTierAsCompletionWindow:
@pytest.mark.parametrize(
"service_tier,expected_window",
[("flex", "flex"), ("balanced", "balanced"), ("priority", "asap"), ("default", "asap")],
)
def test_service_tier_maps_to_completion_window(self, service_tier: str, expected_window: str):
body = _sail_chat_body({"service_tier": service_tier})
assert "service_tier" not in body
assert body["metadata"]["completion_window"] == expected_window
def test_auto_service_tier_dropped_without_window(self):
body = _sail_chat_body({"service_tier": "auto"})
assert "service_tier" not in body
assert "completion_window" not in body.get("metadata", {})
def test_no_service_tier_leaves_body_alone(self):
body = _sail_chat_body({})
assert "metadata" not in body
def test_caller_completion_window_wins_over_service_tier(self):
body = _sail_chat_body(
{
"service_tier": "flex",
"metadata": {"completion_window": "balanced", "trace": "abc"},
}
)
assert "service_tier" not in body
assert body["metadata"] == {"completion_window": "balanced", "trace": "abc"}
def test_optional_params_not_mutated(self):
optional_params = {"service_tier": "flex"}
_sail_chat_body(optional_params)
assert optional_params == {"service_tier": "flex"}
def test_non_sail_provider_keeps_service_tier(self):
config = create_config_class(JSONProviderRegistry.get("parasail"))()
body = config.transform_request(
model="x",
messages=_MESSAGES,
optional_params={"service_tier": "flex"},
litellm_params={},
headers={},
)
assert body["service_tier"] == "flex"
assert "metadata" not in body
@pytest.mark.parametrize(
"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
):
body = _sail_responses_body({"service_tier": service_tier})
assert "service_tier" not in body
assert body["metadata"]["completion_window"] == expected_window
def test_responses_api_caller_completion_window_wins(self):
body = _sail_responses_body(
{
"service_tier": "flex",
"metadata": {"completion_window": "balanced"},
}
)
assert "service_tier" not in body
assert body["metadata"] == {"completion_window": "balanced"}
@pytest.mark.respx()
def test_sail_completion_end_to_end_sends_window_not_tier(self, respx_mock: respx.Router):
respx_mock.post(SAIL_CHAT_COMPLETIONS).respond(json=_chat_completion_payload())
litellm.completion(
model=MODEL,
messages=[{"role": "user", "content": "hi"}],
service_tier="flex",
)
body = json.loads(respx_mock.calls[0].request.content)
assert "service_tier" not in body
assert body["metadata"] == {"completion_window": "flex"}
@pytest.mark.respx()
def test_sail_completion_window_via_extra_body_and_tier(self, respx_mock: respx.Router):
respx_mock.post(SAIL_CHAT_COMPLETIONS).respond(json=_chat_completion_payload())
litellm.completion(
model=MODEL,
messages=[{"role": "user", "content": "hi"}],
service_tier="flex",
extra_body={"metadata": {"completion_window": "balanced"}},
)
body = json.loads(respx_mock.calls[0].request.content)
assert "service_tier" not in body
assert body["metadata"] == {"completion_window": "balanced"}
def _sail_completion_response(prompt_tokens: int, completion_tokens: int) -> litellm.ModelResponse:
return litellm.ModelResponse(
model="zai-org/GLM-5.3",
choices=[{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
class TestSailTierPricing:
@pytest.mark.parametrize("tier", ["flex", "balanced", "priority"])
def test_service_tier_bills_at_tier_rates(self, tier: str):
rates = litellm.model_cost[MODEL]
prompt_tokens, completion_tokens = 1000, 200
suffix = {"flex": "_flex", "balanced": "_balanced"}.get(tier, "")
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
model=MODEL,
custom_llm_provider="sail",
optional_params={"service_tier": tier},
)
expected = (
prompt_tokens * rates[f"input_cost_per_token{suffix}"]
+ completion_tokens * rates[f"output_cost_per_token{suffix}"]
)
assert cost == pytest.approx(expected)
@pytest.mark.parametrize(
"optional_params",
[
{"extra_body": {"metadata": {"completion_window": "balanced"}}},
{"extra_body": {"metadata": {"completion_window": "flex"}}},
{"metadata": {"completion_window": "balanced"}},
],
ids=["extra_body_balanced", "extra_body_flex", "metadata_balanced"],
)
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"]
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
model=MODEL,
custom_llm_provider="sail",
optional_params=optional_params,
)
expected = (
prompt_tokens * rates[f"input_cost_per_token_{window}"]
+ completion_tokens * rates[f"output_cost_per_token_{window}"]
)
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
cost = litellm.completion_cost(
completion_response=_sail_completion_response(prompt_tokens, completion_tokens),
model=MODEL,
custom_llm_provider="sail",
optional_params={"extra_body": {"metadata": {"completion_window": "asap"}}},
)
expected = (
prompt_tokens * rates["input_cost_per_token"]
+ completion_tokens * rates["output_cost_per_token"]
)
assert cost == pytest.approx(expected)
_TIER_COST_BASES = ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost")
def test_sail_tier_prices_are_monotone_and_complete():
problems: list[str] = []
for name, entry in litellm.model_cost.items():
if not name.startswith("sail/"):
continue
for tier in ("balanced", "flex"):
present = [base for base in _TIER_COST_BASES if entry.get(f"{base}_{tier}") is not None]
if not present:
continue
if len(present) != len(_TIER_COST_BASES):
problems.append(f"{name}: {tier} tier has {present}, expected all of {_TIER_COST_BASES}")
for base in present:
if entry.get(base) is None:
problems.append(f"{name}: has {base}_{tier} but no {base}")
elif entry[f"{base}_{tier}"] > entry[base]:
problems.append(f"{name}: {base}_{tier}={entry[f'{base}_{tier}']} exceeds {base}={entry[base]}")
for base in _TIER_COST_BASES:
flex, balanced = entry.get(f"{base}_flex"), entry.get(f"{base}_balanced")
if flex is not None and balanced is not None and flex > balanced:
problems.append(f"{name}: {base}_flex={flex} exceeds {base}_balanced={balanced}")
assert problems == []