fix(azure_ai): map max completion tokens to max tokens

Signed-off-by: Vic Wen (Manpower Services Taiwan Co Ltd) <a-vicwen@microsoft.com>
This commit is contained in:
Vic Wen (Manpower Services Taiwan Co Ltd) 2026-09-20 22:22:45 +08:00
parent 5fa1257b7c
commit 9a1aff1bf5
2 changed files with 29 additions and 3 deletions

View file

@ -1,6 +1,7 @@
import copy
import enum
import re
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
import httpx
@ -96,20 +97,33 @@ 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:
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.

View file

@ -1,4 +1,5 @@
import json
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
@ -146,6 +147,17 @@ 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_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},