test(sail): add unit regression coverage for the Sail provider PR

Covers balanced service tier cost parity, cost-map invariants, suffix fallback, service tier validation, merge_extra_body hooks, Responses logging, transcription provider set, JSON loader merge-base compatibility, get_model_info balanced fields and provider inference before the completion-window gate. MUTATIONS.md records the 15 named production mutations each test kills
This commit is contained in:
shrey kharbanda 2026-09-24 16:19:31 +00:00
parent d7bc17fa47
commit 31ed2482b5
13 changed files with 1228 additions and 5 deletions

590
MUTATIONS.md Normal file
View file

@ -0,0 +1,590 @@
# Mutation evidence for litellm_sail_tests_unit
Tests-only regression branch for PR 42840, created from tip d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 (merge base 0c1c3e18d5250ec3a0e1e3f287e0b93e2906d900). Every new test was run three ways: at the merge base with the changed test files copied in (red unless noted below), at the tip with one named production mutation applied (red), then at the tip with the mutated file restored from d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 (green). Production files were diffed against d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 before every mutation and after every restore, and the final `git status` of the branch shows test files plus this document only
## Merge-base run
Command, run in a detached worktree of the merge base with the eleven changed test files and the new chat transformation test copied in:
```
LITELLM_LOCAL_MODEL_COST_MAP=True pytest -p no:cacheprovider <the tests listed in the PR report> -q
```
Result: 15 failed, 13 passed. The 13 that pass at the merge base are the 12 parametrized `test_completion_cost_balanced_tier_bills_standard_rates_on_rows_without_balanced_keys` cases (the balanced tier does not exist there, so it takes the unknown-tier path and trivially equals the no-tier cost; M01 and M02 are their kill evidence) and `test_completion_cost_ignores_completion_window_for_openai` (the merge base has no completion-window gate at all, so it bills base rates for every provider; M15 is its kill evidence). The 15 failures are import errors for `merge_extra_body`, `KNOWN_REQUEST_SERVICE_TIERS` and `special_handling`, the missing `sail` slug and `_balanced` keys, `sail/` resolving to `LLM Provider NOT provided`, the transcription call routing to OpenAI instead of failing before HTTP, and `extra_body` missing from the Responses logging optional_params
## Named production mutations
Each mutation ran `LITELLM_LOCAL_MODEL_COST_MAP=True pytest -p no:cacheprovider <tests> -q` at the tip, once mutated and once restored. Kill count: 15 of 15
### M01 area 1/3: drop the service-tier suffix fallback in _get_cost_per_unit
File: `litellm/litellm_core_utils/llm_cost_calc/utils.py`
Tests:
```
tests/test_litellm/test_cost_calculator.py::test_completion_cost_balanced_tier_bills_standard_rates_on_rows_without_balanced_keys
tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py::test_get_cost_per_unit_falls_back_from_balanced_key_to_base
```
Mutation (exact string replace):
```
if cost_per_unit is None:
# Check if any service tier suffix exists in the cost key
for suffix in _SERVICE_TIER_SUFFIXES:
```
with:
```
if cost_per_unit is None:
# Check if any service tier suffix exists in the cost key
for suffix in ():
```
Red (mutated), killed: true
```
E AssertionError: (model, no tier, balanced, unknown tier) rows that disagree: (('claude-fable-5', 0.0173, 0.0, 0.0173), ('claude-fable-5-1', 0.017075, 0.0, 0.017075), ('claude-haiku-4-5', 0.00173, 0.0, 0.00173), ('claude-haiku-4-5-20251001', 0.00173, 0.0, 0.00173), ('claude-mythos-5', 0.0173, ...
E assert (('claude-fab....017075), ...) == ()
E
E Left contains 20 more items, first extra item: ('claude-fable-5', 0.0173, 0.0, 0.0173)
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/litellm_core_utils/llm_cost_calc/utils.py`, then the same command
Green (restored), exit 0: `13 passed in 1.64s`
### M02 area 1/3/4: remove BALANCED from the ServiceTier enum
File: `litellm/types/utils.py`
Tests:
```
tests/test_litellm/integrations/test_prometheus_service_tier_label.py::test_requested_tier_recognizes_balanced_and_drops_garbage_without_touching_the_other_tiers
tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py::test_threshold_keys_exclude_balanced_variants_and_still_parse_k_thresholds
```
Mutation (exact string replace):
```
AUTO = "auto"
BALANCED = "balanced"
FLEX = "flex"
```
with:
```
AUTO = "auto"
FLEX = "flex"
```
Red (mutated), killed: true
```
E AttributeError: type object 'ServiceTier' has no attribute 'BALANCED'
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/types/utils.py`, then the same command
Green (restored), exit 0: `2 passed in 0.65s`
### M03 area 3: map the balanced tier onto the flex cost suffix
File: `litellm/litellm_core_utils/llm_cost_calc/utils.py`
Tests:
```
tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py::test_threshold_keys_exclude_balanced_variants_and_still_parse_k_thresholds
```
Mutation (exact string replace):
```
ServiceTier.BALANCED.value: ServiceTier.BALANCED.value,
```
with:
```
ServiceTier.BALANCED.value: ServiceTier.FLEX.value,
```
Red (mutated), killed: true
```
E assert (4e-06, 6e-06) == (7e-06, 1.4e-05)
E
E At index 0 diff: 4e-06 != 7e-06
E Use -v to get more diff
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/litellm_core_utils/llm_cost_calc/utils.py`, then the same command
Green (restored), exit 0: `1 passed in 0.42s`
### M04 area 2: put a balanced key on a non-sail cost-map row
File: `model_prices_and_context_window.json`
Tests:
```
tests/test_litellm/test_model_prices_schema.py::test_balanced_tier_keys_live_only_on_sail_rows_next_to_their_base_key
```
Mutation (regex replace):
```
("gpt-4o-mini": \{\n)
```
with:
```
\1 "input_cost_per_token_balanced": 1e-07,\n
```
Red (mutated), killed: true
```
E AssertionError: (model, balanced key, litellm_provider, has base key) for balanced tier pricing that is off a sail/ row or missing its base key: (('gpt-4o-mini', 'input_cost_per_token_balanced', 'openai', True),)
E assert (('gpt-4o-min...enai', True),) == ()
E
E Left contains one more item: ('gpt-4o-mini', 'input_cost_per_token_balanced', 'openai', True)
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- model_prices_and_context_window.json`, then the same command
Green (restored), exit 0: `1 passed in 0.40s`
### M05 area 5: BaseConfig.merge_extra_body deep-merges nested metadata
File: `litellm/llms/base_llm/chat/transformation.py`
Tests:
```
tests/test_litellm/llms/base_llm/chat/test_base_chat_transformation.py::test_default_merge_extra_body_shallow_merges_and_lets_extra_body_replace_nested_metadata
```
Mutation (exact string replace):
```
return {**request, **extra_body} if extra_body else request # mutable-ok: wire request body is a plain dict
```
with:
```
if not extra_body:
return request
merged = {**request, **extra_body}
if isinstance(request.get('metadata'), dict) and isinstance(extra_body.get('metadata'), dict):
merged['metadata'] = {**request['metadata'], **extra_body['metadata']}
return merged
```
Red (mutated), killed: true
```
E AssertionError: assert {'messages': ...re': 0.2, ...} == {'messages': ...re': 0.2, ...}
E
E Omitting 4 identical items, use -vv to show
E Differing items:
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/base_llm/chat/transformation.py`, then the same command
Green (restored), exit 0: `1 passed in 0.41s`
### M06 area 5: BaseResponsesAPIConfig.merge_extra_body deep-merges nested metadata
File: `litellm/llms/base_llm/responses/transformation.py`
Tests:
```
tests/test_litellm/llms/base_llm/responses/test_transformation.py::test_default_merge_extra_body_shallow_merges_and_lets_extra_body_replace_nested_metadata
```
Mutation (exact string replace):
```
return {**request, **extra_body} if extra_body else request # mutable-ok: wire request body is a plain dict
```
with:
```
if not extra_body:
return request
merged = {**request, **extra_body}
if isinstance(request.get('metadata'), dict) and isinstance(extra_body.get('metadata'), dict):
merged['metadata'] = {**request['metadata'], **extra_body['metadata']}
return merged
```
Red (mutated), killed: true
```
E AssertionError: assert {'input': 'hi...el': 'm', ...} == {'input': 'hi...el': 'm', ...}
E
E Omitting 4 identical items, use -vv to show
E Differing items:
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/base_llm/responses/transformation.py`, then the same command
Green (restored), exit 0: `1 passed in 0.39s`
### M07 area 5: llm_http_handler.completion inlines the shallow merge instead of calling the hook
File: `litellm/llms/custom_httpx/llm_http_handler.py`
Tests:
```
tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py::test_completion_merges_extra_body_through_the_real_config_hook
```
Mutation (exact string replace):
```
data: Final = provider_config.merge_extra_body(transformed, extra_body)
```
with:
```
data: Final = {**transformed, **extra_body} if extra_body else transformed
```
Red (mutated), killed: true
```
E AssertionError: assert {'messages': ...re': 0.2, ...} == {'messages': ...d': 'x'}, ...}
E
E Omitting 3 identical items, use -vv to show
E Differing items:
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/custom_httpx/llm_http_handler.py`, then the same command
Green (restored), exit 0: `1 passed in 0.43s`
### M08 area 5: llm_http_handler response_api_handler (sync + async) inlines the shallow merge instead of calling the hook
File: `litellm/llms/custom_httpx/llm_http_handler.py`
Tests:
```
tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py::test_response_api_handler_merges_extra_body_through_the_real_config_hook
tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py::test_async_response_api_handler_merges_extra_body_through_the_real_config_hook
```
Mutation (exact string replace):
```
data = responses_api_provider_config.merge_extra_body(data, extra_body)
```
with:
```
data = {**data, **extra_body} if extra_body else data
```
Red (mutated), killed: true
```
E AssertionError: assert {'input': 'hi...el': 'm', ...} == {'input': 'hi...el': 'm', ...}
E
E Omitting 3 identical items, use -vv to show
E Differing items:
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/custom_httpx/llm_http_handler.py`, then the same command
Green (restored), exit 0: `2 passed in 0.45s`
### M09 area 6: responses/main.py stops passing extra_body into logging optional_params
File: `litellm/responses/main.py`
Tests:
```
tests/test_litellm/responses/test_responses_api_request_body.py::test_aresponses_logs_extra_body_once_in_optional_params_and_leaves_the_wire_body_alone
```
Mutation (exact string replace):
```
**(
{"extra_body": extra_body} if extra_body else {}
), # mutable-ok: update_from_kwargs stores a plain dict
```
with:
```
```
Red (mutated), killed: true
```
E AssertionError: assert ({'max_output...warg': '0'}},) == ({'extra_body...warg': '0'}},)
E
E At index 0 diff: {'max_output_tokens': 50, 'metadata': {'from_kwarg': '0'}} != {'max_output_tokens': 50, 'metadata': {'from_kwarg': '0'}, 'extra_body': {'metadata': {'from_extra_body': '1'}, 'vendor_only_field': 'x'}}
E Use -v to get more diff
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/responses/main.py`, then the same command
Green (restored), exit 0: `1 passed in 0.84s`
### M10 area 7: keep sail inside OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS
File: `litellm/constants.py`
Tests:
```
tests/test_litellm/test_constants.py::test_openai_transcription_providers_are_the_openai_compatible_set_minus_sail
tests/test_litellm/llms/openai_like/test_sail_provider.py::TestSailRequestShape::test_sail_transcription_rejected_without_hitting_sail
```
Mutation (exact string replace):
```
frozenset(openai_compatible_providers) - frozenset(("sail",))
```
with:
```
frozenset(openai_compatible_providers)
```
Red (mutated), killed: true
```
E AssertionError: assert frozenset({'a...ure_ai', ...}) == frozenset({'a...ure_ai', ...})
E
E Extra items in the left set:
E 'sail'
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/constants.py`, then the same command
Green (restored), exit 0: `2 passed in 0.84s`
### M11 area 8: JSON loader drops the special_handling block of every provider
File: `litellm/llms/openai_like/json_loader.py`
Tests:
```
tests/test_litellm/llms/openai_like/test_json_providers.py::test_json_loader_yields_the_merge_base_params_and_special_handling_for_every_non_sail_slug
```
Mutation (exact string replace):
```
self.special_handling: Mapping[str, object] = data.get("special_handling") or MappingProxyType({})
```
with:
```
self.special_handling: Mapping[str, object] = MappingProxyType({})
```
Red (mutated), killed: true
```
E AssertionError: assert {'abliteratio...', ...]}, ...} == {'abliteratio...', ...]}, ...}
E
E Omitting 26 identical items, use -vv to show
E Differing items:
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/openai_like/json_loader.py`, then the same command
Green (restored), exit 0: `1 passed in 7.49s`
### M12 area 8: dynamic_config drops temperature from every JSON provider's supported params
File: `litellm/llms/openai_like/dynamic_config.py`
Tests:
```
tests/test_litellm/llms/openai_like/test_json_providers.py::test_json_loader_yields_the_merge_base_params_and_special_handling_for_every_non_sail_slug
```
Mutation (exact string replace):
```
excluded_params: Final = frozenset(tool_params) | frozenset(provider.unsupported_params)
```
with:
```
excluded_params: Final = frozenset(tool_params) | frozenset(provider.unsupported_params) | {"temperature"}
```
Red (mutated), killed: true
```
E AssertionError: assert {'abliteratio...', ...]}, ...} == {'abliteratio...', ...]}, ...}
E
E Differing items:
E {'apertis': {'special_handling': {}, 'supported_openai_params': ['audio', 'extra_headers', 'frequency_penalty', 'logit_bias', 'logprobs', 'max_completion_tokens', ...]}} != {'apertis': {'special_handling': {}, 'supported_openai_params': ['audio', 'extra_headers', 'frequency_penalty', 'logi ...
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/llms/openai_like/dynamic_config.py`, then the same command
Green (restored), exit 0: `1 passed in 7.78s`
### M13 area 9: get_model_info never fills input_cost_per_token_balanced
File: `litellm/utils.py`
Tests:
```
tests/test_litellm/test_utils.py::test_get_model_info_reports_balanced_tier_prices_only_where_the_row_has_them
```
Mutation (exact string replace):
```
input_cost_per_token_balanced=_model_info.get("input_cost_per_token_balanced", None),
```
with:
```
input_cost_per_token_balanced=None,
```
Red (mutated), killed: true
```
E assert (None, 2.5e-0...7, None, None) == (5e-07, 2.5e-...7, None, None)
E
E At index 0 diff: None != 5e-07
E Use -v to get more diff
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/utils.py`, then the same command
Green (restored), exit 0: `1 passed, 6 warnings in 0.93s`
### M14 area 10: completion_cost runs the completion-window gate before inferring the provider
File: `litellm/cost_calculator.py`
Tests:
```
tests/test_litellm/test_cost_calculator.py::test_completion_cost_infers_sail_from_model_before_the_completion_window_gate
```
Mutation (exact string replace):
```
if custom_llm_provider is None:
try:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model
) # strip the llm provider from the model name -> for image gen cost calculation
except Exception as e:
verbose_logger.debug(
"litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - %s",
e,
)
service_tier = _service_tier_billed_by_completion_window(
service_tier=service_tier,
optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
)
```
with:
```
service_tier = _service_tier_billed_by_completion_window(
service_tier=service_tier,
optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
)
if custom_llm_provider is None:
try:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model
) # strip the llm provider from the model name -> for image gen cost calculation
except Exception as e:
verbose_logger.debug(
"litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - %s",
e,
)
```
Red (mutated), killed: true
```
E assert 0.001596 == 0.00075999999...9999 ± 1.0e-12
E
E comparison failed
E Obtained: 0.001596
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/cost_calculator.py`, then the same command
Green (restored), exit 0: `1 passed in 0.55s`
### M15 area 10: completion-window gate ignores the provider flag (every provider bills by completion_window)
File: `litellm/cost_calculator.py`
Tests:
```
tests/test_litellm/test_cost_calculator.py::test_completion_cost_ignores_completion_window_for_openai
```
Mutation (exact string replace):
```
if optional_params is None or not _provider_bills_by_completion_window(custom_llm_provider):
```
with:
```
if optional_params is None:
```
Red (mutated), killed: true
```
E assert 0.01 == 0.02 ± 2.0e-11
E
E comparison failed
E Obtained: 0.01
```
Restore: `git checkout d7bc17fa47ee9acacd441e6b3349c7aa9e1c6c62 -- litellm/cost_calculator.py`, then the same command
Green (restored), exit 0: `1 passed in 0.45s`

View file

@ -16,6 +16,7 @@ import pytest
from litellm.integrations.prometheus import PrometheusLogger
from litellm.litellm_core_utils.service_tier_utils import (
KNOWN_REQUEST_SERVICE_TIERS,
get_requested_service_tier,
get_service_tier_from_standard_logging_payload,
)
from litellm.types.integrations.prometheus import (
@ -232,6 +233,40 @@ async def test_success_event_emits_service_tier_on_latency_and_spend_metrics():
_clear_prometheus_registry()
def test_requested_tier_recognizes_balanced_and_drops_garbage_without_touching_the_other_tiers():
tiers_before_balanced = (
"auto",
"batch",
"default",
"fast",
"flex",
"priority",
"scale",
"standard",
"standard_only",
"ultrafast",
)
garbage = ("balanced ", "Balanced", "balanced_tier", "", "asap", 1, None, ["balanced"])
assert KNOWN_REQUEST_SERVICE_TIERS == frozenset((*tiers_before_balanced, "balanced"))
assert (
get_requested_service_tier(_standard_logging_payload(model_parameters={"service_tier": "balanced"}))
== "balanced"
)
assert (
tuple(
get_requested_service_tier(_standard_logging_payload(model_parameters={"service_tier": tier}))
for tier in tiers_before_balanced
)
== tiers_before_balanced
)
assert tuple(
get_requested_service_tier(_standard_logging_payload(model_parameters={"service_tier": value}))
for value in garbage
) == (None,) * len(garbage)
assert get_requested_service_tier(_standard_logging_payload(model_parameters="balanced")) is None
def test_allowlist_covers_every_modeled_service_tier():
"""A tier modeled for cost calculation is real traffic, so it must resolve
rather than being dropped as an unknown caller value."""

View file

@ -2428,6 +2428,44 @@ def test_threshold_keys_exclude_service_tier_variants():
assert prompt_base == 3e-6
def test_get_cost_per_unit_falls_back_from_balanced_key_to_base():
from litellm.litellm_core_utils.llm_cost_calc.utils import _get_cost_per_unit
without_balanced: Final = cast(ModelInfo, {"input_cost_per_token": 2e-6})
with_balanced: Final = cast(ModelInfo, {"input_cost_per_token_balanced": 5e-6, "input_cost_per_token": 2e-6})
assert _get_cost_per_unit(without_balanced, "input_cost_per_token_balanced") == 2e-6
assert _get_cost_per_unit(with_balanced, "input_cost_per_token_balanced") == 5e-6
assert _get_cost_per_unit(with_balanced, "input_cost_per_token") == 2e-6
def test_threshold_keys_exclude_balanced_variants_and_still_parse_k_thresholds():
model_info: Final = cast(
ModelInfo,
{
"input_cost_per_token": 1e-6,
"input_cost_per_token_above_128k_tokens": 3e-6,
"input_cost_per_token_above_200k_tokens": 4e-6,
"input_cost_per_token_above_200k_tokens_balanced": 7e-6,
"input_cost_per_token_above_300k_tokens_balanced": 9e-6,
"output_cost_per_token": 2e-6,
"output_cost_per_token_above_128k_tokens": 5e-6,
"output_cost_per_token_above_200k_tokens": 6e-6,
"output_cost_per_token_above_200k_tokens_balanced": 14e-6,
"output_cost_per_token_above_300k_tokens_balanced": 18e-6,
},
)
above_300k: Final = Usage(prompt_tokens=350_000, completion_tokens=1_000, total_tokens=351_000)
between_128k_and_200k: Final = Usage(prompt_tokens=150_000, completion_tokens=1_000, total_tokens=151_000)
below_128k: Final = Usage(prompt_tokens=100_000, completion_tokens=1_000, total_tokens=101_000)
assert _get_token_base_cost(model_info=model_info, usage=above_300k)[:2] == (4e-6, 6e-6)
assert _get_token_base_cost(model_info=model_info, usage=between_128k_and_200k)[:2] == (3e-6, 5e-6)
assert _get_token_base_cost(model_info=model_info, usage=below_128k)[:2] == (1e-6, 2e-6)
assert _get_token_base_cost(model_info=model_info, usage=above_300k, service_tier="balanced")[:2] == (7e-6, 14e-6)
@pytest.mark.parametrize(
"model,custom_llm_provider,reasoning_tokens,cached_tokens",
[

View file

@ -0,0 +1,21 @@
"""The shared chat config contract."""
from typing import Final
import litellm
def test_default_merge_extra_body_shallow_merges_and_lets_extra_body_replace_nested_metadata() -> None:
cfg: Final = litellm.OpenAIGPTConfig()
request: Final = {"model": "m", "messages": [], "metadata": {"from_request": "1"}, "temperature": 0.2}
extra_body: Final = {"metadata": {"from_extra_body": "2"}, "vendor_only_field": "x"}
assert cfg.merge_extra_body(dict(request), extra_body) == {
"model": "m",
"messages": [],
"metadata": {"from_extra_body": "2"},
"temperature": 0.2,
"vendor_only_field": "x",
}
assert cfg.merge_extra_body(dict(request), None) == request
assert cfg.merge_extra_body(dict(request), {}) == request

View file

@ -33,3 +33,19 @@ async def test_default_async_transform_delegates_to_the_sync_transform():
)
assert async_body == sync_body
assert "cache_control" not in async_body["input"][0]["content"][0]
def test_default_merge_extra_body_shallow_merges_and_lets_extra_body_replace_nested_metadata():
cfg = OpenAIResponsesAPIConfig()
request = {"model": "m", "input": "hi", "metadata": {"from_request": "1"}, "max_output_tokens": 5}
extra_body = {"metadata": {"from_extra_body": "2"}, "vendor_only_field": "x"}
assert cfg.merge_extra_body(dict(request), extra_body) == {
"model": "m",
"input": "hi",
"metadata": {"from_extra_body": "2"},
"max_output_tokens": 5,
"vendor_only_field": "x",
}
assert cfg.merge_extra_body(dict(request), None) == request
assert cfg.merge_extra_body(dict(request), {}) == request

View file

@ -4,6 +4,8 @@ import json
import logging
import threading
import time
from collections.abc import Mapping
from datetime import datetime
from typing import Final
from unittest.mock import AsyncMock, Mock, patch
@ -40,6 +42,7 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
from litellm.types.llms.openai import ResponsesAPIResponse
@ -4176,3 +4179,187 @@ async def test_chat_completion_agentic_followup_does_not_repeat_request_params_f
assert followup_calls[0]["temperature"] == 0.2
assert followup_calls[0]["api_base"] == "https://a"
assert followup_calls[0]["model"] == "openai/gpt-5"
_EXTRA_BODY: Final[Mapping[str, object]] = {"metadata": {"from_extra_body": "1"}, "vendor_only_field": "x"}
_RESPONSES_API_RESPONSE: Final[Mapping[str, object]] = {
"id": "resp_1",
"object": "response",
"created_at": 1734366691,
"status": "completed",
"model": "m",
"output": [
{
"type": "message",
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "Done.", "annotations": []}],
}
],
"parallel_tool_calls": True,
"usage": {
"input_tokens": 10,
"output_tokens": 5,
"total_tokens": 15,
"output_tokens_details": {"reasoning_tokens": 0},
},
"error": None,
"incomplete_details": None,
"instructions": None,
"metadata": None,
"temperature": None,
"tool_choice": "auto",
"tools": [],
"top_p": None,
"max_output_tokens": None,
"previous_response_id": None,
"reasoning": None,
"truncation": None,
"user": None,
}
class _NestingChatConfig(litellm.OpenAIGPTConfig):
def merge_extra_body(
self, request: dict[str, object], extra_body: Mapping[str, object] | None
) -> dict[str, object]:
return {**request, "nested_by_hook": dict(extra_body or {})}
class _NestingResponsesConfig(OpenAIResponsesAPIConfig):
def merge_extra_body(
self, request: dict[str, object], extra_body: Mapping[str, object] | None
) -> dict[str, object]:
return {**request, "nested_by_hook": dict(extra_body or {})}
class _RecordingTransport(httpx.BaseTransport, httpx.AsyncBaseTransport):
def __init__(self, response_json: Mapping[str, object]) -> None:
self._response_json = response_json
self.bodies: tuple[dict, ...] = ()
def _record(self, request: httpx.Request) -> httpx.Response:
self.bodies = (*self.bodies, json.loads(request.content))
return httpx.Response(200, json=dict(self._response_json), request=request)
def handle_request(self, request: httpx.Request) -> httpx.Response:
return self._record(request)
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
return self._record(request)
def _real_logging_obj(call_type: str):
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
return LitellmLogging(
model="m",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type=call_type,
start_time=datetime.now(),
litellm_call_id="extra-body-call",
function_id="extra-body-fn",
)
def _chat_wire_body(config: BaseConfig) -> dict:
transport = _RecordingTransport(A_COMPLETION)
BaseLLMHTTPHandler().completion(
model="m",
messages=[{"role": "user", "content": "hi"}],
api_base="https://llm.example/v1/chat/completions",
custom_llm_provider="openai",
model_response=ModelResponse(),
encoding=None,
logging_obj=_real_logging_obj("completion"),
optional_params={"temperature": 0.2, "metadata": {"from_kwarg": "0"}, "extra_body": dict(_EXTRA_BODY)},
timeout=10.0,
litellm_params={},
acompletion=False,
api_key="sk-test",
client=HTTPHandler(client=httpx.Client(transport=transport)),
provider_config=config,
)
assert len(transport.bodies) == 1
return transport.bodies[0]
def test_completion_merges_extra_body_through_the_real_config_hook():
assert _chat_wire_body(litellm.OpenAIGPTConfig()) == {
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.2,
"metadata": {"from_extra_body": "1"},
"vendor_only_field": "x",
}
assert _chat_wire_body(_NestingChatConfig()) == {
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.2,
"metadata": {"from_kwarg": "0"},
"nested_by_hook": dict(_EXTRA_BODY),
}
def _responses_wire_body(config: BaseResponsesAPIConfig) -> dict:
transport = _RecordingTransport(_RESPONSES_API_RESPONSE)
BaseLLMHTTPHandler().response_api_handler(
model="m",
input="hi",
responses_api_provider_config=config,
response_api_optional_request_params={"max_output_tokens": 50, "metadata": {"from_kwarg": "0"}},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
logging_obj=_real_logging_obj("responses"),
extra_body=dict(_EXTRA_BODY),
client=HTTPHandler(client=httpx.Client(transport=transport)),
)
assert len(transport.bodies) == 1
return transport.bodies[0]
async def _async_responses_wire_body(config: BaseResponsesAPIConfig) -> dict:
transport = _RecordingTransport(_RESPONSES_API_RESPONSE)
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=transport)
await BaseLLMHTTPHandler().async_response_api_handler(
model="m",
input="hi",
responses_api_provider_config=config,
response_api_optional_request_params={"max_output_tokens": 50, "metadata": {"from_kwarg": "0"}},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
logging_obj=_real_logging_obj("responses"),
extra_body=dict(_EXTRA_BODY),
client=client,
)
assert len(transport.bodies) == 1
return transport.bodies[0]
_SHALLOW_MERGED_RESPONSES_BODY: Final[Mapping[str, object]] = {
"model": "m",
"input": "hi",
"max_output_tokens": 50,
"metadata": {"from_extra_body": "1"},
"vendor_only_field": "x",
}
_HOOK_NESTED_RESPONSES_BODY: Final[Mapping[str, object]] = {
"model": "m",
"input": "hi",
"max_output_tokens": 50,
"metadata": {"from_kwarg": "0"},
"nested_by_hook": dict(_EXTRA_BODY),
}
def test_response_api_handler_merges_extra_body_through_the_real_config_hook():
assert _responses_wire_body(OpenAIResponsesAPIConfig()) == _SHALLOW_MERGED_RESPONSES_BODY
assert _responses_wire_body(_NestingResponsesConfig()) == _HOOK_NESTED_RESPONSES_BODY
async def test_async_response_api_handler_merges_extra_body_through_the_real_config_hook():
assert await _async_responses_wire_body(OpenAIResponsesAPIConfig()) == _SHALLOW_MERGED_RESPONSES_BODY
assert await _async_responses_wire_body(_NestingResponsesConfig()) == _HOOK_NESTED_RESPONSES_BODY

View file

@ -2,8 +2,11 @@
Tests for JSON-based provider configuration system.
"""
import json
import os
import subprocess
import sys
from pathlib import Path
from unittest.mock import patch
try:
@ -20,9 +23,6 @@ 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()
)
@ -555,3 +555,58 @@ if __name__ == "__main__":
print("\n" + "=" * 50)
print("✓ All tests passed!")
print("=" * 50)
_LOADER_SNAPSHOT_SCRIPT = """
import json
import os
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
import litellm
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
snapshot = {
slug: {
"supported_openai_params": sorted(create_config_class(provider)().get_supported_openai_params(model="snapshot-model")),
"special_handling": dict(provider.special_handling),
}
for slug in sorted(JSONProviderRegistry.list_providers())
for provider in [JSONProviderRegistry.get(slug)]
}
print(json.dumps({"litellm_file": litellm.__file__, "snapshot": snapshot}, sort_keys=True))
"""
def _loader_snapshot(checkout: Path) -> dict:
completed = subprocess.run(
[sys.executable, "-P", "-c", _LOADER_SNAPSHOT_SCRIPT],
cwd=checkout,
env={**os.environ, "PYTHONPATH": str(checkout), "LITELLM_LOCAL_MODEL_COST_MAP": "True"},
capture_output=True,
text=True,
check=True,
timeout=600,
)
payload = json.loads(completed.stdout.splitlines()[-1])
assert Path(payload["litellm_file"]).resolve() == (checkout / "litellm" / "__init__.py").resolve()
return payload["snapshot"]
def test_json_loader_yields_the_merge_base_params_and_special_handling_for_every_non_sail_slug(tmp_path):
tip = Path(litellm.__file__).parents[1]
merge_base = subprocess.run(
["git", "-C", str(tip), "merge-base", "HEAD", "origin/main"], capture_output=True, text=True, check=True
).stdout.strip()
base = tmp_path / "merge_base"
subprocess.run(["git", "-C", str(tip), "worktree", "add", "--detach", str(base), merge_base], check=True)
try:
base_snapshot = _loader_snapshot(base)
tip_snapshot = _loader_snapshot(tip)
finally:
subprocess.run(["git", "-C", str(tip), "worktree", "remove", "--force", str(base)], check=True)
assert {slug: entry for slug, entry in tip_snapshot.items() if slug != "sail"} == base_snapshot
assert "sail" in tip_snapshot
assert "sail" not in base_snapshot
assert tip_snapshot["sail"]["special_handling"] == {"service_tier_as_completion_window": True}
assert len(base_snapshot) > 0

View file

@ -251,10 +251,13 @@ class TestSailRequestShape:
wav_file.writeframes(b"\x00" * 1600)
wav.seek(0)
with pytest.raises(ValueError, match="Unmapped provider"):
with pytest.raises(ValueError, match="Unmapped provider") as excinfo:
litellm.transcription(model=MODEL, file=wav)
assert not route.called
assert type(excinfo.value) is ValueError
assert str(excinfo.value) == "Unmapped provider passed in. Unable to get the response."
assert route.call_count == 0
assert respx_mock.calls.call_count == 0
@pytest.mark.respx()
def test_sail_unsupported_params_dropped_with_drop_params(self, respx_mock: respx.Router):

View file

@ -4,8 +4,10 @@ over the wire and surface provider errors correctly. Expected JSON bodies are st
in expected_responses_api_request/.
"""
import asyncio
import copy
import json
import re
from pathlib import Path
from importlib import import_module
from typing import Final
@ -16,6 +18,7 @@ import pytest
import respx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -242,6 +245,70 @@ async def test_aresponses_forwards_non_enum_reasoning_effort(
assert response.output[0].content[0].text == "Done."
class _OptionalParamsCapture(CustomLogger):
def __init__(self, litellm_call_id: str) -> None:
super().__init__()
self.litellm_call_id: Final = litellm_call_id
self.logged: tuple[dict, ...] = ()
self.seen: Final = asyncio.Event()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
if kwargs["litellm_call_id"] != self.litellm_call_id:
return
self.logged = (*self.logged, kwargs["optional_params"])
self.seen.set()
def _latest_openai_responses_model() -> str:
return max(
(
name
for name, row in litellm.model_cost.items()
if name.startswith("gpt-")
and row.get("litellm_provider") == "openai"
and row.get("mode") == "chat"
and "/v1/responses" in row.get("supported_endpoints", ())
),
key=lambda name: tuple(float(part) for part in re.findall(r"\d+(?:\.\d+)?", name)),
)
@pytest.mark.asyncio
async def test_aresponses_logs_extra_body_once_in_optional_params_and_leaves_the_wire_body_alone(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
):
monkeypatch.setenv("OPENAI_API_KEY", "fake-api-key")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
capture: Final = _OptionalParamsCapture("extra-body-logging-call")
monkeypatch.setattr(litellm, "callbacks", [capture])
litellm.in_memory_llm_clients_cache.flush_cache()
model: Final = _latest_openai_responses_model()
upstream: Final = respx_mock.post("https://api.openai.com/v1/responses").mock(
return_value=httpx.Response(200, json=_minimal_responses_api_payload("resp_extra_body", model))
)
extra_body: Final = {"metadata": {"from_extra_body": "1"}, "vendor_only_field": "x"}
await litellm.aresponses(
model=f"openai/{model}",
input="hi",
max_output_tokens=50,
metadata={"from_kwarg": "0"},
extra_body=copy.deepcopy(extra_body),
litellm_call_id=capture.litellm_call_id,
)
await asyncio.wait_for(capture.seen.wait(), timeout=10)
assert upstream.call_count == 1
assert json.loads(upstream.calls[0].request.read()) == {
"model": model,
"input": "hi",
"max_output_tokens": 50,
"metadata": {"from_extra_body": "1"},
"vendor_only_field": "x",
}
assert capture.logged == ({"max_output_tokens": 50, "metadata": {"from_kwarg": "0"}, "extra_body": extra_body},)
@pytest.mark.asyncio
async def test_acompletion_with_tools_forwards_non_enum_reasoning_effort_over_the_bridge(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter

View file

@ -68,3 +68,11 @@ def _build_constant_env_var_map() -> dict[str, str]:
env_var_map[constant_name] = env_var_name
return env_var_map
def test_openai_transcription_providers_are_the_openai_compatible_set_minus_sail():
assert "sail" in constants.openai_compatible_providers
assert constants.OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS == frozenset(("openai",)) | (
frozenset(constants.openai_compatible_providers) - frozenset(("sail",))
)
assert "sail" not in constants.OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS

View file

@ -1,5 +1,9 @@
import datetime
import json
import math
import re
import time
from collections.abc import Callable, Mapping
from pathlib import Path
from types import MappingProxyType, SimpleNamespace
from typing import Final, cast
@ -4697,3 +4701,142 @@ def test_completion_cost_is_zero_when_explicit_rates_are_zero(monkeypatch: pytes
)
assert cost == 0.0
_PRICES_PATH: Final = Path(__file__).parents[2] / "model_prices_and_context_window.json"
_PRICES: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType(json.loads(_PRICES_PATH.read_text()))
_BALANCED_PARITY_FAMILIES: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"openai": ("openai",),
"anthropic": ("anthropic",),
"gemini": ("gemini",),
"bedrock": ("bedrock", "bedrock_converse"),
"groq": ("groq",),
"azure": ("azure",),
}
)
def _version_key(name: str) -> tuple[float, ...]:
return tuple(float(part) for part in re.findall(r"\d+(?:\.\d+)?", name))
def _priced_chat_rows(providers: tuple[str, ...]) -> tuple[str, ...]:
return tuple(
sorted(
name
for name, row in _PRICES.items()
if row.get("litellm_provider") in providers
and row.get("mode") == "chat"
and isinstance(row.get("input_cost_per_token"), float)
and row["input_cost_per_token"] > 0
and isinstance(row.get("cache_read_input_token_cost"), float)
)
)
def _latest_row(prefix: str, provider: str, *required_keys: str) -> str:
return max(
(
name
for name, row in _PRICES.items()
if name.startswith(prefix)
and row.get("litellm_provider") == provider
and row.get("mode") == "chat"
and all(key in row for key in required_keys)
),
key=_version_key,
)
def _plain_usage() -> Usage:
return Usage(prompt_tokens=1000, completion_tokens=200, total_tokens=1200)
def _cached_tier_usage() -> Usage:
return Usage(
prompt_tokens=1000,
completion_tokens=200,
total_tokens=1200,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=300),
)
def _cost_for_tier(model: str, custom_llm_provider: str, usage: Usage, service_tier: str | None) -> float:
return completion_cost(
completion_response=ModelResponse(model=model, choices=[], usage=usage),
model=model,
custom_llm_provider=custom_llm_provider,
optional_params={"service_tier": service_tier} if service_tier is not None else None,
)
@pytest.mark.parametrize("usage_factory", [_plain_usage, _cached_tier_usage], ids=["plain", "cached"])
@pytest.mark.parametrize("family", sorted(_BALANCED_PARITY_FAMILIES))
def test_completion_cost_balanced_tier_bills_standard_rates_on_rows_without_balanced_keys(
_local_model_cost_map, family: str, usage_factory: Callable[[], Usage]
):
"""A row without ``*_balanced`` keys bills ``balanced`` and an unknown tier exactly like no tier."""
rows: Final = _priced_chat_rows(_BALANCED_PARITY_FAMILIES[family])
assert rows, f"no priced chat rows for {family}"
costs: Final = tuple(
(
name,
_cost_for_tier(name, family, usage_factory(), None),
_cost_for_tier(name, family, usage_factory(), "balanced"),
_cost_for_tier(name, family, usage_factory(), "not-a-tier"),
)
for name in rows
)
drifted: Final = tuple(
entry
for entry in costs
if not math.isclose(entry[2], entry[1], rel_tol=1e-9) or not math.isclose(entry[3], entry[1], rel_tol=1e-9)
)
assert drifted == (), f"(model, no tier, balanced, unknown tier) rows that disagree: {drifted}"
assert min(entry[1] for entry in costs if entry[1] > 0) > 0
assert len([entry for entry in costs if entry[1] > 0]) >= len(costs) - 1
def test_completion_cost_infers_sail_from_model_before_the_completion_window_gate(_local_model_cost_map):
model: Final = _latest_row("sail/", "sail", "input_cost_per_token_flex", "output_cost_per_token_flex")
row: Final = _PRICES[model]
usage: Final = _plain_usage()
cost: Final = completion_cost(
completion_response=ModelResponse(model=model, choices=[], usage=usage),
model=model,
optional_params={"metadata": {"completion_window": "flex"}},
)
expected: Final = usage.prompt_tokens * cast(
float, row["input_cost_per_token_flex"]
) + usage.completion_tokens * cast(float, row["output_cost_per_token_flex"])
assert cost == pytest.approx(expected, rel=1e-9)
assert cost != pytest.approx(
usage.prompt_tokens * cast(float, row["input_cost_per_token"])
+ usage.completion_tokens * cast(float, row["output_cost_per_token"]),
rel=1e-9,
)
def test_completion_cost_ignores_completion_window_for_openai(_local_model_cost_map):
model: Final = _latest_row("gpt-", "openai", "input_cost_per_token_flex", "output_cost_per_token_flex")
row: Final = _PRICES[model]
usage: Final = _plain_usage()
cost: Final = completion_cost(
completion_response=ModelResponse(model=model, choices=[], usage=usage),
model=f"openai/{model}",
optional_params={"metadata": {"completion_window": "flex"}},
)
expected: Final = usage.prompt_tokens * cast(float, row["input_cost_per_token"]) + usage.completion_tokens * cast(
float, row["output_cost_per_token"]
)
assert cost == pytest.approx(expected, rel=1e-9)
assert cost != pytest.approx(
usage.prompt_tokens * cast(float, row["input_cost_per_token_flex"])
+ usage.completion_tokens * cast(float, row["output_cost_per_token_flex"]),
rel=1e-9,
)

View file

@ -232,6 +232,27 @@ def test_dated_variants_carry_base_alias_service_tier_pricing(prices: dict):
)
def test_balanced_tier_keys_live_only_on_sail_rows_next_to_their_base_key(prices: dict):
balanced_keys = tuple(
(name, tier_key)
for name, entry in prices.items()
if isinstance(entry, dict)
for tier_key in entry
if tier_key.endswith("_balanced")
)
misplaced = tuple(
(name, tier_key, entry.get("litellm_provider"), tier_anchor(tier_key) in entry)
for name, tier_key in balanced_keys
for entry in [prices[name]]
if not name.startswith("sail/") or entry.get("litellm_provider") != "sail" or tier_anchor(tier_key) not in entry
)
assert misplaced == (), (
"(model, balanced key, litellm_provider, has base key) for balanced tier pricing that is "
f"off a sail/ row or missing its base key: {misplaced}"
)
assert balanced_keys != ()
OPENAI_REASONING_FAMILY_MARKERS = ("codex", "deep-research", "chat-latest")

View file

@ -7,6 +7,7 @@ import json
import logging
import os
import queue
import re
import threading
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import Future, ThreadPoolExecutor
@ -238,6 +239,44 @@ def test_get_model_info_prefers_exact_dated_key_over_stripped(
assert info["key"] == expected_key
_BALANCED_MODEL_INFO_FIELDS: Final = (
"input_cost_per_token_balanced",
"output_cost_per_token_balanced",
"cache_read_input_token_cost_balanced",
"cache_creation_input_token_cost_balanced",
"output_cost_per_reasoning_token_balanced",
)
def _latest_chat_row(prefix: str, custom_llm_provider: str, *required_keys: str) -> str:
return max(
(
name
for name, row in litellm.model_cost.items()
if name.startswith(prefix)
and row.get("litellm_provider") == custom_llm_provider
and row.get("mode") == "chat"
and all(key in row for key in required_keys)
),
key=lambda name: tuple(float(part) for part in re.findall(r"\d+(?:\.\d+)?", name)),
)
def test_get_model_info_reports_balanced_tier_prices_only_where_the_row_has_them(local_model_cost_map: None) -> None:
gpt_row: Final = _latest_chat_row("gpt-", "openai")
sail_row: Final = _latest_chat_row("sail/", "sail", "input_cost_per_token_balanced")
assert all(field not in litellm.model_cost[gpt_row] for field in _BALANCED_MODEL_INFO_FIELDS)
gpt_info: Final = litellm.get_model_info(model=gpt_row, custom_llm_provider="openai")
sail_info: Final = litellm.get_model_info(model=sail_row, custom_llm_provider="sail")
assert tuple(gpt_info[field] for field in _BALANCED_MODEL_INFO_FIELDS) == (None,) * len(_BALANCED_MODEL_INFO_FIELDS)
assert tuple(sail_info[field] for field in _BALANCED_MODEL_INFO_FIELDS) == tuple(
litellm.model_cost[sail_row].get(field) for field in _BALANCED_MODEL_INFO_FIELDS
)
assert sail_info["input_cost_per_token_balanced"] != sail_info["input_cost_per_token"]
def test_get_model_info_internal_failure_is_not_reported_as_unmapped() -> None:
with patch("litellm.utils._get_potential_model_names", side_effect=RuntimeError("malformed metadata")):
with pytest.raises(Exception, match="This model isn't mapped yet") as exc_info: