Merge pull request #26275 from BerriAI/litellm_fix-ag-not-resolved

This commit is contained in:
ryan-crabbe-berri 2026-05-01 18:37:38 -07:00 • committed by GitHub
commit 85d426c6b5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 407 additions and 14 deletions

View file

@ -502,23 +502,28 @@ 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,
team_object=team_object,
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:
@ -2975,6 +2980,77 @@ async def can_team_access_model(
raise
async def _key_access_group_grants_model(
model: Union[str, List[str]],
valid_token: Optional[UserAPIKeyAuth],
team_object: Optional[LiteLLM_TeamTable],
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`.
A key's access group only counts if the access group itself authorizes the
caller as an owner — that is, the group's `assigned_team_ids` includes the
key's `team_id`, or the group's `assigned_key_ids` includes the key's
token. This preserves the team-as-owner boundary (a team member cannot
escalate by naming a group assigned to a different team) while still
letting a group reach the key without first being added to the team's
`access_group_ids` list.
"""
if valid_token is None:
return False
key_access_group_ids = list(valid_token.access_group_ids or [])
if not key_access_group_ids:
return False
from litellm.proxy.proxy_server import prisma_client as _prisma_client
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj
from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache
if _prisma_client is None or _user_api_key_cache is None:
return False
key_team_id = valid_token.team_id or (
team_object.team_id if team_object is not None else None
)
key_token = valid_token.token
authorized_models: List[str] = []
for ag_id in key_access_group_ids:
try:
ag = await get_access_object(
access_group_id=ag_id,
prisma_client=_prisma_client,
user_api_key_cache=_user_api_key_cache,
proxy_logging_obj=_proxy_logging_obj,
)
except Exception:
continue
team_authorized = bool(
key_team_id and key_team_id in (ag.assigned_team_ids or [])
)
key_authorized = bool(key_token and key_token in (ag.assigned_key_ids or []))
if team_authorized or key_authorized:
authorized_models.extend(ag.access_model_names or [])
if not authorized_models:
return False
try:
_can_object_call_model(
model=model,
llm_router=llm_router,
models=list(set(authorized_models)),
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,

View file

@ -1153,3 +1153,320 @@ 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)
# ---------------------------------------------------------------------------
def _patch_proxy_server_globals():
"""Patch proxy_server's prisma_client and user_api_key_cache to non-None mocks
so the helper's None-guard doesn't short-circuit. The actual values don't
matter because get_access_object is patched separately to return fixtures."""
from unittest.mock import MagicMock, patch
return [
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
]
def _fake_access_group(
access_group_id: str,
access_model_names=None,
assigned_team_ids=None,
assigned_key_ids=None,
):
from litellm.proxy._types import LiteLLM_AccessGroupTable
return LiteLLM_AccessGroupTable(
access_group_id=access_group_id,
access_group_name=access_group_id,
access_model_names=access_model_names or [],
assigned_team_ids=assigned_team_ids or [],
assigned_key_ids=assigned_key_ids or [],
)
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_team_authorized():
"""Group's assigned_team_ids includes the key's team and grants the model → True.
This is the happy path equivalent of Andres's report: admin creates an
access group with assigned_team_ids=[team-a], grants claude-haiku-4-5,
attaches it to a key on team-a. Override fires.
"""
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=["premium-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[], # deliberately not synced — the access group itself authorizes
)
fake_ag = _fake_access_group(
access_group_id="premium-group",
access_model_names=["claude-haiku-4-5"],
assigned_team_ids=["team-a"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is True
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_key_directly_authorized():
"""Group's assigned_key_ids includes the key's token and grants the model → True.
Per-key authorization path: an admin scopes a group directly to a key
(assigned_key_ids) without listing the team.
"""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token-hashed",
models=[],
access_group_ids=["per-key-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="per-key-group",
access_model_names=["claude-haiku-4-5"],
assigned_team_ids=[],
assigned_key_ids=["test-token-hashed"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is True
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_key_has_no_groups():
"""Key with no access_group_ids → False (early return, no DB read)."""
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=[],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=["any-group"],
)
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_group_does_not_cover_model():
"""Group authorizes the team but does not grant the requested model → False."""
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=["basic-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="basic-group",
access_model_names=["gpt-4o-mini"],
assigned_team_ids=["team-a"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_group_authorizes_neither():
"""
Bypass regression test: a team member sets a foreign access group on their
key. The group grants the requested model but its assigned_team_ids /
assigned_key_ids do not include this caller's team or token. Override is
denied — the team's 401 propagates.
"""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="team-a-token",
models=[],
access_group_ids=["team-b-premium"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="team-b-premium",
access_model_names=["claude-opus-4-5"],
assigned_team_ids=["team-b"],
assigned_key_ids=["team-b-token"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-opus-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_get_access_object_raises():
"""Group lookup failure (404, network, etc.) is treated as no authorization."""
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=["missing-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
side_effect=Exception("not found"),
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()