mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
97dbd8efcb
commit
6a469c2159
2 changed files with 195 additions and 28 deletions
|
|
@ -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 [])
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue