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

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 

6 

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 

17 

18_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields) 

19 

20 

21_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields) 

22 

23 

24def _check_wildcard_routing(model: str) -> bool: 

25 """ 

26 Returns True if a model is a provider wildcard. 

27 

28 eg: 

29 - anthropic/* 

30 - openai/* 

31 - * 

32 """ 

33 if "*" in model: 

34 return True 

35 return False 

36 

37 

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) 

44 

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 

49 

50 

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 

69 

70 

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 

78 

79 if prisma_client is None: 

80 return [] 

81 

82 if user_api_key_dict.object_permission_id is None: 

83 return [] 

84 

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 [] 

95 

96 

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()) 

126 

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 ) 

133 

134 # deduplicate while preserving order 

135 all_models = list(dict.fromkeys(all_models)) 

136 

137 verbose_proxy_logger.debug("ALL KEY MODELS - %s", len(all_models)) 

138 return all_models 

139 

140 

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()) 

168 

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 ) 

175 

176 # deduplicate while preserving order 

177 all_models = list(dict.fromkeys(all_models)) 

178 

179 verbose_proxy_logger.debug("ALL TEAM MODELS - %s", len(all_models)) 

180 return all_models 

181 

182 

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""" 

197 

198 """ 

199 - If key list is empty -> defer to team list 

200 - If team list is empty -> defer to proxy model list 

201 

202 If list contains wildcard -> return known provider models 

203 """ 

204 

205 unique_models: Final = [] 

206 

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) 

211 

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 

220 

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]) 

223 

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) 

227 

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 

234 

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 ) 

241 

242 complete_model_list: Final = unique_models + all_wildcard_models 

243 

244 return complete_model_list 

245 

246 

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 

252 

253 credential_values: Final = CredentialAccessor.get_credential_values(litellm_params.litellm_credential_name) 

254 if not credential_values: 

255 return litellm_params 

256 

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 

263 

264 

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 [] 

278 

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 

288 

289 litellm_params = _hydrate_litellm_credential_name(litellm_params) 

290 

291 wildcard_models = get_provider_models(provider=provider, litellm_params=litellm_params) 

292 

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("*", "") 

298 

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] 

306 

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 [] 

324 

325 

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. 

330 

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 "") 

342 

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 

355 

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 

368 

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) 

376 

377 return expanded 

378 

379 

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) 

392 

393 models_to_remove.add(model) 

394 

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)) 

406 

407 for model in models_to_remove: 

408 unique_models.remove(model) 

409 

410 return all_wildcard_models 

411 

412 

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. 

420 

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") 

425 

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 [] 

431 

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 [] 

443 

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 [] 

446 

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) 

450 

451 if fallback_model_group is None: 

452 return [] 

453 

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 []