fix(cost): price Azure PTU spillover requests at standard token rates

Azure PTU deployments carry zeroed per-token pricing because the reservation
is billed flat by the hour. When Azure spills a request onto pay-as-you-go
capacity it returns x-ms-is-spilled-over: true, and that traffic was still
priced at zero. The response cost calculator now detects the spillover header
on the result's hidden params or the logged provider response headers and
skips the zeroed custom pricing only for genuine PTU deployments while the
feature flag is on. Azure sync streaming now also records response headers on
the logging object, matching the async paths.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-17 05:24:05 +00:00
parent 5ef40a630b
commit a57483d1c8
6 changed files with 293 additions and 5 deletions

View file

@ -90,6 +90,7 @@ from litellm.litellm_core_utils.logging_utils import (
truncate_base64_in_messages_async,
)
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
from litellm.litellm_core_utils.ptu_pricing import is_spilled_over_ptu_request
from litellm.litellm_core_utils.redact_messages import (
redact_message_input_output_from_custom_logger,
redact_message_input_output_from_logging,
@ -1746,8 +1747,14 @@ class Logging(LiteLLMLoggingBaseClass):
if transformed_result is not None:
result = transformed_result
result_hidden_params: Final = getattr(result, "_hidden_params", None) or MappingProxyType({})
result_additional_headers: Final = (
result_hidden_params.get("additional_headers")
if isinstance(result_hidden_params, dict)
else getattr(result_hidden_params, "additional_headers", None)
)
if isinstance(result, (BaseModel, HttpxBinaryResponseContent)) and hasattr(result, "_hidden_params"):
hidden_params: Final = getattr(result, "_hidden_params", {})
hidden_params: Final = result_hidden_params
if (
"response_cost" in hidden_params and hidden_params["response_cost"] is not None
): # use cost if already calculated
@ -1762,8 +1769,17 @@ class Logging(LiteLLMLoggingBaseClass):
router_model_id = self.get_router_model_id()
## RESPONSE COST ##
custom_pricing: Final = use_custom_pricing_for_model(
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None)
spilled_over: Final = is_spilled_over_ptu_request(
model_info=_deployment_model_info(self.litellm_params if hasattr(self, "litellm_params") else None),
response_headers=self.model_call_details.get("response_headers"),
additional_headers=result_additional_headers,
)
custom_pricing: Final = (
False
if spilled_over
else use_custom_pricing_for_model(
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None)
)
)
prompt = self._prompt_for_cost_calculation()
@ -5257,6 +5273,18 @@ def _get_custom_logger_settings_from_proxy_server(callback_name: str) -> dict:
return {}
def _deployment_model_info(litellm_params: dict | None) -> Mapping[str, object]:
"""The router-stamped deployment model_info from whichever metadata field carries it."""
if litellm_params is None:
return MappingProxyType({})
for metadata_key in ("metadata", "litellm_metadata"):
if not isinstance(metadata := litellm_params.get(metadata_key), Mapping):
continue
if model_info := metadata.get("model_info"):
return model_info
return MappingProxyType({})
def use_custom_pricing_for_model(litellm_params: dict | None) -> bool:
"""
Check if the model uses custom pricing

View file

@ -17,6 +17,7 @@ from litellm.types.router import ModelInfo
from litellm.types.utils import CustomPricingLiteLLMParams, MirroredPricingParams
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
AZURE_SPILLOVER_HEADER: Final = "x-ms-is-spilled-over"
def is_ptu_cost_attribution_enabled() -> bool:
@ -235,3 +236,22 @@ def zeroed_ptu_pricing(
),
}
)
def is_spilled_over_ptu_request(
model_info: Mapping[str, object],
response_headers: Mapping[str, object] | None,
additional_headers: Mapping[str, object] | None,
) -> bool:
"""Whether Azure served this request from pay-as-you-go capacity, so the zeroed PTU rates must not apply."""
if ptu_terms(model_info) is None:
return False
if not is_ptu_cost_attribution_enabled():
return False
for headers, key in (
(response_headers, AZURE_SPILLOVER_HEADER),
(additional_headers, f"llm_provider-{AZURE_SPILLOVER_HEADER}"),
):
if headers is not None and str(headers.get(key)).lower() == "true":
return True
return False

View file

@ -561,6 +561,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
headers, response = self.make_sync_azure_openai_chat_completion_request(
azure_client=azure_client, data=data, timeout=timeout
)
logging_obj.model_call_details["response_headers"] = headers
streamwrapper: Final = CustomStreamWrapper(
completion_stream=response,
model=model,

View file

@ -7229,3 +7229,155 @@ def test_add_dynamic_callback_registers_once_per_list_without_touching_the_calle
assert logging_obj.dynamic_async_failure_callbacks == [callback]
assert LitellmLogging._with_dynamic_callback(None, callback) == [callback]
assert LitellmLogging._with_dynamic_callback((callback,), callback) == [callback]
class TestAzurePTUSpilloverCost:
"""Azure PTU deployments price tokens at zero because the reservation is billed flat.
A request Azure spills onto pay-as-you-go capacity must bill per token instead, so
the zeroed custom pricing has to be skipped when the provider reports spillover.
"""
ROUTER_MODEL_ID: Final = "ptu-spill-router-model-id"
SERVED_MODEL: Final = "azure/spill-served-model-ptu"
PTU_MODEL_INFO: Final = {
"id": ROUTER_MODEL_ID,
"team_id": "team-1",
"ptu_count": 100,
"cost_per_ptu_per_hour": 1.0,
"ptu_effective_from": "2026-01-01",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
}
EXPECTED_SPILL_COST: Final = 100 * 2e-6 + 50 * 8e-6
@staticmethod
def _register_models() -> None:
litellm.register_model(
model_cost={
TestAzurePTUSpilloverCost.ROUTER_MODEL_ID: {
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"litellm_provider": "azure",
"mode": "chat",
},
TestAzurePTUSpilloverCost.SERVED_MODEL: {
"input_cost_per_token": 2e-6,
"output_cost_per_token": 8e-6,
"litellm_provider": "azure",
"mode": "chat",
},
}
)
@staticmethod
def _unregister_models() -> None:
litellm.model_cost.pop(TestAzurePTUSpilloverCost.ROUTER_MODEL_ID, None)
litellm.model_cost.pop(TestAzurePTUSpilloverCost.SERVED_MODEL, None)
def _logging_obj(self, model_info: dict, *, flag: str, litellm_rate: float, monkeypatch) -> LitellmLogging:
monkeypatch.setenv("LITELLM_ENABLE_PTU_COST_ATTRIBUTION", flag)
obj = LitellmLogging(
model=self.SERVED_MODEL,
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="completion",
start_time=time.time(),
litellm_call_id="ptu-spill-1",
function_id="f",
)
obj.update_environment_variables(
model=self.SERVED_MODEL,
user="",
optional_params={},
litellm_params={
"api_base": "",
"metadata": {"model_info": model_info},
"input_cost_per_token": litellm_rate,
"output_cost_per_token": litellm_rate,
},
custom_llm_provider="azure",
)
return obj
@staticmethod
def _response() -> ModelResponse:
from litellm.types.utils import Usage
return ModelResponse(
id="chatcmpl-spill-1",
created=1234567890,
model="spill-served-model-ptu",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "ok"},
"finish_reason": "stop",
}
],
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
)
def test_spillover_via_response_additional_headers_bills_per_token(self, monkeypatch) -> None:
self._register_models()
try:
obj = self._logging_obj(dict(self.PTU_MODEL_INFO), flag="True", litellm_rate=0.0, monkeypatch=monkeypatch)
response = self._response()
response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"}
assert obj._response_cost_calculator(result=response) == pytest.approx(self.EXPECTED_SPILL_COST)
finally:
self._unregister_models()
def test_spillover_via_streaming_response_headers_bills_per_token(self, monkeypatch) -> None:
self._register_models()
try:
obj = self._logging_obj(dict(self.PTU_MODEL_INFO), flag="True", litellm_rate=0.0, monkeypatch=monkeypatch)
obj.model_call_details["response_headers"] = {
"x-ms-is-spilled-over": "true",
"x-ms-spillover-from-deployment": "ptu-dep",
}
assert obj._response_cost_calculator(result=self._response()) == pytest.approx(self.EXPECTED_SPILL_COST)
finally:
self._unregister_models()
def test_non_spilled_ptu_request_stays_zero_priced(self, monkeypatch) -> None:
self._register_models()
try:
obj = self._logging_obj(dict(self.PTU_MODEL_INFO), flag="True", litellm_rate=0.0, monkeypatch=monkeypatch)
assert obj._response_cost_calculator(result=self._response()) == 0.0
finally:
self._unregister_models()
def test_spillover_header_without_the_flag_stays_zero_priced(self, monkeypatch) -> None:
self._register_models()
try:
obj = self._logging_obj(dict(self.PTU_MODEL_INFO), flag="", litellm_rate=0.0, monkeypatch=monkeypatch)
response = self._response()
response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"}
assert obj._response_cost_calculator(result=response) == 0.0
finally:
self._unregister_models()
def test_spillover_header_does_not_touch_non_ptu_custom_pricing(self, monkeypatch) -> None:
self._register_models()
custom_model_id: Final = "non-ptu-custom-router-model-id"
litellm.model_cost[custom_model_id] = {
"input_cost_per_token": 1e-6,
"output_cost_per_token": 1e-6,
"litellm_provider": "azure",
"mode": "chat",
}
try:
model_info: Final = {"id": custom_model_id, "input_cost_per_token": 1e-6}
obj = self._logging_obj(model_info, flag="True", litellm_rate=1e-6, monkeypatch=monkeypatch)
response = self._response()
response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"}
assert obj._response_cost_calculator(result=response) == pytest.approx(150 * 1e-6)
finally:
litellm.model_cost.pop(custom_model_id, None)
self._unregister_models()

View file

@ -7,13 +7,14 @@ from unittest.mock import patch
import pytest
from litellm.litellm_core_utils.ptu_pricing import (
ptu_config_error,
ptu_identity_error,
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
PTU_ZEROED_PRICING_FIELDS,
PTU_ZEROED_TABLE_FIELDS,
SEARCH_CONTEXT_SIZES,
is_spilled_over_ptu_request,
ptu_config_error,
ptu_identity_error,
ptu_terms,
zeroed_ptu_pricing,
)
@ -294,3 +295,35 @@ def test_an_empty_id_is_no_id():
assert error is not None
assert error.startswith("model_info.id is required")
def test_the_spillover_header_marks_the_request_as_pay_as_you_go():
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True"}, clear=False):
assert (
is_spilled_over_ptu_request(
model_info=_VALID,
response_headers={"x-ms-is-spilled-over": "True"},
additional_headers=None,
)
is True
)
def test_no_spillover_marker_keeps_the_zeroed_ptu_rates():
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True"}, clear=False):
assert (
is_spilled_over_ptu_request(
model_info=_VALID,
response_headers={"x-ms-is-spilled-over": "false"},
additional_headers=None,
)
is False
)
assert (
is_spilled_over_ptu_request(
model_info=_VALID,
response_headers=None,
additional_headers={"llm_provider-x-ms-is-spilled-over": "absent"},
)
is False
)

View file

@ -0,0 +1,54 @@
"""Tests for litellm/llms/azure/azure.py AzureChatCompletion handler behaviour."""
import time
from typing import Final
from openai import AzureOpenAI
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure.azure import AzureChatCompletion
class _FakeRawResponse:
headers: Final = {"x-ms-is-spilled-over": "true"}
def parse(self):
return iter(())
class _FakeRawCompletions:
def create(self, **kwargs):
return _FakeRawResponse()
def test_sync_streaming_stamps_response_headers_on_the_logging_obj() -> None:
"""Sync streaming must mirror async_streaming and record the provider response
headers on model_call_details, or downstream consumers (spillover-aware cost
calculation) cannot see them."""
client = AzureOpenAI(api_key="fake", api_version="2024-02-01", azure_endpoint="https://fake.openai.azure.com")
client.chat.completions.with_raw_response = _FakeRawCompletions()
logging_obj = LiteLLMLoggingObj(
model="azure/gpt-4o-spill-test",
messages=[{"role": "user", "content": "Hi"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="spill-sync-1",
function_id="f",
)
AzureChatCompletion().streaming(
logging_obj=logging_obj,
api_base="https://fake.openai.azure.com",
api_key="fake",
api_version="2024-02-01",
dynamic_params=False,
data={"messages": [{"role": "user", "content": "Hi"}], "stream": True},
model="gpt-4o-spill-test",
timeout=30.0,
max_retries=0,
client=client,
)
assert logging_obj.model_call_details["response_headers"] == {"x-ms-is-spilled-over": "true"}