diff --git a/litellm/constants.py b/litellm/constants.py index ba4ae397e09..91c3ff1ec90 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1011,9 +1011,12 @@ openai_compatible_providers: Final[list] = [ "meta", # Meta Model API (Muse Spark) - JSON-configured provider "cognition", "scx-ai", + "sail", ] -OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS: Final = frozenset({"openai"} | frozenset(openai_compatible_providers)) +OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS: Final = frozenset(("openai",)) | ( + frozenset(openai_compatible_providers) - frozenset(("sail",)) +) openai_text_completion_compatible_providers: Final[list] = [ # providers that support `/v1/completions` "together_ai", diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 95a0ba003d2..82ab7d6ba16 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -966,13 +966,12 @@ def _extract_service_tier(source: object) -> str | None: def _completion_window_value(metadata: object) -> str | None: - """Return ``metadata["completion_window"]`` only when it names a billable tier - ("flex" or "balanced"). "asap" is the provider default and bills at standard - pricing, so it returns None.""" + """Return ``metadata["completion_window"]`` only when it names a completion + window ("asap", "flex" or "balanced").""" if not isinstance(metadata, dict): return None window: Final[object] = metadata.get("completion_window") - if isinstance(window, str) and window in (ServiceTier.FLEX.value, ServiceTier.BALANCED.value): + if isinstance(window, str) and window in ("asap", ServiceTier.FLEX.value, ServiceTier.BALANCED.value): return window return None @@ -986,7 +985,7 @@ def _provider_bills_by_completion_window(custom_llm_provider: str | None) -> boo return provider is not None and provider.special_handling.get("service_tier_as_completion_window") is True -def _service_tier_from_completion_window(optional_params: dict[str, object]) -> str | None: +def _service_tier_from_completion_window(optional_params: Mapping[str, object]) -> str | None: """Read ``metadata.completion_window`` from ``extra_body`` or a top-level ``metadata`` param (the two shapes callers use to pick a provider completion window directly).""" extra_body: Final = optional_params.get("extra_body") @@ -994,6 +993,19 @@ def _service_tier_from_completion_window(optional_params: dict[str, object]) -> return _completion_window_value(extra_metadata) or _completion_window_value(optional_params.get("metadata")) +def _service_tier_billed_by_completion_window( + service_tier: str | None, + optional_params: Mapping[str, object] | None, + custom_llm_provider: str | None, +) -> str | None: + if optional_params is None or not _provider_bills_by_completion_window(custom_llm_provider): + return service_tier + window: Final = _service_tier_from_completion_window(optional_params) + if window is None: + return service_tier + return None if window == "asap" else window + + def get_usage_object( completion_response: object, ) -> Usage | None: @@ -1403,13 +1415,6 @@ def completion_cost( ) rerank_billed_units: RerankBilledUnits | None = None - # Providers that bill by completion window: an explicit window on the request wins - # over service_tier, matching what the provider actually sees on the wire - if optional_params is not None and _provider_bills_by_completion_window(custom_llm_provider): - window_tier: Final = _service_tier_from_completion_window(optional_params) - if window_tier is not None: - service_tier = window_tier - # Extract service_tier from optional_params if not provided directly if service_tier is None and optional_params is not None: service_tier = _normalize_service_tier(optional_params.get("service_tier")) @@ -1556,6 +1561,11 @@ def completion_cost( "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, + ) if CostCalculatorUtils._call_type_has_image_response(call_type) and isinstance( completion_response, ImageResponse ): diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 7decf1b4186..94ab6f01297 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -4,7 +4,7 @@ Common base config for all LLM providers import types from abc import ABC, abstractmethod -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final, Union import httpx @@ -311,6 +311,13 @@ class BaseConfig(ABC): ) -> dict: pass + def merge_extra_body( + self, + request: dict[str, object], # mutable-ok: wire request body is a plain dict + extra_body: Mapping[str, object] | None, + ) -> dict[str, object]: # mutable-ok: wire request body is a plain dict + return {**request, **extra_body} if extra_body else request # mutable-ok: wire request body is a plain dict + async def async_transform_request( self, model: str, diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 3834d19ec2b..e9844e1787b 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -1,5 +1,6 @@ import types from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast import httpx @@ -146,6 +147,13 @@ class BaseResponsesAPIConfig(ABC): headers=headers, ) + def merge_extra_body( + self, + request: dict[str, object], # mutable-ok: wire request body is a plain dict + extra_body: Mapping[str, object] | None, + ) -> dict[str, object]: # mutable-ok: wire request body is a plain dict + return {**request, **extra_body} if extra_body else request # mutable-ok: wire request body is a plain dict + @abstractmethod def transform_response_api_response( self, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 052978c2680..fb2a5f6552f 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -618,7 +618,7 @@ class BaseLLMHTTPHandler: def sign_and_log( transformed: dict[str, object], # mutable-ok: async_completion takes dict ) -> tuple[dict[str, object], dict[str, object], bytes | None]: # mutable-ok: async_completion takes dict - data: Final = {**transformed, **extra_body} if extra_body is not None else transformed + data: Final = provider_config.merge_extra_body(transformed, extra_body) signed: Final = cast( # cast-ok: sign_request is declared as a bare dict "tuple[dict[str, object], bytes | None]", provider_config.sign_request( @@ -2702,8 +2702,7 @@ class BaseLLMHTTPHandler: ) data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data) - if extra_body: - data.update(extra_body) + data = responses_api_provider_config.merge_extra_body(data, extra_body) stream = bool(stream or data.get("stream")) # Preserve the OpenAI-style request context (not sent to the provider) for streaming @@ -2890,8 +2889,7 @@ class BaseLLMHTTPHandler: ) data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data) - if extra_body: - data.update(extra_body) + data = responses_api_provider_config.merge_extra_body(data, extra_body) stream = bool(stream or data.get("stream")) # Preserve the OpenAI-style request context (not sent to the provider) for streaming diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index 01b381b26f7..989f633268a 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -28,22 +28,72 @@ _SERVICE_TIER_TO_COMPLETION_WINDOW: Final[Mapping[str, str]] = MappingProxyType( ) -def _apply_service_tier_as_completion_window(body: Mapping[str, object]) -> dict[str, object]: +def _completion_window_str(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _apply_service_tier_as_completion_window( + body: Mapping[str, object], +) -> dict[str, object]: # mutable-ok: transform_request returns a plain dict service_tier: Final = body.get("service_tier") - window: Final[str | None] = ( + mapped_window: Final[str | None] = ( _SERVICE_TIER_TO_COMPLETION_WINDOW.get(service_tier.lower()) if isinstance(service_tier, str) else None ) raw_metadata: Final = body.get("metadata") metadata: Final[Mapping[str, object]] = raw_metadata if isinstance(raw_metadata, dict) else MappingProxyType({}) + raw_extra_body: Final = body.get("extra_body") + extra_body: Final[Mapping[str, object]] = ( + raw_extra_body if isinstance(raw_extra_body, dict) else MappingProxyType({}) + ) + raw_extra_metadata: Final = extra_body.get("metadata") + extra_metadata: Final[Mapping[str, object]] = ( + raw_extra_metadata if isinstance(raw_extra_metadata, dict) else MappingProxyType({}) + ) + caller_window: Final[str | None] = _completion_window_str( + extra_metadata.get("completion_window") + ) or _completion_window_str(metadata.get("completion_window")) + window: Final[str | None] = caller_window or mapped_window new_body: Final[dict[str, object]] = { # mutable-ok: transform_request returns a plain dict key: value for key, value in body.items() if key != "service_tier" } - if window is None or "completion_window" in metadata: + if window is None or (caller_window is None and window == "asap" and body.get("background") is True): return new_body - return { # mutable-ok: transform_request returns a plain dict + merged: Final[dict[str, object]] = { # mutable-ok: transform_request returns a plain dict **new_body, "metadata": {**metadata, "completion_window": window}, # mutable-ok: transform_request returns a plain dict } + if isinstance(raw_extra_metadata, dict) and "completion_window" not in extra_metadata: + return { # mutable-ok: transform_request returns a plain dict + **merged, + "extra_body": { # mutable-ok: SDK merges extra_body into the wire body + **extra_body, + "metadata": { # mutable-ok: SDK merges extra_body into the wire body + **extra_metadata, + "completion_window": window, + }, + }, + } + return merged + + +def _merge_extra_body_keeping_metadata( + request: dict[str, object], # mutable-ok: wire request body is a plain dict + extra_body: Mapping[str, object] | None, +) -> dict[str, object]: # mutable-ok: wire request body is a plain dict + """Shallow merge where ``metadata`` merges one level deep so a caller window + inside ``extra_body.metadata`` wins over a mapped one but cannot wipe it out + by replacing the whole metadata dict.""" + if not extra_body: + return request + request_metadata: Final = request.get("metadata") + extra_metadata: Final = extra_body.get("metadata") + if not isinstance(request_metadata, dict) or not isinstance(extra_metadata, dict): + return {**request, **extra_body} # mutable-ok: request body sent over the wire + return { # mutable-ok: request body sent over the wire + **request, + **extra_body, + "metadata": {**request_metadata, **extra_metadata}, # mutable-ok: request body sent over the wire + } def _service_tier_as_completion_window_enabled(provider: SimpleProviderConfig) -> bool: @@ -125,11 +175,11 @@ def create_config_class(provider: SimpleProviderConfig): def transform_request( self, model: str, - messages: list[AllMessageValues], - optional_params: dict[str, object], - litellm_params: dict[str, object], - headers: dict[str, object], - ) -> dict[str, object]: + messages: list[AllMessageValues], # mutable-ok: matches base signature + optional_params: dict[str, object], # mutable-ok: matches base signature + litellm_params: dict[str, object], # mutable-ok: matches base signature + headers: dict[str, object], # mutable-ok: matches base signature + ) -> dict[str, object]: # mutable-ok: matches base signature body: Final = super().transform_request( model=model, messages=messages, @@ -141,6 +191,15 @@ def create_config_class(provider: SimpleProviderConfig): return _apply_service_tier_as_completion_window(body) return body + def merge_extra_body( + self, + request: dict[str, object], # mutable-ok: wire request body is a plain dict + extra_body: Mapping[str, object] | None, + ) -> dict[str, object]: # mutable-ok: wire request body is a plain dict + if _service_tier_as_completion_window_enabled(provider): + return _merge_extra_body_keeping_metadata(request, extra_body) + return super().merge_extra_body(request, extra_body) + def get_supported_openai_params(self, model: str) -> list: """Get supported OpenAI params, excluding tool-related params for models that don't support function calling.""" @@ -277,10 +336,10 @@ def create_responses_config_class(provider: SimpleProviderConfig): self, model: str, input: str | ResponseInputParam, - response_api_optional_request_params: dict[str, object], + response_api_optional_request_params: dict[str, object], # mutable-ok: matches base signature litellm_params: GenericLiteLLMParams, - headers: dict[str, object], - ) -> dict[str, object]: + headers: dict[str, object], # mutable-ok: matches base signature + ) -> dict[str, object]: # mutable-ok: matches base signature if provider.special_handling.get("force_store_false"): response_api_optional_request_params["store"] = False body: Final = super().transform_responses_api_request( @@ -294,5 +353,14 @@ def create_responses_config_class(provider: SimpleProviderConfig): return _apply_service_tier_as_completion_window(body) return body + def merge_extra_body( + self, + request: dict[str, object], # mutable-ok: wire request body is a plain dict + extra_body: Mapping[str, object] | None, + ) -> dict[str, object]: # mutable-ok: wire request body is a plain dict + if _service_tier_as_completion_window_enabled(provider): + return _merge_extra_body_keeping_metadata(request, extra_body) + return super().merge_extra_body(request, extra_body) + _responses_config_cache[provider.slug] = JSONProviderResponsesConfig return JSONProviderResponsesConfig diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 884ea7e217d..126dfa90685 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -1398,7 +1398,12 @@ def responses( kwargs=kwargs, model=model, user=user, - optional_params=dict(responses_api_request_params), + optional_params={ # mutable-ok: update_from_kwargs stores a plain dict + **responses_api_request_params, + **( + {"extra_body": extra_body} if extra_body else {} + ), # mutable-ok: update_from_kwargs stores a plain dict + }, litellm_params={ **responses_api_request_params, "aresponses": _is_async, @@ -2232,7 +2237,12 @@ def compact_responses( litellm_logging_obj.update_from_kwargs( kwargs=local_vars, model=model, - optional_params=dict(responses_api_request_params), + optional_params={ # mutable-ok: update_from_kwargs stores a plain dict + **responses_api_request_params, + **( + {"extra_body": extra_body} if extra_body else {} + ), # mutable-ok: update_from_kwargs stores a plain dict + }, litellm_params={ **responses_api_request_params, "litellm_call_id": litellm_call_id, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 957ce7ea197..5846266ed8d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -267,7 +267,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): max_input_tokens: Required[int | None] max_output_tokens: Required[int | None] input_cost_per_token: Required[float | None] - input_cost_per_token_balanced: float | None # balanced service tier pricing + input_cost_per_token_balanced: ReadOnly[float | None] # balanced service tier pricing input_cost_per_token_flex: float | None # OpenAI flex service tier pricing input_cost_per_token_priority: float | None # OpenAI priority service tier pricing input_cost_per_token_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing @@ -278,14 +278,14 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_creation_input_token_cost_above_272k_tokens_flex: float | None cache_creation_input_token_cost_above_1hr: float | None cache_creation_input_token_cost_flex: float | None # OpenAI flex service tier pricing - cache_creation_input_token_cost_balanced: float | None # balanced service tier pricing + cache_creation_input_token_cost_balanced: ReadOnly[float | None] # balanced service tier pricing cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing cache_creation_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing cache_read_input_token_cost: float | None cache_read_input_audio_token_cost: ReadOnly[float | None] cache_read_input_image_token_cost: ReadOnly[float | None] cache_read_input_token_cost_flex: float | None # OpenAI flex service tier pricing - cache_read_input_token_cost_balanced: float | None # balanced service tier pricing + cache_read_input_token_cost_balanced: ReadOnly[float | None] # balanced service tier pricing cache_read_input_token_cost_priority: float | None # OpenAI priority service tier pricing cache_read_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing cache_read_input_token_cost_above_200k_tokens: float | None @@ -327,7 +327,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): output_cost_per_token_batches: float | None output_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None] output_cost_per_token: Required[float | None] - output_cost_per_token_balanced: float | None # balanced service tier pricing + output_cost_per_token_balanced: ReadOnly[float | None] # balanced service tier pricing output_cost_per_token_flex: float | None # OpenAI flex service tier pricing output_cost_per_token_priority: float | None # OpenAI priority service tier pricing output_cost_per_token_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing @@ -357,7 +357,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): output_vector_size: int | None output_cost_per_reasoning_token: float | None output_cost_per_reasoning_token_flex: float | None - output_cost_per_reasoning_token_balanced: float | None + output_cost_per_reasoning_token_balanced: ReadOnly[float | None] output_cost_per_reasoning_token_priority: float | None output_cost_per_video_per_second: float | None # only for vertex ai models output_cost_per_audio_per_second: float | None # only for vertex ai models diff --git a/tests/code_coverage_tests/check_sail_registry_against_vendor.py b/tests/code_coverage_tests/check_sail_registry_against_vendor.py index e91266b6d44..5555cf42b20 100644 --- a/tests/code_coverage_tests/check_sail_registry_against_vendor.py +++ b/tests/code_coverage_tests/check_sail_registry_against_vendor.py @@ -9,6 +9,7 @@ Exits 1 with a per-model diff table on any mismatch. """ import json +import math import os import re import sys @@ -50,7 +51,9 @@ def vendor_model_ids(client: httpx.Client, api_key: str) -> tuple[set[str], list return {m["id"] for m in data if isinstance(m, dict) and isinstance(m.get("id"), str)}, [] -def vendor_context_windows(client: httpx.Client, api_key: str, model_ids: list[str]) -> tuple[dict[str, int], list[str]]: +def vendor_context_windows( + client: httpx.Client, api_key: str, model_ids: list[str] +) -> tuple[dict[str, int], list[str]]: windows: dict[str, int] = {} unverified: list[str] = [] for model_id in model_ids: @@ -79,16 +82,16 @@ def vendor_prices(client: httpx.Client) -> dict[str, dict[str, dict[str, float]] text: Final = client.get(PRICING_URL, timeout=60).text labels: Final = ARIA_LABEL_RE.findall(text) prices: dict[str, dict[str, dict[str, float]]] = {} - last_copy: str | None = None + copy_labels: Final = tuple( + (index, label[len(COPY_PREFIX) :]) for index, label in enumerate(labels) if label.startswith(COPY_PREFIX) + ) for index, label in enumerate(labels): match: Final = PRICING_LABEL_RE.search(f'aria-label="{label}"') if match is None: continue - next_label: Final = labels[index + 1] if index + 1 < len(labels) else "" - model_id: Final = next_label[len(COPY_PREFIX):] if next_label.startswith(COPY_PREFIX) else last_copy - if model_id is None: + if not copy_labels: continue - last_copy = model_id + model_id: Final = min(copy_labels, key=lambda entry: abs(entry[0] - index))[1] tiers: Final = prices.setdefault(model_id, {}) tiers[TIER_TO_SUFFIX[match.group(2)]] = { "input_cost_per_token": float(match.group(3)) / 1e6, @@ -122,6 +125,10 @@ def main() -> int: diffs.append(f"{model_id}: context window cost map={expected} vendor={window}") prices: Final = vendor_prices(client) + if not prices: + diffs.append(f"no tier prices parsed from {PRICING_URL}") + for model_id in sorted(set(rows) - set(prices)): + diffs.append(f"{model_id}: sail/ row in cost map but no tier prices on {PRICING_URL}") for model_id, tiers in sorted(prices.items()): row: Final = rows.get(model_id) if row is None: @@ -130,10 +137,8 @@ def main() -> int: for field in ("input_cost_per_token", "cache_read_input_token_cost", "output_cost_per_token"): ours: Final = row.get(f"{field}{suffix}") theirs: Final = vendor_fields[field] - if ours != theirs: - diffs.append( - f"{model_id}: {field}{suffix or ' (asap)'} cost map={ours} vendor={theirs}" - ) + if ours is None or not math.isclose(ours, theirs, rel_tol=1e-9): + diffs.append(f"{model_id}: {field}{suffix or ' (asap)'} cost map={ours} vendor={theirs}") if unverified: print("unverified models (no context window error or 503):", ", ".join(unverified)) 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 008097a417f..35a9d9704ee 100644 --- a/tests/test_litellm/llms/openai_like/test_sail_provider.py +++ b/tests/test_litellm/llms/openai_like/test_sail_provider.py @@ -365,7 +365,7 @@ def _sail_chat_body(optional_params: dict) -> dict: return config.transform_request( model="zai-org/GLM-5.3", messages=_MESSAGES, - optional_params=dict(optional_params), + optional_params=optional_params, litellm_params={}, headers={}, ) @@ -601,6 +601,173 @@ class TestSailTierPricing: assert cost == pytest.approx(expected) +class TestSailWireAndBillingConsistency: + @staticmethod + def _expected_cost(suffix: str, prompt_tokens: int = 2, completion_tokens: int = 2) -> float: + rates = litellm.model_cost[MODEL] + return ( + prompt_tokens * rates[f"input_cost_per_token{suffix}"] + + completion_tokens * rates[f"output_cost_per_token{suffix}"] + ) + + @pytest.mark.parametrize("service_tier", ["flex", "balanced"]) + @pytest.mark.respx() + def test_asap_window_wins_over_tier_and_bills_base(self, respx_mock: respx.Router, service_tier: str): + respx_mock.post(SAIL_CHAT_COMPLETIONS).respond(json=_chat_completion_payload()) + + response = litellm.completion( + model=MODEL, + messages=[{"role": "user", "content": "hi"}], + service_tier=service_tier, + extra_body={"metadata": {"completion_window": "asap"}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"completion_window": "asap"} + assert "service_tier" not in body + assert response._hidden_params["response_cost"] == pytest.approx(self._expected_cost("")) + + @pytest.mark.respx() + def test_tier_window_survives_extra_body_metadata_merge(self, respx_mock: respx.Router): + respx_mock.post(SAIL_CHAT_COMPLETIONS).respond(json=_chat_completion_payload()) + + response = litellm.completion( + model=MODEL, + messages=[{"role": "user", "content": "hi"}], + service_tier="flex", + extra_body={"metadata": {"trace_id": "x"}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"trace_id": "x", "completion_window": "flex"} + assert response._hidden_params["response_cost"] == pytest.approx(self._expected_cost("_flex")) + + @pytest.mark.asyncio + @pytest.mark.respx() + async def test_aresponses_tier_window_survives_extra_body_metadata_merge(self, respx_mock: respx.Router): + respx_mock.post(SAIL_RESPONSES).respond(json=_responses_payload()) + + response = await litellm.aresponses( + model=MODEL, + input="hi", + service_tier="flex", + extra_body={"metadata": {"trace_id": "x"}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"trace_id": "x", "completion_window": "flex"} + assert response._hidden_params["response_cost"] == pytest.approx(self._expected_cost("_flex")) + + @pytest.mark.asyncio + @pytest.mark.respx() + async def test_aresponses_extra_body_window_bills_at_window(self, respx_mock: respx.Router): + respx_mock.post(SAIL_RESPONSES).respond(json=_responses_payload()) + + response = await litellm.aresponses( + model=MODEL, + input="hi", + extra_body={"metadata": {"completion_window": "balanced"}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"completion_window": "balanced"} + assert response._hidden_params["response_cost"] == pytest.approx(self._expected_cost("_balanced")) + + @pytest.mark.asyncio + @pytest.mark.respx() + async def test_aresponses_extra_body_window_overrides_tier(self, respx_mock: respx.Router): + respx_mock.post(SAIL_RESPONSES).respond(json=_responses_payload()) + + response = await litellm.aresponses( + model=MODEL, + input="hi", + service_tier="flex", + extra_body={"metadata": {"completion_window": "balanced"}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"completion_window": "balanced"} + assert response._hidden_params["response_cost"] == pytest.approx(self._expected_cost("_balanced")) + + def test_window_override_applies_when_provider_inferred_from_model(self): + rates = litellm.model_cost[MODEL] + prompt_tokens, completion_tokens = 1000, 200 + + cost = litellm.completion_cost( + completion_response=_sail_completion_response(prompt_tokens, completion_tokens), + model=MODEL, + custom_llm_provider=None, + optional_params={ + "service_tier": "balanced", + "extra_body": {"metadata": {"completion_window": "flex"}}, + }, + ) + + expected = ( + prompt_tokens * rates["input_cost_per_token_flex"] + completion_tokens * rates["output_cost_per_token_flex"] + ) + assert cost == pytest.approx(expected) + + @pytest.mark.parametrize( + "optional_params,expected_window", + [ + ({"service_tier": "default", "background": True}, None), + ({"service_tier": "flex", "background": True}, "flex"), + ], + ids=["background_asap_suppressed", "background_flex_kept"], + ) + def test_background_suppresses_asap_window(self, optional_params: dict, expected_window: str | None): + body = _sail_chat_body(optional_params) + assert "service_tier" not in body + if expected_window is None: + assert "completion_window" not in body.get("metadata", {}) + else: + assert body["metadata"]["completion_window"] == expected_window + + @pytest.mark.respx() + def test_vendor_kwarg_folds_into_request_body(self, respx_mock: respx.Router): + respx_mock.post(SAIL_CHAT_COMPLETIONS).respond(json=_chat_completion_payload()) + + litellm.completion( + model=MODEL, + messages=[{"role": "user", "content": "hi"}], + reasoning_budget=128, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["reasoning_budget"] == 128 + + @pytest.mark.respx() + def test_non_sail_extra_body_metadata_stays_shallow(self, respx_mock: respx.Router): + respx_mock.post("https://api.openai.com/v1/chat/completions").respond( + json={ + "id": "chatcmpl-openai", + "object": "chat.completion", + "created": 1234567890, + "model": "gpt-4.1-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4}, + } + ) + + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + api_key="sk-test", + metadata={"b": 2}, + extra_body={"metadata": {"a": 1}}, + ) + + body = json.loads(respx_mock.calls[0].request.content) + assert body["metadata"] == {"a": 1} + + _TIER_COST_BASES = ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost")