fix(sail): bill by the completion window that reaches the wire and keep it through extra_body merges

This commit is contained in:
shrey kharbanda 2026-09-24 07:29:18 +00:00
parent 17c0dd4a0a
commit 45b9843b3f
10 changed files with 325 additions and 49 deletions

View file

@ -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",

View file

@ -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
):

View file

@ -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,

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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))

View file

@ -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")