mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
test(sail): completion_window transform, tier pricing, registry consistency
This commit is contained in:
parent
4756fd6964
commit
7f58a7cd23
3 changed files with 297 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue