diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 779a86629e2..0ab386d544e 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -1,6 +1,7 @@ import copy import enum import re +from types import MappingProxyType from typing import TYPE_CHECKING, Final, cast import httpx @@ -27,7 +28,7 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse, ProviderField -from litellm.utils import _add_path_to_api_base, supports_tool_choice +from litellm.utils import _add_path_to_api_base, supports_reasoning, supports_tool_choice if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -96,20 +97,35 @@ class AzureAIStudioConfig(OpenAIConfig): model: str, drop_params: bool, ) -> dict[str, object]: # mutable-ok: OpenAIConfig.map_openai_params signature - if not azureAIGPT5Config.is_model_gpt_5_model(model): - return super().map_openai_params( + if azureAIGPT5Config.is_model_gpt_5_model(model): + return azureAIGPT5Config.map_openai_params( non_default_params=non_default_params, optional_params=optional_params, model=model, drop_params=drop_params, ) - return azureAIGPT5Config.map_openai_params( + super().map_openai_params( non_default_params=non_default_params, optional_params=optional_params, model=model, drop_params=drop_params, ) + if "max_completion_tokens" not in optional_params or supports_reasoning( + model=model, custom_llm_provider="azure_ai" + ): + return optional_params + + max_tokens: Final = optional_params["max_completion_tokens"] + token_limit_params: Final = frozenset(("max_tokens", "max_completion_tokens")) + params_without_token_limits: Final = MappingProxyType( + {key: value for key, value in optional_params.items() if key not in token_limit_params} + ) + return { # mutable-ok: BaseConfig mapping contract requires a request dict + **params_without_token_limits, + "max_tokens": max_tokens, + } + def _supports_stop_reason(self, model: str) -> bool: """ Check if the model supports stop tokens. diff --git a/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py index 47609261a25..28568958569 100644 --- a/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -1,4 +1,5 @@ import json +from typing import Final from unittest.mock import MagicMock, patch import pytest @@ -146,6 +147,51 @@ def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url)) +def test_azure_ai_maps_max_completion_tokens_to_max_tokens(): + mapped_params: Final = litellm.get_optional_params( + model="mistral-large-3", + custom_llm_provider="azure_ai", + max_completion_tokens=256, + ) + + assert mapped_params["max_tokens"] == 256 + assert "max_completion_tokens" not in mapped_params + + +def test_azure_ai_keeps_max_completion_tokens_for_gpt_5(): + mapped_params: Final = litellm.get_optional_params( + model="gpt-5", + custom_llm_provider="azure_ai", + max_completion_tokens=256, + ) + + assert mapped_params["max_completion_tokens"] == 256 + assert "max_tokens" not in mapped_params + + +def test_azure_ai_keeps_max_completion_tokens_for_reasoning_models(_local_model_cost_map: None) -> None: + mapped_params: Final = litellm.get_optional_params( + model="o3", + custom_llm_provider="azure_ai", + max_completion_tokens=256, + ) + + assert mapped_params["max_completion_tokens"] == 256 + assert "max_tokens" not in mapped_params + + +def test_azure_ai_keeps_params_without_max_completion_tokens(): + mapped_params: Final = litellm.get_optional_params( + model="mistral-large-3", + custom_llm_provider="azure_ai", + temperature=0.2, + ) + + assert mapped_params["temperature"] == 0.2 + assert "max_tokens" not in mapped_params + assert "max_completion_tokens" not in mapped_params + + def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none(_local_model_cost_map): optional_params = AzureAIStudioConfig().map_openai_params( non_default_params={"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9},