diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 0b30999aa21..224c936a1d5 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2891,18 +2891,43 @@ async def can_key_call_model( Raises: - Exception: If token not allowed to call model """ + key_access_group_ids = valid_token.access_group_ids or [] + key_models = valid_token.models or [] + key_defers_to_team_models = ( + len(key_models) == 0 or SpecialModelNames.all_team_models.value in key_models + ) + + if key_access_group_ids and key_defers_to_team_models: + models_from_groups = await _get_models_from_access_groups( + access_group_ids=key_access_group_ids, + ) + if models_from_groups: + return _can_object_call_model( + model=model, + llm_router=llm_router, + models=models_from_groups, + team_model_aliases=valid_token.team_model_aliases, + team_id=valid_token.team_id, + object_type="key", + ) + raise ProxyException( + message=f"key not allowed to access model. This key has access_group_ids={key_access_group_ids}, but those groups do not grant any models. Tried to access {model}", + type=ProxyErrorTypes.key_model_access_denied, + param="model", + code=status.HTTP_403_FORBIDDEN, + ) + try: return _can_object_call_model( model=model, llm_router=llm_router, - models=valid_token.models, + models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, object_type="key", ) except ProxyException: # Fallback: check key's access_group_ids - key_access_group_ids = valid_token.access_group_ids or [] if key_access_group_ids: models_from_groups = await _get_models_from_access_groups( access_group_ids=key_access_group_ids, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 03167c5a2dc..6e29bd9c78a 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2339,7 +2339,9 @@ async def _enforce_key_and_fallback_model_access( new_model_list = model_list verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}") elif ( - isinstance(valid_token.models, list) and "all-team-models" in valid_token.models + isinstance(valid_token.models, list) + and "all-team-models" in valid_token.models + and not valid_token.access_group_ids ): pass else: