From f92594f2c68a67d77d744c1344ace26e2a575efa Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Wed, 22 Apr 2026 14:28:58 -0700 Subject: [PATCH] fix: honor key access_group_ids when team restricts models Two model-access gates run per request in `common_checks` and they're asymmetric: `can_key_call_model` falls back to the key's `access_group_ids`, but `can_team_access_model` only looks at `team.models` + `team.access_group_ids`. A key granted a model via its own access group on a model-restricted team is silently denied at the team gate. Wrap `can_team_access_model` in try/except in `common_checks`: on `team_model_access_denied`, consult a new `_key_access_group_grants_model` helper that expands `valid_token.access_group_ids` via the existing `_get_models_from_access_groups` and checks via `_can_object_call_model`. Re-raise if the key's access groups don't grant the model. Any other exception propagates unchanged. Effect: request allowed if `team allows X` OR `key's access group grants X`, making the two gates symmetric. Test: add three unit tests for `_key_access_group_grants_model` covering: group covers model, key has no groups, group resolves but does not cover model. --- litellm/proxy/auth/auth_checks.py | 66 ++++++++++++++---- tests/proxy_unit_tests/test_auth_checks.py | 81 ++++++++++++++++++++++ 2 files changed, 133 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 2c8299e77a..e959867091 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -494,23 +494,27 @@ async def common_checks( # noqa: PLR0915 f"Team={team_object.team_id} is blocked. Update via `/team/unblock` if you're an admin." ) - # 2. If team can call model + # 2. If team can call model (or key's access_group_ids grant it) if _model and team_object: with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"): - if not await can_team_access_model( - model=_model, - team_object=team_object, - llm_router=llm_router, - team_model_aliases=( - valid_token.team_model_aliases if valid_token else None - ), - ): - raise ProxyException( - message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", - type=ProxyErrorTypes.team_model_access_denied, - param="model", - code=status.HTTP_401_UNAUTHORIZED, + try: + await can_team_access_model( + model=_model, + team_object=team_object, + llm_router=llm_router, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ) + except ProxyException as team_denial: + if team_denial.type != ProxyErrorTypes.team_model_access_denied: + raise + if not await _key_access_group_grants_model( + model=_model, + valid_token=valid_token, + llm_router=llm_router, + ): + raise # 2.2. If team member has per-member model scope, enforce it if _model and team_object and valid_token and valid_token.user_id: @@ -2863,6 +2867,40 @@ async def can_team_access_model( raise +async def _key_access_group_grants_model( + model: Union[str, List[str]], + valid_token: Optional[UserAPIKeyAuth], + llm_router: Optional[Router], +) -> bool: + """ + Returns True if the key's `access_group_ids` expand to models that grant + access to `model`. Used to let a key's access group override a team's + model restriction in `common_checks`. + """ + if valid_token is None: + return False + key_access_group_ids = valid_token.access_group_ids or [] + if not key_access_group_ids: + return False + models_from_groups = await _get_models_from_access_groups( + access_group_ids=key_access_group_ids, + ) + if not models_from_groups: + return False + try: + _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", + ) + return True + except ProxyException: + return False + + def can_project_access_model( model: Union[str, List[str]], project_object: LiteLLM_ProjectTableCachedObj, diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 86cd5c0c41..6a404712c1 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -1146,3 +1146,84 @@ async def test_can_key_call_model_via_access_group_ids(): valid_token=user_api_key_object, llm_router=router, ) + + +# --------------------------------------------------------------------------- +# _key_access_group_grants_model (key access group overriding team restriction) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_group_covers_model(): + """Key's access_group_ids expand to a set that includes the requested model.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["ryan-access-group"], + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + new_callable=AsyncMock, + return_value=["claude-haiku-4-5"], + ): + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + llm_router=None, + ) + is True + ) + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_key_has_no_groups(): + """Key with no access_group_ids cannot override team denial.""" + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=[], + ) + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + llm_router=None, + ) + is False + ) + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_group_does_not_cover_model(): + """Key's access_group_ids expand to models that do not include the request.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["other-group"], + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + new_callable=AsyncMock, + return_value=["gpt-4o-mini"], + ): + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + llm_router=None, + ) + is False + )