fix(access_groups): derive attached teams from the team table and reject unknown team ids

GET /v1/access_group and GET /v1/access_group/{id} (and the /v1/unified_access_group aliases) used to
return the assigned_team_ids column verbatim. That column is a denormalized mirror of
LiteLLM_TeamTable.access_group_ids and can be stale or hold ids of teams that no longer exist, so the
Attached Teams view drifted from reality.

The read path now runs one team find_many per request, unioning teams whose access_group_ids carry any
group in the response with teams listed in the stored columns. Only real team rows come back, so ghost
ids drop out and teams the mirror missed are added. The stored order is kept for ids that survive and
newly discovered teams are appended.

Create and update now resolve the requested assigned_team_ids inside the transaction and answer 400
with the missing ids before anything is written, instead of silently storing ids that point nowhere.

Refs LIT-6593

Claude-Session: https://claude.ai/code/session_01QvQzYztinxj8ZuD5YxbVdL
This commit is contained in:
ryan-crabbe-berri 2026-09-01 15:46:09 -07:00
parent 97dbd8efcb
commit 6a469c2159
2 changed files with 195 additions and 28 deletions

View file

@ -1,4 +1,5 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final, Protocol
from fastapi import APIRouter, Depends, HTTPException, status
@ -119,8 +120,57 @@ def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
)
def _record_to_response(record: _AccessGroupRecord) -> AccessGroupResponse:
return AccessGroupResponse.model_validate(record.dict())
def _record_to_response(
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str] | None = None
) -> AccessGroupResponse:
stored: Final = record.dict()
payload: Final = (
stored if assigned_team_ids is None else MappingProxyType({**stored, "assigned_team_ids": assigned_team_ids})
)
return AccessGroupResponse.model_validate(payload)
def _attached_team_ids_by_group(
records: Sequence[_AccessGroupRecord], teams: Sequence[_TeamRecord]
) -> Mapping[str, tuple[str, ...]]:
"""Teams really attached to each group: the stored column minus ghosts, plus teams the mirror missed."""
real_team_ids: Final = frozenset(team.team_id for team in teams)
def attached(record: _AccessGroupRecord) -> tuple[str, ...]:
stored: Final = (team_id for team_id in (record.assigned_team_ids or ()) if team_id in real_team_ids)
carrying: Final = (team.team_id for team in teams if record.access_group_id in (team.access_group_ids or ()))
return tuple(dict.fromkeys((*stored, *carrying)))
return MappingProxyType({record.access_group_id: attached(record) for record in records})
async def _attached_team_ids_for(
team_table: _TeamTable, records: Sequence[_AccessGroupRecord]
) -> Mapping[str, tuple[str, ...]]:
if not records:
return MappingProxyType({})
group_ids: Final = tuple(record.access_group_id for record in records)
stored_team_ids: Final = tuple(
dict.fromkeys(team_id for record in records for team_id in (record.assigned_team_ids or ()))
)
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where must be a dict
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where must be a dict
clauses: Final = (carrying, listed) if stored_team_ids else (carrying,)
where: Final = {"OR": clauses} # mutable-ok: prisma where must be a dict
return _attached_team_ids_by_group(records, await team_table.find_many(where=where))
async def _require_teams_exist(tx: _AccessGroupTx, team_ids: Sequence[str]) -> None:
if not team_ids:
return
where: Final = {"team_id": {"in": team_ids}} # mutable-ok: prisma where must be a dict
found: Final = await tx.litellm_teamtable.find_many(where=where)
missing: Final = frozenset(team_ids) - frozenset(team.team_id for team in found)
if missing:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unknown team ids: {', '.join(sorted(missing))}",
)
def _record_to_access_group_table(record: _AccessGroupRecord) -> LiteLLM_AccessGroupTable:
@ -330,6 +380,7 @@ async def create_access_group(
status_code=status.HTTP_409_CONFLICT,
detail=f"Access group '{data.access_group_name}' already exists",
)
await _require_teams_exist(tx, data.assigned_team_ids or ())
record: Final = await tx.litellm_accessgrouptable.create(
data={
@ -390,7 +441,8 @@ async def list_access_groups(
table: Final = AccessGroupRepository(prisma_client).table
records: Final = await table.find_many(order={"created_at": "desc"})
return [_record_to_response(r) for r in records]
attached: Final = await _attached_team_ids_for(prisma_client.db.litellm_teamtable, records)
return [_record_to_response(r, assigned_team_ids=attached[r.access_group_id]) for r in records]
@router.get(
@ -411,7 +463,8 @@ async def get_access_group(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
return _record_to_response(record)
attached: Final = await _attached_team_ids_for(prisma_client.db.litellm_teamtable, (record,))
return _record_to_response(record, assigned_team_ids=attached[record.access_group_id])
@router.put(
@ -461,6 +514,7 @@ async def update_access_group(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
await _require_teams_exist(tx, data.assigned_team_ids or ())
old_team_ids: Final[set[str]] = set(existing.assigned_team_ids or [])
old_key_ids: Final[set[str]] = set(existing.assigned_key_ids or [])

View file

@ -12,13 +12,12 @@ from fastapi.testclient import TestClient
from prisma.errors import PrismaError
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import app
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.proxy_server import app
def _make_access_group_record(
@ -58,6 +57,10 @@ def _make_access_group_record(
return record
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None):
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [])
@pytest.fixture
def client_and_mocks(monkeypatch):
"""Setup mock prisma and admin auth for access group endpoints."""
@ -185,7 +188,8 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
)
def test_create_access_group_success(client_and_mocks, base_path, payload):
"""Create access group with various payloads returns 201."""
client, _, mock_table, *_ = client_and_mocks
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[_make_team_record("team-1")])
resp = client.post(base_path, json=payload)
assert resp.status_code == 201
@ -277,13 +281,44 @@ def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_success_empty(client_and_mocks, base_path):
"""List access groups returns empty list when none exist."""
client, _, mock_table, *_ = client_and_mocks
"""List access groups returns empty list when none exist, without querying teams."""
client, mock_prisma, mock_table, *_ = client_and_mocks
resp = client.get(base_path)
assert resp.status_code == 200
assert resp.json() == []
mock_table.find_many.assert_awaited_once()
mock_prisma.db.litellm_teamtable.find_many.assert_not_awaited()
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_attributes_teams_per_group_with_one_query(client_and_mocks, base_path):
"""List derives each group's teams from the team table in a single query, attributed per group."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
records = [
_make_access_group_record(access_group_id="ag-1", access_group_name="group-1"),
_make_access_group_record(access_group_id="ag-2", access_group_name="group-2"),
]
mock_table.find_many = AsyncMock(return_value=records)
mock_team_table.find_many = AsyncMock(
return_value=[
_make_team_record("team-x", ["ag-1"]),
_make_team_record("team-y", ["ag-2"]),
_make_team_record("team-z", ["ag-1", "ag-2"]),
]
)
resp = client.get(base_path)
assert resp.status_code == 200
body = resp.json()
assert body[0]["assigned_team_ids"] == ["team-x", "team-z"]
assert body[1]["assigned_team_ids"] == ["team-y", "team-z"]
mock_team_table.find_many.assert_awaited_once()
(carrying,) = mock_team_table.find_many.call_args.kwargs["where"]["OR"]
assert list(carrying["access_group_ids"]["hasSome"]) == ["ag-1", "ag-2"]
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
@ -373,6 +408,43 @@ def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
assert resp.json()["access_group_id"] == access_group_id
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_get_access_group_derives_assigned_teams_from_team_table(client_and_mocks, base_path):
"""Get drops ghost ids from the stored column and adds teams that carry the group but were never mirrored."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
record = _make_access_group_record(access_group_id="ag-123", assigned_team_ids=["team-a", "ghost-team"])
mock_table.find_unique = AsyncMock(return_value=record)
mock_team_table.find_many = AsyncMock(
return_value=[
_make_team_record("team-a", ["ag-123"]),
_make_team_record("team-b", ["ag-123"]),
_make_team_record("team-c", ["ag-123"]),
]
)
resp = client.get(f"{base_path}/ag-123")
assert resp.status_code == 200
assert resp.json()["assigned_team_ids"] == ["team-a", "team-b", "team-c"]
carrying, listed = mock_team_table.find_many.call_args.kwargs["where"]["OR"]
assert list(carrying["access_group_ids"]["hasSome"]) == ["ag-123"]
assert list(listed["team_id"]["in"]) == ["team-a", "ghost-team"]
def test_get_access_group_empty_column_and_no_teams_returns_empty(client_and_mocks):
"""Get returns [] when the column is empty and no team carries the group."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=_make_access_group_record(access_group_id="ag-123"))
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
resp = client.get("/v1/access_group/ag-123")
assert resp.status_code == 200
assert resp.json()["assigned_team_ids"] == []
def test_get_access_group_not_found(client_and_mocks):
"""Get access group returns 404 when not found."""
client, _, mock_table, *_ = client_and_mocks
@ -985,6 +1057,28 @@ def test_record_to_access_group_table():
assert result.access_agent_ids == ["agent-1"]
def test_attached_team_ids_by_group_keeps_column_order_then_appends_unmirrored_teams():
"""Stored ids that resolve keep their order, ghosts drop, carriers the mirror missed append once, per group."""
from litellm.proxy.management_endpoints.access_group_endpoints import (
_attached_team_ids_by_group,
)
records = [
_make_access_group_record(access_group_id="ag-1", assigned_team_ids=["team-b", "ghost", "team-a"]),
_make_access_group_record(access_group_id="ag-2", assigned_team_ids=[]),
]
teams = [
_make_team_record("team-a", ["ag-1"]),
_make_team_record("team-b", []),
_make_team_record("team-c", ["ag-1"]),
_make_team_record("team-d", ["ag-2"]),
]
result = _attached_team_ids_by_group(records, teams)
assert dict(result) == {"ag-1": ("team-b", "team-a", "team-c"), "ag-2": ("team-d",)}
# ---------------------------------------------------------------------------
# Sync tests: CREATE
# ---------------------------------------------------------------------------
@ -997,9 +1091,8 @@ def test_create_access_group_syncs_assigned_teams(client_and_mocks):
)
mock_team_table = mock_prisma.db.litellm_teamtable
team_record = MagicMock()
team_record.team_id = "team-1"
team_record.access_group_ids = []
team_record = _make_team_record("team-1")
mock_team_table.find_many = AsyncMock(return_value=[team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.post(
@ -1043,20 +1136,22 @@ def test_create_access_group_syncs_assigned_keys(client_and_mocks):
assert "ag-new" in call_kwargs["data"]["access_group_ids"]
def test_create_access_group_skips_sync_for_nonexistent_team(client_and_mocks):
"""Create skips updating a team that doesn't exist in DB."""
client, mock_prisma, _, mock_cache, mock_proxy_logging = client_and_mocks
def test_create_access_group_rejects_nonexistent_team(client_and_mocks):
"""Create refuses to store a team id that does not resolve to a team row."""
client, mock_prisma, mock_access_group_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
mock_team_table.find_unique = AsyncMock(return_value=None)
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-real")])
resp = client.post(
"/v1/access_group",
json={
"access_group_name": "new-group",
"assigned_team_ids": ["nonexistent-team"],
"assigned_team_ids": ["team-real", "nonexistent-team", "also-missing"],
},
)
assert resp.status_code == 201
assert resp.status_code == 400
assert resp.json()["detail"] == "Unknown team ids: also-missing, nonexistent-team"
mock_access_group_table.create.assert_not_awaited()
mock_team_table.update.assert_not_awaited()
@ -1065,9 +1160,8 @@ def test_create_access_group_idempotent_team_sync(client_and_mocks):
client, mock_prisma, _, mock_cache, mock_proxy_logging = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
team_record = MagicMock()
team_record.team_id = "team-1"
team_record.access_group_ids = ["ag-new"] # already synced
team_record = _make_team_record("team-1", ["ag-new"])
mock_team_table.find_many = AsyncMock(return_value=[team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.post(
@ -1095,9 +1189,8 @@ def test_update_access_group_syncs_added_teams(client_and_mocks):
)
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
team_record = MagicMock()
team_record.team_id = "team-new"
team_record.access_group_ids = []
team_record = _make_team_record("team-new")
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-existing", ["ag-update"]), team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.put(
@ -1113,6 +1206,25 @@ def test_update_access_group_syncs_added_teams(client_and_mocks):
assert "ag-update" in call_kwargs["data"]["access_group_ids"]
def test_update_access_group_rejects_nonexistent_team(client_and_mocks):
"""Update refuses to store a team id that does not resolve to a team row and leaves the group untouched."""
client, mock_prisma, mock_access_group_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
existing = _make_access_group_record(access_group_id="ag-update", assigned_team_ids=["team-existing"])
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-existing", ["ag-update"])])
resp = client.put(
"/v1/access_group/ag-update",
json={"assigned_team_ids": ["team-existing", "team-ghost"]},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Unknown team ids: team-ghost"
mock_access_group_table.update.assert_not_awaited()
mock_team_table.update.assert_not_awaited()
def test_update_access_group_syncs_removed_teams(client_and_mocks):
"""Update removes access_group_id from de-assigned teams."""
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = (
@ -1125,9 +1237,8 @@ def test_update_access_group_syncs_removed_teams(client_and_mocks):
)
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
team_to_remove = MagicMock()
team_to_remove.team_id = "team-remove"
team_to_remove.access_group_ids = ["ag-update"]
team_to_remove = _make_team_record("team-remove", ["ag-update"])
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-keep", ["ag-update"])])
mock_team_table.find_unique = AsyncMock(return_value=team_to_remove)
resp = client.put(
@ -1160,6 +1271,7 @@ def test_update_access_group_no_team_sync_when_ids_not_in_payload(client_and_moc
resp = client.put("/v1/access_group/ag-update", json={"description": "new desc"})
assert resp.status_code == 200
mock_team_table.find_many.assert_not_awaited()
mock_team_table.find_unique.assert_not_awaited()
mock_team_table.update.assert_not_awaited()
@ -1293,7 +1405,7 @@ def test_delete_access_group_handles_out_of_sync_assigned_keys(client_and_mocks)
def test_update_access_group_null_assigned_ids_treated_as_empty(client_and_mocks):
"""Update with explicit null for assigned_*_ids clears the list and writes [] to DB."""
client, _, mock_table, *_ = client_and_mocks
client, mock_prisma, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(
access_group_id="ag-update",
@ -1313,3 +1425,4 @@ def test_update_access_group_null_assigned_ids_treated_as_empty(client_and_mocks
update_call_kwargs = mock_table.update.call_args.kwargs
assert update_call_kwargs["data"]["assigned_team_ids"] == []
assert update_call_kwargs["data"]["assigned_key_ids"] == []
mock_prisma.db.litellm_teamtable.find_many.assert_not_awaited()