Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/auth/model_checks.py: 47%
224 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1# What is this?
2## Common checks for /v1/models and `/model/info`
3import copy
4from collections.abc import Sequence
5from typing import Any, Final
7import litellm
8from litellm._logging import verbose_proxy_logger
9from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
10from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
11from litellm.repositories.object_permission_repository import ObjectPermissionRepository
12from litellm.router import Router
13from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
14from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params
15from litellm.types.utils import LlmProviders
16from litellm.utils import get_valid_models
18_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
21_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
24def _check_wildcard_routing(model: str) -> bool:
25 """
26 Returns True if a model is a provider wildcard.
28 eg:
29 - anthropic/*
30 - openai/*
31 - *
32 """
33 if "*" in model:
34 return True
35 return False
38def get_provider_models(provider: str, litellm_params: LiteLLM_Params | None = None) -> list[str] | None:
39 """
40 Returns the list of known models by provider
41 """
42 if provider == "*":
43 return get_valid_models(litellm_params=litellm_params)
45 if provider in litellm.models_by_provider:
46 provider_models: Final = get_valid_models(custom_llm_provider=provider, litellm_params=litellm_params)
47 return provider_models
48 return None
51def _get_models_from_access_groups(
52 model_access_groups: dict[str, list[str]],
53 all_models: list[str],
54 include_model_access_groups: bool | None = False,
55 proxy_model_list: Sequence[str] | None = None,
56) -> list[str]:
57 # a grant naming both a deployed model and an access group means both at runtime
58 # (_check_model_access_helper unions them), so listings must keep the literal too
59 deployed_model_names: Final = frozenset(proxy_model_list or ())
60 kept_models: Final = [
61 model
62 for model in all_models
63 if model not in model_access_groups or include_model_access_groups or model in deployed_model_names
64 ]
65 member_models: Final = [
66 member for model in all_models if model in model_access_groups for member in model_access_groups[model]
67 ]
68 return kept_models + member_models
71async def get_mcp_server_ids(
72 user_api_key_dict: UserAPIKeyAuth,
73) -> list[str]:
74 """
75 Returns the list of MCP server ids for a given key by querying the object_permission table
76 """
77 from litellm.proxy.proxy_server import prisma_client
79 if prisma_client is None:
80 return []
82 if user_api_key_dict.object_permission_id is None:
83 return []
85 # Make a direct SQL query to get just the mcp_servers
86 try:
87 result: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
88 where={"object_permission_id": user_api_key_dict.object_permission_id},
89 )
90 if result and result.mcp_servers:
91 return result.mcp_servers
92 return []
93 except Exception:
94 return []
97def get_key_models(
98 user_api_key_dict: UserAPIKeyAuth,
99 proxy_model_list: list[str],
100 model_access_groups: dict[str, list[str]],
101 include_model_access_groups: bool | None = False,
102 only_model_access_groups: bool | None = False,
103) -> list[str]:
104 """
105 Returns:
106 - List of model name strings
107 - Empty list if no models set
108 - If model_access_groups is provided, only return models that are in the access groups
109 - If include_model_access_groups is True, it includes the 'keys' of the model_access_groups
110 in the response - {"beta-models": ["gpt-4", "claude-v1"]} -> returns 'beta-models'
111 """
112 all_models: list[str] = []
113 if len(user_api_key_dict.models) > 0: 113 ↛ 114line 113 didn't jump to line 114 because the condition on line 113 was never true
114 all_models = list(user_api_key_dict.models) # copy to avoid mutating cached objects
115 if SpecialModelNames.all_team_models.value in all_models:
116 all_models = list(user_api_key_dict.team_models)
117 if SpecialModelNames.all_team_models.value in all_models:
118 all_models = [model for model in all_models if model != SpecialModelNames.all_team_models.value]
119 all_models.extend(proxy_model_list)
120 if include_model_access_groups:
121 all_models.extend(model_access_groups.keys())
122 if SpecialModelNames.all_proxy_models.value in all_models:
123 all_models = list(proxy_model_list) # copy to avoid mutating caller's list
124 if include_model_access_groups:
125 all_models.extend(model_access_groups.keys())
127 all_models = _get_models_from_access_groups(
128 model_access_groups=model_access_groups,
129 all_models=all_models,
130 include_model_access_groups=include_model_access_groups,
131 proxy_model_list=proxy_model_list,
132 )
134 # deduplicate while preserving order
135 all_models = list(dict.fromkeys(all_models))
137 verbose_proxy_logger.debug("ALL KEY MODELS - %s", len(all_models))
138 return all_models
141def get_team_models(
142 team_models: list[str],
143 proxy_model_list: list[str],
144 model_access_groups: dict[str, list[str]],
145 include_model_access_groups: bool | None = False,
146) -> list[str]:
147 """
148 Returns:
149 - List of model name strings
150 - Empty list if no models set
151 - If model_access_groups is provided, only return models that are in the access groups
152 """
153 all_models_set: Final[set[str]] = set()
154 if len(team_models) > 0:
155 all_models_set.update(team_models)
156 if SpecialModelNames.all_team_models.value in all_models_set: 156 ↛ 157line 156 didn't jump to line 157 because the condition on line 156 was never true
157 all_models_set.update(team_models)
158 # GH#30619: expand all-team-models sentinel
159 # to the actual proxy model list
160 all_models_set.discard(SpecialModelNames.all_team_models.value)
161 all_models_set.update(proxy_model_list)
162 if include_model_access_groups:
163 all_models_set.update(model_access_groups.keys())
164 if SpecialModelNames.all_proxy_models.value in all_models_set:
165 all_models_set.update(proxy_model_list)
166 if include_model_access_groups:
167 all_models_set.update(model_access_groups.keys())
169 all_models = _get_models_from_access_groups(
170 model_access_groups=model_access_groups,
171 all_models=list(all_models_set),
172 include_model_access_groups=include_model_access_groups,
173 proxy_model_list=proxy_model_list,
174 )
176 # deduplicate while preserving order
177 all_models = list(dict.fromkeys(all_models))
179 verbose_proxy_logger.debug("ALL TEAM MODELS - %s", len(all_models))
180 return all_models
183def get_complete_model_list(
184 key_models: Sequence[str],
185 team_models: Sequence[str],
186 proxy_model_list: list[str],
187 user_model: str | None,
188 infer_model_from_keys: bool | None,
189 return_wildcard_routes: bool | None = False,
190 llm_router: Router | None = None,
191 model_access_groups: dict[str, list[str]] = {},
192 include_model_access_groups: bool | None = False,
193 only_model_access_groups: bool | None = False,
194 team_id: str | None = None,
195) -> list[str]:
196 """Logic for returning complete model list for a given key + team pair"""
198 """
199 - If key list is empty -> defer to team list
200 - If team list is empty -> defer to proxy model list
202 If list contains wildcard -> return known provider models
203 """
205 unique_models: Final = []
207 def append_unique(models):
208 for model in models:
209 if model not in unique_models and model != SpecialModelNames.no_default_models.value: 209 ↛ 208line 209 didn't jump to line 208 because the condition on line 209 was always true
210 unique_models.append(model)
212 if key_models: 212 ↛ 213line 212 didn't jump to line 213 because the condition on line 212 was never true
213 append_unique(key_models)
214 elif team_models:
215 append_unique(team_models)
216 else:
217 append_unique(proxy_model_list)
218 if include_model_access_groups:
219 append_unique(list(model_access_groups.keys())) # TODO: keys order
221 if user_model: 221 ↛ 222line 221 didn't jump to line 222 because the condition on line 221 was never true
222 append_unique([user_model])
224 if infer_model_from_keys: 224 ↛ 225line 224 didn't jump to line 225 because the condition on line 224 was never true
225 valid_models: Final = get_valid_models()
226 append_unique(valid_models)
228 if only_model_access_groups:
229 model_access_groups_to_return: Final[list[str]] = []
230 for model in unique_models:
231 if model in model_access_groups: 231 ↛ 232line 231 didn't jump to line 232 because the condition on line 231 was never true
232 model_access_groups_to_return.append(model)
233 return model_access_groups_to_return
235 all_wildcard_models: Final = _get_wildcard_models(
236 unique_models=unique_models,
237 return_wildcard_routes=return_wildcard_routes,
238 llm_router=llm_router,
239 team_id=team_id,
240 )
242 complete_model_list: Final = unique_models + all_wildcard_models
244 return complete_model_list
247def _hydrate_litellm_credential_name(
248 litellm_params: LiteLLM_Params | None,
249) -> LiteLLM_Params | None:
250 if litellm_params is None or litellm_params.litellm_credential_name is None:
251 return litellm_params
253 credential_values: Final = CredentialAccessor.get_credential_values(litellm_params.litellm_credential_name)
254 if not credential_values:
255 return litellm_params
257 litellm_params = litellm_params.model_copy()
258 for key, value in credential_values.items():
259 if key in _CREDENTIAL_LITELLM_PARAM_FIELDS and getattr(litellm_params, key, None) is None:
260 setattr(litellm_params, key, value)
261 litellm_params.litellm_credential_name = None
262 return litellm_params
265def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_Params | None = None) -> list[str]:
266 wildcard_model_to_expand: Final = (
267 litellm_params.model
268 if wildcard_model == "*"
269 and litellm_params is not None
270 and _check_wildcard_routing(litellm_params.model)
271 and "/" in litellm_params.model
272 else wildcard_model
273 )
274 try:
275 wildcard_provider_prefix, wildcard_suffix = wildcard_model_to_expand.split("/", 1)
276 except ValueError: # safely fail
277 return []
279 # Use provider from litellm_params when available, otherwise from wildcard prefix
280 # (e.g., "openai" from "openai/*" - needed for BYOK where wildcard isn't in router)
281 if litellm_params is not None:
282 try:
283 provider = litellm_params.model.split("/", 1)[0]
284 except ValueError:
285 provider = wildcard_provider_prefix
286 else:
287 provider = wildcard_provider_prefix
289 litellm_params = _hydrate_litellm_credential_name(litellm_params)
291 wildcard_models = get_provider_models(provider=provider, litellm_params=litellm_params)
293 if wildcard_models is None:
294 return []
295 if wildcard_suffix != "*":
296 ## CHECK IF PARTIAL FILTER e.g. `gemini-*`
297 model_prefix: Final = wildcard_suffix.replace("*", "")
299 is_partial_filter: Final = any(wc_model.startswith(model_prefix) for wc_model in wildcard_models)
300 if is_partial_filter:
301 filtered_wildcard_models = [wc_model for wc_model in wildcard_models if wc_model.startswith(model_prefix)]
302 wildcard_models = filtered_wildcard_models
303 else:
304 # add model prefix to wildcard models
305 wildcard_models = [f"{model_prefix}{model}" for model in wildcard_models]
307 known_providers: Final = {provider.value for provider in LlmProviders}
308 suffix_appended_wildcard_models: Final = []
309 for model in wildcard_models:
310 if not model.startswith(wildcard_provider_prefix):
311 # `get_provider_models` returns provider-prefixed ids (e.g. "ollama/gemma3:1b").
312 # When the wildcard uses a custom prefix (e.g. "ollama_server1/*" to distinguish
313 # multiple instances), replace that existing provider prefix instead of stacking
314 # both, which would otherwise yield an uncallable "ollama_server1/ollama/gemma3:1b".
315 # Only strip the leading segment when it is a known provider, so ids whose first
316 # segment is an org rather than a provider (e.g. "meta-llama/Llama-3-8B") keep it.
317 leading, sep, model_suffix = model.partition("/")
318 if sep and leading in known_providers:
319 model = f"{wildcard_provider_prefix}/{model_suffix}"
320 else:
321 model = f"{wildcard_provider_prefix}/{model}"
322 suffix_appended_wildcard_models.append(model)
323 return suffix_appended_wildcard_models or []
326def expand_wildcard_deployments_for_model_info(
327 deployments: list[dict[str, Any]],
328) -> list[dict[str, Any]]:
329 """Expand wildcard deployments into one row per known provider model.
331 PR #30025 changed /model/info to read from llm_router.model_list (correct,
332 so team-scoped rows are included). This function restores wildcard expansion
333 on top of that: a wildcard deployment like model_name="*" / litellm_params.model="openai/*"
334 becomes one entry per known openai model, matching /v1/models behaviour.
335 """
336 expanded: Final[list[dict[str, Any]]] = []
337 for deployment in deployments:
338 model_name = str(deployment.get("model_name") or "")
339 raw_params = deployment.get("litellm_params")
340 litellm_params_dict: dict[str, Any] = raw_params if isinstance(raw_params, dict) else {}
341 litellm_model = str(litellm_params_dict.get("model") or "")
343 # Determine the wildcard pattern to expand.
344 # Branch order matters: only fall to litellm_model when model_name is
345 # also a wildcard, so a concrete model_name is never overwritten.
346 if _check_wildcard_routing(model_name) and "/" in model_name: 346 ↛ 347line 346 didn't jump to line 347 because the condition on line 346 was never true
347 wildcard_pattern = model_name
348 elif _check_wildcard_routing(model_name) and _check_wildcard_routing(litellm_model): 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true
349 wildcard_pattern = litellm_model
350 elif _check_wildcard_routing(model_name): 350 ↛ 351line 350 didn't jump to line 351 because the condition on line 350 was never true
351 wildcard_pattern = model_name
352 else:
353 expanded.append(deployment)
354 continue
356 try:
357 litellm_params = LiteLLM_Params.model_validate(litellm_params_dict) if litellm_params_dict else None
358 except Exception:
359 expanded.append(deployment)
360 continue
361 expanded_names = get_known_models_from_wildcard(
362 wildcard_model=wildcard_pattern,
363 litellm_params=litellm_params,
364 )
365 if not expanded_names:
366 expanded.append(deployment)
367 continue
369 for name in expanded_names:
370 row = copy.deepcopy(deployment)
371 row["model_name"] = name
372 params = row.get("litellm_params")
373 if isinstance(params, dict):
374 params["model"] = name
375 expanded.append(row)
377 return expanded
380def _get_wildcard_models(
381 unique_models: list[str],
382 return_wildcard_routes: bool | None = False,
383 llm_router: Router | None = None,
384 team_id: str | None = None,
385) -> list[str]:
386 models_to_remove: Final = set()
387 all_wildcard_models: Final = []
388 for model in unique_models:
389 if _check_wildcard_routing(model=model):
390 if return_wildcard_routes: 390 ↛ 391line 390 didn't jump to line 391 because the condition on line 390 was never true
391 all_wildcard_models.append(model)
393 models_to_remove.add(model)
395 model_list = llm_router.get_model_list(model_name=model, team_id=team_id) if llm_router else None
396 if model_list: 396 ↛ 397line 396 didn't jump to line 397 because the condition on line 396 was never true
397 for router_model in model_list:
398 all_wildcard_models.extend(
399 get_known_models_from_wildcard(
400 wildcard_model=model,
401 litellm_params=LiteLLM_Params(**router_model["litellm_params"]),
402 )
403 )
404 else:
405 all_wildcard_models.extend(get_known_models_from_wildcard(wildcard_model=model, litellm_params=None))
407 for model in models_to_remove:
408 unique_models.remove(model)
410 return all_wildcard_models
413def get_all_fallbacks(
414 model: str,
415 llm_router: Router | None = None,
416 fallback_type: str = "general",
417) -> list[str]:
418 """
419 Get all fallbacks for a given model from the router's fallback configuration.
421 Args:
422 model: The model name to get fallbacks for
423 llm_router: The LiteLLM router instance
424 fallback_type: Type of fallback ("general", "context_window", "content_policy")
426 Returns:
427 List of fallback model names. Empty list if no fallbacks found.
428 """
429 if llm_router is None: 429 ↛ 430line 429 didn't jump to line 430 because the condition on line 429 was never true
430 return []
432 # Get the appropriate fallback list based on type
433 fallbacks_config: list = []
434 if fallback_type == "general":
435 fallbacks_config = getattr(llm_router, "fallbacks", [])
436 elif fallback_type == "context_window":
437 fallbacks_config = getattr(llm_router, "context_window_fallbacks", [])
438 elif fallback_type == "content_policy": 438 ↛ 441line 438 didn't jump to line 441 because the condition on line 438 was always true
439 fallbacks_config = getattr(llm_router, "content_policy_fallbacks", [])
440 else:
441 verbose_proxy_logger.warning("Unknown fallback_type: %s", fallback_type)
442 return []
444 if not fallbacks_config: 444 ↛ 447line 444 didn't jump to line 447 because the condition on line 444 was always true
445 return []
447 try:
448 # Use existing function to get fallback model group
449 fallback_model_group, _ = get_fallback_model_group(fallbacks=fallbacks_config, model_group=model)
451 if fallback_model_group is None:
452 return []
454 return fallback_model_group
455 except Exception as e:
456 verbose_proxy_logger.error("Error getting fallbacks for model %s: %s", model, e)
457 return []