From 7f58a7cd23badab899a004b42714be214b95abb2 Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 04:57:26 +0000 Subject: [PATCH] test(sail): completion_window transform, tier pricing, registry consistency --- .../llm_cost_calc/test_llm_cost_calc_utils.py | 18 ++ .../llms/openai_like/test_json_providers.py | 46 ++++ .../llms/openai_like/test_sail_provider.py | 233 ++++++++++++++++++ 3 files changed, 297 insertions(+) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 781a3a7c4ed..903afb0b3dc 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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 diff --git a/tests/test_litellm/llms/openai_like/test_json_providers.py b/tests/test_litellm/llms/openai_like/test_json_providers.py index d84cc8d3237..4696becfc1b 100644 --- a/tests/test_litellm/llms/openai_like/test_json_providers.py +++ b/tests/test_litellm/llms/openai_like/test_json_providers.py @@ -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""" 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 6fffe0c81f0..7cafab12262 100644 --- a/tests/test_litellm/llms/openai_like/test_sail_provider.py +++ b/tests/test_litellm/llms/openai_like/test_sail_provider.py @@ -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 == []