fix(model_checks.py): handle custom values in wildcard model name (e.g. genai/test/*) (#13116)

Fixes https://github.com/BerriAI/litellm/issues/13078
This commit is contained in:
Krish Dholakia 2025-07-29 21:42:15 -07:00 • committed by GitHub
parent 7e5bc8af28
commit a34206f67e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 43 additions and 25 deletions

View file

@ -705,8 +705,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
if api_key is None:
api_key = get_secret_str("OPENAI_API_KEY")
# Strip api_base to just the base URL (scheme + host + port)
parsed_url = httpx.URL(api_base)
base_url = f"{parsed_url.scheme}://{parsed_url.host}"
if parsed_url.port:
base_url += f":{parsed_url.port}"
response = litellm.module_level_client.get(
url=f"{api_base}/v1/models",
url=f"{base_url}/v1/models",
headers={"Authorization": f"Bearer {api_key}"},
)

View file

@ -1,5 +1,9 @@
model_list:
- model_name: bedrock-claude-3.7-sonnet
- model_name: genai/test/*
litellm_params:
model: bedrock/us.anthropic.claude-3-7-sonnet-20250219-v1:0
model: openai/*
api_base: https://api.openai.com
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
check_provider_endpoint: true

View file

@ -75,16 +75,14 @@ async def get_mcp_server_ids(
if prisma_client is None:
return []
if user_api_key_dict.object_permission_id is None:
return []
# Make a direct SQL query to get just the mcp_servers
try:
result = await prisma_client.db.litellm_objectpermissiontable.find_unique(
where={"object_permission_id": user_api_key_dict.object_permission_id},
where={"object_permission_id": user_api_key_dict.object_permission_id},
)
if result and result.mcp_servers:
return result.mcp_servers
@ -229,19 +227,30 @@ def get_known_models_from_wildcard(
provider = wildcard_provider_prefix
# get all known provider models
wildcard_models = get_provider_models(
provider=provider, litellm_params=litellm_params
)
if wildcard_models is None:
return []
if wildcard_suffix != "*":
## CHECK IF PARTIAL FILTER e.g. `gemini-*`
model_prefix = wildcard_suffix.replace("*", "")
filtered_wildcard_models = [
wc_model
for wc_model in wildcard_models
if wc_model.startswith(model_prefix)
]
wildcard_models = filtered_wildcard_models
is_partial_filter = any(
wc_model.startswith(model_prefix) for wc_model in wildcard_models
)
if is_partial_filter:
filtered_wildcard_models = [
wc_model
for wc_model in wildcard_models
if wc_model.startswith(model_prefix)
]
wildcard_models = filtered_wildcard_models
else:
# add model prefix to wildcard models
wildcard_models = [f"{model_prefix}{model}" for model in wildcard_models]
suffix_appended_wildcard_models = []
for model in wildcard_models:
@ -298,18 +307,18 @@ def get_all_fallbacks(
) -> List[str]:
"""
Get all fallbacks for a given model from the router's fallback configuration.
Args:
model: The model name to get fallbacks for
llm_router: The LiteLLM router instance
fallback_type: Type of fallback ("general", "context_window", "content_policy")
Returns:
List of fallback model names. Empty list if no fallbacks found.
"""
if llm_router is None:
return []
# Get the appropriate fallback list based on type
fallbacks_config: list = []
if fallback_type == "general":
@ -321,20 +330,19 @@ def get_all_fallbacks(
else:
verbose_proxy_logger.warning(f"Unknown fallback_type: {fallback_type}")
return []
if not fallbacks_config:
return []
try:
# Use existing function to get fallback model group
fallback_model_group, _ = get_fallback_model_group(
fallbacks=fallbacks_config,
model_group=model
fallbacks=fallbacks_config, model_group=model
)
if fallback_model_group is None:
return []
return fallback_model_group
except Exception as e:
verbose_proxy_logger.error(f"Error getting fallbacks for model {model}: {e}")

View file

@ -5283,7 +5283,7 @@ def validate_environment( # noqa: PLR0915
keys_in_environment = True
else:
missing_keys.append("GOOGLE_API_KEY")
missing_keys.append("GEMINI_API_KEY")
missing_keys.append("GEMINI_API_KEY")
elif custom_llm_provider == "groq":
if "GROQ_API_KEY" in os.environ:
keys_in_environment = True
@ -5640,7 +5640,7 @@ def _calculate_retry_after(
min_timeout: int = 0,
) -> Union[float, int]:
retry_after = _get_retry_after_from_exception_header(response_headers)
# Add some jitter (default JITTER is 0.75 - so upto 0.75s)
jitter = JITTER * random.random()
@ -5654,8 +5654,8 @@ def _calculate_retry_after(
# Make sure sleep_seconds is boxed between min_timeout and MAX_RETRY_DELAY
sleep_seconds = max(sleep_seconds, min_timeout)
sleep_seconds = min(sleep_seconds, MAX_RETRY_DELAY)
sleep_seconds = min(sleep_seconds, MAX_RETRY_DELAY)
return sleep_seconds + jitter