Removed stop param from unsupported azure models (#15229)

* Removed stop param from unsupported model

* Use better handling for stop method

* Use better handling for stop method
This commit is contained in:
Sameer Kankute 2025-10-07 08:26:18 +05:30 • committed by GitHub
parent 6e538033ed
commit 8d7f39798c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 54 additions and 50 deletions

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
from litellm.llms.openai.openai import OpenAIConfig
from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, ProviderField
@ -35,9 +36,24 @@ class AzureAIStudioConfig(OpenAIConfig):
for param in supported_params:
if param != "tool_choice":
filtered_supported_params.append(param)
return filtered_supported_params
supported_params = filtered_supported_params
# Filter out unsupported parameters for specific models
if not self._supports_stop_reason(model):
supported_params = [param for param in supported_params if param != "stop"]
return supported_params
def _supports_stop_reason(self, model: str) -> bool:
"""
Check if the model supports stop tokens.
"""
if "grok" in model:
# Reuse Xai method for Grok model
xai_config = XAIChatConfig()
return xai_config._supports_stop_reason(model)
return True
def validate_environment(
self,
headers: dict,
@ -53,9 +69,7 @@ class AzureAIStudioConfig(OpenAIConfig):
else:
headers["Authorization"] = f"Bearer {api_key}"
headers["Content-Type"] = (
"application/json" # tell Azure AI Studio to expect JSON
)
headers["Content-Type"] = "application/json" # tell Azure AI Studio to expect JSON
return headers
@ -65,10 +79,7 @@ class AzureAIStudioConfig(OpenAIConfig):
"""
parsed_url = urlparse(api_base)
host = parsed_url.hostname
if host and (
host.endswith(".services.ai.azure.com")
or host.endswith(".openai.azure.com")
):
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
return True
return False
@ -115,13 +126,9 @@ class AzureAIStudioConfig(OpenAIConfig):
# Add the path to the base URL
if "services.ai.azure.com" in api_base:
new_url = _add_path_to_api_base(
api_base=api_base, ending_path="/models/chat/completions"
)
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions")
else:
new_url = _add_path_to_api_base(
api_base=api_base, ending_path="/chat/completions"
)
new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions")
# Use the new query_params dictionary
final_url = httpx.URL(new_url).copy_with(params=query_params)
@ -191,11 +198,7 @@ class AzureAIStudioConfig(OpenAIConfig):
dynamic_api_key = api_key or get_secret_str("AZURE_AI_API_KEY")
if self._is_azure_openai_model(model=model, api_base=api_base):
verbose_logger.debug(
"Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(
model
)
)
verbose_logger.debug("Model={} is Azure OpenAI model. Setting custom_llm_provider='azure'.".format(model))
custom_llm_provider = "azure"
return api_base, dynamic_api_key, custom_llm_provider
@ -211,9 +214,7 @@ class AzureAIStudioConfig(OpenAIConfig):
if extra_body and isinstance(extra_body, dict):
optional_params.update(extra_body)
optional_params.pop("max_retries", None)
return super().transform_request(
model, messages, optional_params, litellm_params, headers
)
return super().transform_request(model, messages, optional_params, litellm_params, headers)
def transform_response(
self,
@ -252,47 +253,30 @@ class AzureAIStudioConfig(OpenAIConfig):
if should_drop_params and "Extra inputs are not permitted" in error_text:
return True
elif (
"unknown field: parameter index is not a valid field" in error_text
): # remove index from tool calls
elif "unknown field: parameter index is not a valid field" in error_text: # remove index from tool calls
return True
elif (
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
in error_text
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in error_text
): # remove extra-parameters from tool calls
return True
return super().should_retry_llm_api_inside_llm_translation_on_http_error(
e=e, litellm_params=litellm_params
)
return super().should_retry_llm_api_inside_llm_translation_on_http_error(e=e, litellm_params=litellm_params)
@property
def max_retry_on_unprocessable_entity_error(self) -> int:
return 2
def transform_request_on_unprocessable_entity_error(
self, e: httpx.HTTPStatusError, request_data: dict
) -> dict:
def transform_request_on_unprocessable_entity_error(self, e: httpx.HTTPStatusError, request_data: dict) -> dict:
_messages = cast(Optional[List[AllMessageValues]], request_data.get("messages"))
if (
"unknown field: parameter index is not a valid field" in e.response.text
and _messages is not None
):
if "unknown field: parameter index is not a valid field" in e.response.text and _messages is not None:
litellm.remove_index_from_tool_calls(
messages=_messages,
)
elif (
AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value
in e.response.text
):
request_data = self._drop_extra_params_from_request_data(
request_data, e.response.text
)
elif AzureFoundryErrorStrings.SET_EXTRA_PARAMETERS_TO_PASS_THROUGH.value in e.response.text:
request_data = self._drop_extra_params_from_request_data(request_data, e.response.text)
data = drop_params_from_unprocessable_entity_error(e=e, data=request_data)
return data
def _drop_extra_params_from_request_data(
self, request_data: dict, error_text: str
) -> dict:
def _drop_extra_params_from_request_data(self, request_data: dict, error_text: str) -> dict:
params_to_drop = self._extract_params_to_drop_from_error_text(error_text)
if params_to_drop:
for param in params_to_drop:
@ -300,9 +284,7 @@ class AzureAIStudioConfig(OpenAIConfig):
request_data.pop(param, None)
return request_data
def _extract_params_to_drop_from_error_text(
self, error_text: str
) -> Optional[List[str]]:
def _extract_params_to_drop_from_error_text(self, error_text: str) -> Optional[List[str]]:
"""
Error text looks like this"
"Extra parameters ['stream_options', 'extra-parameters'] are not allowed when extra-parameters is not set or set to be 'error'.

View file

@ -43,3 +43,25 @@ def test_azure_ai_validate_environment():
litellm_params={},
)
assert headers["Content-Type"] == "application/json"
def test_azure_ai_grok_stop_parameter_handling():
"""
Test that Grok models properly handle stop parameter filtering in Azure AI Studio.
"""
config = AzureAIStudioConfig()
# Test Grok model detection
assert config._supports_stop_reason("grok-4-fast") == False
assert config._supports_stop_reason("grok-4") == False
assert config._supports_stop_reason("grok-3-mini") == False
assert config._supports_stop_reason("grok-code-fast") == False
assert config._supports_stop_reason("gpt-4") == True
# Test supported parameters for Grok models
grok_params = config.get_supported_openai_params("grok-4-fast")
assert "stop" not in grok_params, "Grok models should not support stop parameter"
# Test supported parameters for non-Grok models
gpt_params = config.get_supported_openai_params("gpt-4")
assert "stop" in gpt_params, "GPT models should support stop parameter"