mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
5ef40a630b
commit
a57483d1c8
6 changed files with 293 additions and 5 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
54
tests/test_litellm/llms/azure/test_azure.py
Normal file
54
tests/test_litellm/llms/azure/test_azure.py
Normal 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"}
|
||||
Loading…
Add table
Reference in a new issue