mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
29b83cb016
commit
15783555f7
2 changed files with 60 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue