fix: use azure provider for container cost tracking

Override transform_container_create_response in AzureOpenAIContainerConfig
to pass provider="azure" instead of inheriting the hardcoded "openai" value.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
zooneon 2026-03-27 23:30:17 +09:00
parent 29b83cb016
commit 15783555f7
2 changed files with 60 additions and 0 deletions

View file

@ -1,8 +1,15 @@
from typing import Optional
import httpx
from litellm.constants import AZURE_DEFAULT_CONTAINERS_API_VERSION
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
from litellm.types.containers.main import ContainerObject
from litellm.types.router import GenericLiteLLMParams
@ -46,3 +53,33 @@ class AzureOpenAIContainerConfig(OpenAIContainerConfig):
return BaseAzureLLM._base_validate_azure_environment(
headers=headers, litellm_params=litellm_params
)
def transform_container_create_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ContainerObject:
"""Transform the Azure container creation response.
Overrides OpenAI's method to use provider="azure" for cost tracking.
"""
response_data = raw_response.json()
container_obj = ContainerObject(**response_data) # type: ignore[arg-type]
container_cost = StandardBuiltInToolCostTracking.get_cost_for_code_interpreter(
sessions=1,
provider="azure",
)
if (
not hasattr(container_obj, "_hidden_params")
or container_obj._hidden_params is None
):
container_obj._hidden_params = {}
if "additional_headers" not in container_obj._hidden_params:
container_obj._hidden_params["additional_headers"] = {}
container_obj._hidden_params["additional_headers"][
"llm_provider-x-litellm-response-cost"
] = container_cost
return container_obj

View file

@ -156,6 +156,29 @@ class TestAzureContainerTransformations:
assert container.id == "cntr_azure_123"
assert container.name == "Azure Container"
def test_transform_container_create_response_uses_azure_provider_for_cost(self):
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {
"id": "cntr_azure_cost",
"object": "container",
"created_at": 1747857508,
"status": "running",
"expires_after": {"anchor": "last_active_at", "minutes": 20},
"last_active_at": 1747857508,
"name": "Azure Cost Test",
}
from unittest.mock import patch
with patch(
"litellm.llms.azure.containers.transformation.StandardBuiltInToolCostTracking.get_cost_for_code_interpreter"
) as mock_cost:
mock_cost.return_value = 0.03
self.config.transform_container_create_response(
raw_response=mock_response,
logging_obj=self.logging_obj,
)
mock_cost.assert_called_once_with(sessions=1, provider="azure")
def test_transform_container_list_response(self):
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {