From 15783555f78291b4f9511270837c4a2307b895aa Mon Sep 17 00:00:00 2001 From: zooneon Date: Fri, 27 Mar 2026 23:30:17 +0900 Subject: [PATCH] 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 --- .../llms/azure/containers/transformation.py | 37 +++++++++++++++++++ .../test_azure_container_transformation.py | 23 ++++++++++++ 2 files changed, 60 insertions(+) diff --git a/litellm/llms/azure/containers/transformation.py b/litellm/llms/azure/containers/transformation.py index 8f02c6a57e8..b1e7632ef0c 100644 --- a/litellm/llms/azure/containers/transformation.py +++ b/litellm/llms/azure/containers/transformation.py @@ -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 diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index 042b4a0c767..c0260542dec 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -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 = {