mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: gate key access_group override on group's own assignment
Replaces the previous intersect-with-team.access_group_ids check, which made the override unreachable in practice (the team-gate fallback already covered every case the intersection allowed). The override now resolves each of the key's access_group_ids via get_access_object and accepts the group only if its assigned_team_ids includes the key's team_id, or its assigned_key_ids includes the key's token. This fulfills the original ask (a key can extend a team's allow-list via a group the admin granted to that team or that specific key) while still rejecting foreign groups referenced by team members of other teams.
This commit is contained in:
parent
f17d779666
commit
bf4c250d86
2 changed files with 232 additions and 56 deletions
|
|
@ -2879,28 +2879,57 @@ async def _key_access_group_grants_model(
|
|||
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 it is also assigned to the key's team
|
||||
(i.e., present in `team_object.access_group_ids`). This preserves the
|
||||
team-as-owner boundary: a team member cannot escalate by naming an access
|
||||
group that belongs to a different team.
|
||||
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 or team_object is None:
|
||||
if valid_token is None:
|
||||
return False
|
||||
key_access_group_ids = set(valid_token.access_group_ids or [])
|
||||
team_access_group_ids = set(team_object.access_group_ids or [])
|
||||
allowed_group_ids = key_access_group_ids & team_access_group_ids
|
||||
if not allowed_group_ids:
|
||||
key_access_group_ids = list(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=list(allowed_group_ids),
|
||||
|
||||
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
|
||||
)
|
||||
if not models_from_groups:
|
||||
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=models_from_groups,
|
||||
models=list(set(authorized_models)),
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_id=valid_token.team_id,
|
||||
object_type="key",
|
||||
|
|
|
|||
|
|
@ -1153,9 +1153,44 @@ async def test_can_key_call_model_via_access_group_ids():
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
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_group_covers_model():
|
||||
"""Key's access_group_ids expand to a set that includes the requested model."""
|
||||
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
|
||||
|
|
@ -1163,20 +1198,31 @@ async def test_key_access_group_grants_model_when_group_covers_model():
|
|||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
models=[],
|
||||
access_group_ids=["ryan-access-group"],
|
||||
access_group_ids=["premium-group"],
|
||||
team_id="team-a",
|
||||
)
|
||||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-a",
|
||||
models=["mock-success"],
|
||||
access_group_ids=["ryan-access-group"],
|
||||
access_group_ids=[], # deliberately not synced — the access group itself authorizes
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["claude-haiku-4-5"],
|
||||
):
|
||||
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",
|
||||
|
|
@ -1186,11 +1232,68 @@ async def test_key_access_group_grants_model_when_group_covers_model():
|
|||
)
|
||||
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 cannot override team denial."""
|
||||
"""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(
|
||||
|
|
@ -1202,7 +1305,7 @@ async def test_key_access_group_grants_model_when_key_has_no_groups():
|
|||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-a",
|
||||
models=["mock-success"],
|
||||
access_group_ids=["ryan-access-group"],
|
||||
access_group_ids=["any-group"],
|
||||
)
|
||||
assert (
|
||||
await _key_access_group_grants_model(
|
||||
|
|
@ -1217,7 +1320,7 @@ async def test_key_access_group_grants_model_when_key_has_no_groups():
|
|||
|
||||
@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."""
|
||||
"""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
|
||||
|
|
@ -1225,20 +1328,31 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model():
|
|||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
models=[],
|
||||
access_group_ids=["other-group"],
|
||||
access_group_ids=["basic-group"],
|
||||
team_id="team-a",
|
||||
)
|
||||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-a",
|
||||
models=["mock-success"],
|
||||
access_group_ids=["other-group"],
|
||||
access_group_ids=[],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["gpt-4o-mini"],
|
||||
):
|
||||
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",
|
||||
|
|
@ -1248,21 +1362,25 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model():
|
|||
)
|
||||
is False
|
||||
)
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_access_group_grants_model_when_group_not_assigned_to_team():
|
||||
async def test_key_access_group_grants_model_when_group_authorizes_neither():
|
||||
"""
|
||||
Regression test: a team member naming a foreign access group on their key
|
||||
must NOT escalate to that group's models. The group expands to the requested
|
||||
model, but it isn't assigned to the key's team — so the override is denied.
|
||||
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="test-token",
|
||||
token="team-a-token",
|
||||
models=[],
|
||||
access_group_ids=["team-b-premium"],
|
||||
team_id="team-a",
|
||||
|
|
@ -1270,14 +1388,26 @@ async def test_key_access_group_grants_model_when_group_not_assigned_to_team():
|
|||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-a",
|
||||
models=["mock-success"],
|
||||
access_group_ids=["team-a-basic"],
|
||||
access_group_ids=[],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["claude-opus-4-5"],
|
||||
) as mocked_expand:
|
||||
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",
|
||||
|
|
@ -1287,19 +1417,22 @@ async def test_key_access_group_grants_model_when_group_not_assigned_to_team():
|
|||
)
|
||||
is False
|
||||
)
|
||||
# Foreign group must be filtered out before expansion ever runs.
|
||||
mocked_expand.assert_not_called()
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_access_group_grants_model_when_team_has_no_groups():
|
||||
"""Team with no access_group_ids leaves the intersection empty → denied."""
|
||||
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=["ryan-access-group"],
|
||||
access_group_ids=["missing-group"],
|
||||
team_id="team-a",
|
||||
)
|
||||
team_object = LiteLLM_TeamTable(
|
||||
|
|
@ -1307,12 +1440,26 @@ async def test_key_access_group_grants_model_when_team_has_no_groups():
|
|||
models=["mock-success"],
|
||||
access_group_ids=[],
|
||||
)
|
||||
assert (
|
||||
await _key_access_group_grants_model(
|
||||
model="claude-haiku-4-5",
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
llm_router=None,
|
||||
|
||||
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
|
||||
)
|
||||
is False
|
||||
)
|
||||
finally:
|
||||
for p in patches:
|
||||
p.stop()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue