mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(sail): bill by the completion window that reaches the wire and keep it through extra_body merges
This commit is contained in:
parent
17c0dd4a0a
commit
45b9843b3f
10 changed files with 325 additions and 49 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue