diff --git a/litellm/llms/openai_like/chat/transformation.py b/litellm/llms/openai_like/chat/transformation.py index 895d5a99971..96a6c3172ea 100644 --- a/litellm/llms/openai_like/chat/transformation.py +++ b/litellm/llms/openai_like/chat/transformation.py @@ -166,3 +166,11 @@ class OpenAILikeChatConfig(OpenAIGPTConfig): mapped_params.pop("max_completion_tokens", None) return mapped_params + + +class CustomOpenAIChatConfig(OpenAILikeChatConfig): + def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> List[str]: + if api_base is None: + raise ValueError("api_base must be set to discover models for the custom_openai provider") + models = super().get_models(api_key=api_key, api_base=api_base) + return [f"custom_openai/{model}" for model in models] diff --git a/litellm/utils.py b/litellm/utils.py index 4bbaba83d2f..02a14c7ad56 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8359,7 +8359,11 @@ class ProviderConfigManager: elif LlmProviders.OPENAI == provider: return litellm.OpenAIGPTConfig() elif LlmProviders.CUSTOM_OPENAI == provider: - return litellm.OpenAIGPTConfig() + from litellm.llms.openai_like.chat.transformation import ( + CustomOpenAIChatConfig, + ) + + return CustomOpenAIChatConfig() elif LlmProviders.GEMINI == provider: return litellm.GeminiModelInfo() elif LlmProviders.VERTEX_AI == provider: diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index be5d1fc324e..71e743f2dee 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -1683,7 +1683,14 @@ def test_get_valid_models_custom_openai(monkeypatch): api_key="sk-1234", api_base="https://my-openai-compatible-endpoint/v1", ) - assert "my-custom-model" in valid_models + assert "custom_openai/my-custom-model" in valid_models + + +def test_custom_openai_get_models_requires_api_base(): + from litellm.llms.openai_like.chat.transformation import CustomOpenAIChatConfig + + with pytest.raises(ValueError): + CustomOpenAIChatConfig().get_models(api_base=None) def test_get_valid_models_fireworks_ai(monkeypatch):