diff --git a/litellm/main.py b/litellm/main.py index 6102fe3ccce..70f55125507 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1256,6 +1256,7 @@ def completion( # type: ignore # noqa: PLR0915 additional_drop_params=kwargs.get("additional_drop_params"), remove_sensitive_keys=True, add_provider_specific_params=True, + provider_config=provider_config, ) if litellm.add_function_to_prompt and optional_params.get( diff --git a/litellm/utils.py b/litellm/utils.py index aa3c00735ec..40a1438b3fe 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3088,6 +3088,7 @@ def pre_process_non_default_params( model: str, remove_sensitive_keys: bool = False, add_provider_specific_params: bool = False, + provider_config: Optional[BaseConfig] = None, ) -> dict: """ Pre-process non-default params to a standardized format @@ -3103,14 +3104,6 @@ def pre_process_non_default_params( additional_endpoint_specific_params=["messages"], ) - provider_config: Optional[BaseConfig] = None - if custom_llm_provider is not None and custom_llm_provider in [ - provider.value for provider in LlmProviders - ]: - provider_config = ProviderConfigManager.get_provider_chat_config( - model=model, provider=LlmProviders(custom_llm_provider) - ) - if "response_format" in non_default_params: if provider_config is not None: non_default_params[ diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index bd39fbfc9c4..dba40e2214a 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -957,7 +957,12 @@ def test_get_model_info_shows_supports_computer_use(): def test_pre_process_non_default_params(model, custom_llm_provider): from pydantic import BaseModel - from litellm.utils import pre_process_non_default_params + from litellm.utils import ProviderConfigManager, pre_process_non_default_params + + provider_config = ProviderConfigManager.get_provider_chat_config( + model=model, + provider=LlmProviders(custom_llm_provider) + ) class ResponseFormat(BaseModel): x: str @@ -974,6 +979,7 @@ def test_pre_process_non_default_params(model, custom_llm_provider): special_params=special_params, custom_llm_provider=custom_llm_provider, additional_drop_params=None, + provider_config=provider_config, ) print(processed_non_default_params) assert processed_non_default_params == {