From 31ed2482b5f746bcc76c7f03407fa947b80bc356 Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 16:19:31 +0000 Subject: [PATCH] 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 --- MUTATIONS.md | 590 ++++++++++++++++++ .../test_prometheus_service_tier_label.py | 35 ++ .../llm_cost_calc/test_llm_cost_calc_utils.py | 38 ++ .../chat/test_base_chat_transformation.py | 21 + .../base_llm/responses/test_transformation.py | 16 + .../custom_httpx/test_llm_http_handler.py | 187 ++++++ .../llms/openai_like/test_json_providers.py | 61 +- .../llms/openai_like/test_sail_provider.py | 7 +- .../test_responses_api_request_body.py | 67 ++ tests/test_litellm/test_constants.py | 8 + tests/test_litellm/test_cost_calculator.py | 143 +++++ .../test_litellm/test_model_prices_schema.py | 21 + tests/test_litellm/test_utils.py | 39 ++ 13 files changed, 1228 insertions(+), 5 deletions(-) create mode 100644 MUTATIONS.md create mode 100644 tests/test_litellm/llms/base_llm/chat/test_base_chat_transformation.py diff --git a/MUTATIONS.md b/MUTATIONS.md new file mode 100644 index 00000000000..0027eb990ca --- /dev/null +++ b/MUTATIONS.md @@ -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 -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 -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` diff --git a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py b/tests/test_litellm/integrations/test_prometheus_service_tier_label.py index b2212c4ff41..db4c993db1e 100644 --- a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py +++ b/tests/test_litellm/integrations/test_prometheus_service_tier_label.py @@ -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.""" 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 903afb0b3dc..b2f83714017 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 @@ -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", [ diff --git a/tests/test_litellm/llms/base_llm/chat/test_base_chat_transformation.py b/tests/test_litellm/llms/base_llm/chat/test_base_chat_transformation.py new file mode 100644 index 00000000000..b52be79b415 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/chat/test_base_chat_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/base_llm/responses/test_transformation.py b/tests/test_litellm/llms/base_llm/responses/test_transformation.py index c6142685661..a6efa157e28 100644 --- a/tests/test_litellm/llms/base_llm/responses/test_transformation.py +++ b/tests/test_litellm/llms/base_llm/responses/test_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 114dc221cdc..7cfdb05291e 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 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 4696becfc1b..ec916275819 100644 --- a/tests/test_litellm/llms/openai_like/test_json_providers.py +++ b/tests/test_litellm/llms/openai_like/test_json_providers.py @@ -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 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 35a9d9704ee..3ee4e24d015 100644 --- a/tests/test_litellm/llms/openai_like/test_sail_provider.py +++ b/tests/test_litellm/llms/openai_like/test_sail_provider.py @@ -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): diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index 98e74955c6f..5e693791212 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -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 diff --git a/tests/test_litellm/test_constants.py b/tests/test_litellm/test_constants.py index 12e473f68a4..fb7b534b795 100644 --- a/tests/test_litellm/test_constants.py +++ b/tests/test_litellm/test_constants.py @@ -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 diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index f7d6cfaf079..d797780652c 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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, + ) diff --git a/tests/test_litellm/test_model_prices_schema.py b/tests/test_litellm/test_model_prices_schema.py index 7e589096325..d936f6c1a45 100644 --- a/tests/test_litellm/test_model_prices_schema.py +++ b/tests/test_litellm/test_model_prices_schema.py @@ -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") diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 7aab68bf859..525785b0a82 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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: