diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index de2ca4762f1..55b4670d14e 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -13,7 +13,7 @@ from litellm.router import Router from litellm.router_utils.fallback_event_handlers import get_fallback_model_group from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params from litellm.types.utils import LlmProviders -from litellm.utils import get_valid_models +from litellm.utils import ProviderConfigManager, get_valid_models _CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields) @@ -35,6 +35,16 @@ def _check_wildcard_routing(model: str) -> bool: return False +def _provider_supports_model_discovery(provider: str) -> bool: + if provider in litellm.models_by_provider: + return True + try: + llm_provider: Final = LlmProviders(provider) + except ValueError: + return False + return ProviderConfigManager.get_provider_model_info(model=None, provider=llm_provider) is not None + + def get_provider_models(provider: str, litellm_params: LiteLLM_Params | None = None) -> list[str] | None: """ Returns the list of known models by provider @@ -42,7 +52,7 @@ def get_provider_models(provider: str, litellm_params: LiteLLM_Params | None = N if provider == "*": return get_valid_models(litellm_params=litellm_params) - if provider in litellm.models_by_provider: + if _provider_supports_model_discovery(provider): provider_models: Final = get_valid_models(custom_llm_provider=provider, litellm_params=litellm_params) return provider_models return None diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 36bfc4c5dd3..f580df2dad6 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -1,4 +1,4 @@ -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -432,6 +432,47 @@ def test_wildcard_credential_hydration_preserves_deployment_params( } +def test_get_known_models_from_wildcard_hosted_vllm_uses_provider_endpoint(): + import litellm + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + response = MagicMock() + response.json.return_value = { + "data": [ + {"id": "meta-llama/Llama-3.1-8B-Instruct"}, + {"id": "qwen2.5"}, + ] + } + original_check_provider_endpoint = litellm.check_provider_endpoint + try: + litellm.check_provider_endpoint = True # test-quality-ok: required to exercise provider endpoint discovery + with ( + patch("litellm.module_level_client.get", return_value=response), # test-quality-ok: required HTTP boundary + patch( # test-quality-ok: VLLM model listing has no injectable API key seam + "litellm.llms.vllm.common_utils.VLLMModelInfo.get_api_key", + return_value="test-key", + ), + ): + result = get_known_models_from_wildcard( + "hosted_vllm/*", + LiteLLM_Params(model="hosted_vllm/*", api_base="http://localhost:8000/v1"), + ) + finally: + litellm.check_provider_endpoint = original_check_provider_endpoint # test-quality-ok: restore test global + + assert result == [ + "hosted_vllm/meta-llama/Llama-3.1-8B-Instruct", + "hosted_vllm/qwen2.5", + ] + + +def test_get_known_models_from_wildcard_unknown_provider_returns_empty(): + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + + assert get_known_models_from_wildcard("not_a_real_provider/*") == [] + + def test_wildcard_custom_prefix_does_not_stack_provider_prefix(monkeypatch): """Regression test for #30358.