diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index ce470f04aca..868d02ee1e7 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -14,6 +14,8 @@ from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI if TYPE_CHECKING: from aiohttp import ClientSession +import inspect + import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -22,6 +24,14 @@ from litellm.llms.custom_httpx.http_handler import ( get_ssl_configuration, ) +def _get_client_init_params(cls: type) -> List[str]: + """Extract __init__ parameter names (excluding 'self') from a class.""" + return [p for p in inspect.signature(cls.__init__).parameters if p != "self"] + + +_OPENAI_INIT_PARAMS: List[str] = _get_client_init_params(OpenAI) +_AZURE_OPENAI_INIT_PARAMS: List[str] = _get_client_init_params(AzureOpenAI) + class OpenAIError(BaseLLMException): def __init__( @@ -183,18 +193,10 @@ class BaseOpenAILLM: client_type: Literal["openai", "azure"] ) -> List[str]: """Returns a list of fields that are used to initialize the OpenAI client""" - import inspect - - from openai import AzureOpenAI, OpenAI - if client_type == "openai": - signature = inspect.signature(OpenAI.__init__) + return _OPENAI_INIT_PARAMS else: - signature = inspect.signature(AzureOpenAI.__init__) - - # Extract parameter names, excluding 'self' - param_names = [param for param in signature.parameters if param != "self"] - return param_names + return _AZURE_OPENAI_INIT_PARAMS @staticmethod def _get_async_http_client( @@ -230,3 +232,5 @@ class BaseOpenAILLM: verify=ssl_config, follow_redirects=True, ) + + diff --git a/tests/test_litellm/llms/openai/test_openai_common_utils.py b/tests/test_litellm/llms/openai/test_openai_common_utils.py index 469005d103f..f2740be642d 100644 --- a/tests/test_litellm/llms/openai/test_openai_common_utils.py +++ b/tests/test_litellm/llms/openai/test_openai_common_utils.py @@ -129,3 +129,38 @@ async def test_openai_client_reuse(function_name, is_async, args): # Verify we tried to get from cache 10 times (once per request) assert mock_get_cache.call_count == 10, "Should check cache for each request" + + +def test_precomputed_init_params_match_inspect_signature(): + """ + Verify that the pre-computed _OPENAI_INIT_PARAMS and _AZURE_OPENAI_INIT_PARAMS + match what inspect.signature() returns. If the OpenAI SDK changes its __init__ + params, this test will fail — signaling the constants need updating. + """ + import inspect + + from openai import AzureOpenAI, OpenAI + + from litellm.llms.openai.common_utils import ( + _AZURE_OPENAI_INIT_PARAMS, + _OPENAI_INIT_PARAMS, + ) + + expected_openai = [ + p for p in inspect.signature(OpenAI.__init__).parameters if p != "self" + ] + expected_azure = [ + p for p in inspect.signature(AzureOpenAI.__init__).parameters if p != "self" + ] + + assert _OPENAI_INIT_PARAMS == expected_openai + assert _AZURE_OPENAI_INIT_PARAMS == expected_azure + + +@pytest.mark.parametrize("client_type", ["openai", "azure"]) +def test_get_openai_client_initialization_param_fields(client_type): + """Verify the method returns the correct pre-computed params for each client type.""" + result = BaseOpenAILLM.get_openai_client_initialization_param_fields(client_type) + assert isinstance(result, list) + assert len(result) > 0 + assert "self" not in result