mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): refuse a team admin's budget write when the budget changed mid-request
The keep-or-lower check compares against the budget update_team read, so the write now only lands while the stored max_budget still matches it and answers 409 otherwise. A concurrent proxy admin cut can no longer be overwritten with a higher value.
This commit is contained in:
parent
d3f0607820
commit
fc13cea479
3 changed files with 198 additions and 96 deletions
|
|
@ -16,6 +16,7 @@ import math
|
|||
import traceback
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
|
|
@ -340,6 +341,14 @@ class _ErrorDetail(TypedDict):
|
|||
error: ReadOnly[str]
|
||||
|
||||
|
||||
class _TeamIdWhere(TypedDict):
|
||||
team_id: ReadOnly[str]
|
||||
|
||||
|
||||
class _TeamIdAndBudgetWhere(_TeamIdWhere):
|
||||
max_budget: ReadOnly[float | None]
|
||||
|
||||
|
||||
class _TeamCreateTx(AccessGroupSyncTx, Protocol):
|
||||
@property
|
||||
def litellm_teamtable(self) -> "TableActions[prisma_models.LiteLLM_TeamTable]": ...
|
||||
|
|
@ -1200,11 +1209,18 @@ async def _check_user_team_limits(
|
|||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _MaxBudgetGuard:
|
||||
"""The team write only lands while the stored max_budget still equals `expected`."""
|
||||
|
||||
expected: float | None
|
||||
|
||||
|
||||
def _check_team_budget_update_authority(
|
||||
data: UpdateTeamRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
existing_team_max_budget: float | None,
|
||||
) -> None:
|
||||
) -> _MaxBudgetGuard | None:
|
||||
"""
|
||||
Restrict who can grow a team's spend ceiling on /team/update.
|
||||
|
||||
|
|
@ -1213,13 +1229,19 @@ def _check_team_budget_update_authority(
|
|||
removing the cap (setting it to None). Setting a finite budget on a team
|
||||
that has no cap is a restriction and is allowed. Org admins editing
|
||||
org-scoped teams are governed by _check_org_team_limits() instead.
|
||||
|
||||
The verdict holds only for the budget it was checked against, so a restricted
|
||||
caller's budget write gets a guard; without it, a concurrent budget cut could
|
||||
be overwritten with a higher value.
|
||||
"""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return
|
||||
if existing_team_max_budget is None:
|
||||
return
|
||||
return None
|
||||
|
||||
budget_explicitly_set: Final = "max_budget" in (getattr(data, "model_fields_set", None) or set())
|
||||
guard: Final = _MaxBudgetGuard(expected=existing_team_max_budget) if budget_explicitly_set else None
|
||||
if existing_team_max_budget is None:
|
||||
return guard
|
||||
|
||||
if budget_explicitly_set and data.max_budget is None:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -1235,6 +1257,37 @@ def _check_team_budget_update_authority(
|
|||
"error": f"Only a proxy admin can raise a team's max_budget. Team's current max_budget={existing_team_max_budget}, requested={data.max_budget}."
|
||||
},
|
||||
)
|
||||
return guard
|
||||
|
||||
|
||||
_TEAM_UPDATE_INCLUDE: Final = MappingProxyType(
|
||||
{
|
||||
"litellm_model_table": True,
|
||||
# `object_permission` is included so `_refresh_cached_team`
|
||||
# doesn't write a cached team with the relation nulled out.
|
||||
# See team_model_add for the full rationale.
|
||||
"object_permission": True,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def _write_team_update(
|
||||
prisma_client: PrismaClient | None,
|
||||
team_id: str,
|
||||
team_update_data: Mapping[str, object],
|
||||
max_budget_guard: _MaxBudgetGuard | None,
|
||||
) -> "prisma_models.LiteLLM_TeamTable | None":
|
||||
by_id: Final[_TeamIdWhere] = {"team_id": team_id}
|
||||
if max_budget_guard is None:
|
||||
return await _team_db(prisma_client).update(where=by_id, data=team_update_data, include=_TEAM_UPDATE_INCLUDE)
|
||||
by_id_and_budget: Final[_TeamIdAndBudgetWhere] = {"team_id": team_id, "max_budget": max_budget_guard.expected}
|
||||
written: Final = await _team_db(prisma_client).update_many(where=by_id_and_budget, data=team_update_data)
|
||||
if written == 0:
|
||||
conflict: Final[_ErrorDetail] = {
|
||||
"error": "The team's max_budget changed during this update. Reload the team and try again."
|
||||
}
|
||||
raise HTTPException(status_code=409, detail=conflict)
|
||||
return await _team_db(prisma_client).find_unique(where=by_id, include=_TEAM_UPDATE_INCLUDE)
|
||||
|
||||
|
||||
def _existing_model_cap(raw_budget_config: object) -> BudgetConfig | None:
|
||||
|
|
@ -2341,12 +2394,15 @@ async def update_team(
|
|||
|
||||
# A team admin never grows its own team's spend ceiling. Org admins grow org-scoped teams
|
||||
# within the org limits _check_org_team_limits() enforced above.
|
||||
if org_id_to_check is None or access_role == "team_admin":
|
||||
max_budget_guard: Final = (
|
||||
_check_team_budget_update_authority(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
existing_team_max_budget=existing_team_row.max_budget,
|
||||
)
|
||||
if org_id_to_check is None or access_role == "team_admin"
|
||||
else None
|
||||
)
|
||||
_check_team_model_budget_update_authority(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -2493,17 +2549,7 @@ async def update_team(
|
|||
|
||||
updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv)
|
||||
team_update_data: Final[Mapping[str, object]] = updated_kv
|
||||
team_row: Final = await _team_db(prisma_client).update(
|
||||
where={"team_id": data.team_id},
|
||||
data=team_update_data,
|
||||
# `object_permission` is included so `_refresh_cached_team`
|
||||
# doesn't write a cached team with the relation nulled out.
|
||||
# See team_model_add for the full rationale.
|
||||
include={
|
||||
"litellm_model_table": True,
|
||||
"object_permission": True,
|
||||
},
|
||||
)
|
||||
team_row: Final = await _write_team_update(prisma_client, data.team_id, team_update_data, max_budget_guard)
|
||||
|
||||
if team_row is None or team_row.team_id is None:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -609,10 +609,14 @@ class TestTeamAdminWithRpmLimitAndMaxBudgetEnabled:
|
|||
lower the team's budget. Raising or removing the budget stays with the proxy admin."""
|
||||
|
||||
@pytest.mark.covers("mgmt.team.update.team_admin_limited_to_enabled_fields")
|
||||
def test_team_admin_saves_a_new_rpm_limit_and_a_lower_budget(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
@pytest.mark.parametrize(
|
||||
"current_budget",
|
||||
[pytest.param(_TEAM_MAX_BUDGET, id="lower"), pytest.param(None, id="first-budget")],
|
||||
)
|
||||
def test_team_admin_saves_a_new_rpm_limit_and_a_tighter_budget(
|
||||
self, client: ManagementClient, resources: ResourceManager, current_budget: float | None
|
||||
) -> None:
|
||||
team_id, admin_key = _team_with_admin(client, resources, max_budget=_TEAM_MAX_BUDGET)
|
||||
team_id, admin_key = _team_with_admin(client, resources, max_budget=current_budget)
|
||||
access = _read_team(client, team_id, admin_key).team_info.caller_edit_access
|
||||
assert access == CallerEditAccess(kind="team_admin", editable_fields=["max_budget", "rpm_limit"]), (
|
||||
f"/team/info should list max_budget and rpm_limit as the team admin's editable fields, got {access}"
|
||||
|
|
@ -624,8 +628,8 @@ class TestTeamAdminWithRpmLimitAndMaxBudgetEnabled:
|
|||
)
|
||||
|
||||
assert outcome.status_code == 200, (
|
||||
f"a team admin setting an RPM limit and lowering the budget must succeed, got {outcome.status_code}: "
|
||||
f"{outcome.body[:300]}"
|
||||
f"a team admin setting an RPM limit and tightening the budget from {current_budget} must succeed, "
|
||||
f"got {outcome.status_code}: {outcome.body[:300]}"
|
||||
)
|
||||
after = _poll_team(
|
||||
client,
|
||||
|
|
|
|||
|
|
@ -6653,40 +6653,18 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed(
|
|||
"litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()
|
||||
),
|
||||
):
|
||||
mock_existing_team = MagicMock()
|
||||
mock_existing_team.team_id = "standalone-uncapped-123"
|
||||
mock_existing_team.organization_id = None
|
||||
mock_existing_team.max_budget = None # team has no cap
|
||||
mock_existing_team.model_id = None
|
||||
mock_existing_team.model_dump.return_value = {
|
||||
"team_id": "standalone-uncapped-123",
|
||||
"organization_id": None,
|
||||
"max_budget": None,
|
||||
"members_with_roles": [
|
||||
{"user_id": "uncapped-team-admin", "role": "admin"}
|
||||
],
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_existing_team
|
||||
_TeamRowStore(
|
||||
mock_prisma.db.litellm_teamtable,
|
||||
{
|
||||
"team_id": "standalone-uncapped-123",
|
||||
"max_budget": None,
|
||||
"members_with_roles": [{"user_id": "uncapped-team-admin", "role": "admin"}],
|
||||
},
|
||||
)
|
||||
mock_prisma.jsonify_team_object = lambda db_data: db_data
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
mock_updated_team = MagicMock()
|
||||
mock_updated_team.team_id = "standalone-uncapped-123"
|
||||
mock_updated_team.organization_id = None
|
||||
mock_updated_team.max_budget = 1000.0
|
||||
mock_updated_team.litellm_model_table = None
|
||||
mock_updated_team.model_dump.return_value = {
|
||||
"team_id": "standalone-uncapped-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 1000.0,
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=mock_updated_team
|
||||
)
|
||||
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
http_request=dummy_request,
|
||||
|
|
@ -6847,21 +6825,13 @@ async def test_update_team_standalone_lower_budget_allowed(
|
|||
"litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()
|
||||
) as mock_audit,
|
||||
):
|
||||
mock_existing_team = MagicMock()
|
||||
mock_existing_team.team_id = "standalone-lower-budget-123"
|
||||
mock_existing_team.organization_id = None
|
||||
mock_existing_team.max_budget = 500.0
|
||||
mock_existing_team.model_id = None
|
||||
mock_existing_team.model_dump.return_value = {
|
||||
"team_id": "standalone-lower-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 500.0,
|
||||
"members_with_roles": [
|
||||
{"user_id": "standalone-lower-budget-admin", "role": "admin"}
|
||||
],
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_existing_team
|
||||
_TeamRowStore(
|
||||
mock_prisma.db.litellm_teamtable,
|
||||
{
|
||||
"team_id": "standalone-lower-budget-123",
|
||||
"max_budget": 500.0,
|
||||
"members_with_roles": [{"user_id": "standalone-lower-budget-admin", "role": "admin"}],
|
||||
},
|
||||
)
|
||||
mock_prisma.jsonify_team_object = lambda db_data: db_data
|
||||
|
||||
|
|
@ -6872,20 +6842,6 @@ async def test_update_team_standalone_lower_budget_allowed(
|
|||
mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
mock_updated_team = MagicMock()
|
||||
mock_updated_team.team_id = "standalone-lower-budget-123"
|
||||
mock_updated_team.organization_id = None
|
||||
mock_updated_team.max_budget = 300.0
|
||||
mock_updated_team.litellm_model_table = None
|
||||
mock_updated_team.model_dump.return_value = {
|
||||
"team_id": "standalone-lower-budget-123",
|
||||
"organization_id": None,
|
||||
"max_budget": 300.0,
|
||||
}
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=mock_updated_team
|
||||
)
|
||||
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
http_request=dummy_request,
|
||||
|
|
@ -14968,6 +14924,49 @@ def _update_request_stub():
|
|||
return Mock(spec=Request)
|
||||
|
||||
|
||||
class _TeamRowStore:
|
||||
"""One team row whose writes honor their where clause, as Postgres does.
|
||||
|
||||
`budget_set_after_read` is a proxy admin's budget change that commits after update_team read the row."""
|
||||
|
||||
def __init__(self, table: MagicMock, row: dict[str, object], budget_set_after_read: float | None = None) -> None:
|
||||
self.row: Final = {
|
||||
"organization_id": None,
|
||||
"soft_budget": None,
|
||||
"model_id": None,
|
||||
"model_max_budget": None,
|
||||
"litellm_model_table": None,
|
||||
"metadata": {},
|
||||
**row,
|
||||
}
|
||||
self._budget_set_after_read = budget_set_after_read
|
||||
table.find_unique = self.find_unique
|
||||
table.update = self.update
|
||||
table.update_many = self.update_many
|
||||
|
||||
def _snapshot(self) -> MagicMock:
|
||||
snapshot: Final = MagicMock(**self.row)
|
||||
snapshot.model_dump.return_value = dict(self.row)
|
||||
return snapshot
|
||||
|
||||
async def find_unique(self, where, include=None):
|
||||
snapshot: Final = self._snapshot()
|
||||
if self._budget_set_after_read is not None:
|
||||
self.row["max_budget"] = self._budget_set_after_read
|
||||
self._budget_set_after_read = None
|
||||
return snapshot
|
||||
|
||||
async def update(self, where, data, include=None):
|
||||
self.row.update(data)
|
||||
return self._snapshot()
|
||||
|
||||
async def update_many(self, where, data):
|
||||
if any(self.row.get(column) != value for column, value in where.items()):
|
||||
return 0
|
||||
self.row.update(data)
|
||||
return 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_admin_is_refused_before_any_write_when_no_fields_are_enabled():
|
||||
import contextlib
|
||||
|
|
@ -15191,23 +15190,18 @@ async def test_update_team_stops_a_team_admin_raising_an_org_team_budget_under_t
|
|||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
|
||||
)
|
||||
org_team = MagicMock()
|
||||
org_team.metadata = {}
|
||||
org_team.organization_id = "budgeted-org"
|
||||
org_team.max_budget = 10.0
|
||||
org_team.model_max_budget = None
|
||||
org_team.model_dump.return_value = {
|
||||
"team_id": "test_team_id",
|
||||
"team_alias": "test_team",
|
||||
"organization_id": "budgeted-org",
|
||||
"max_budget": 10.0,
|
||||
"metadata": {},
|
||||
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
|
||||
}
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=org_team)
|
||||
store = _TeamRowStore(
|
||||
prisma.db.litellm_teamtable,
|
||||
{
|
||||
"team_id": "test_team_id",
|
||||
"team_alias": "test_team",
|
||||
"organization_id": "budgeted-org",
|
||||
"max_budget": 10.0,
|
||||
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
|
||||
},
|
||||
)
|
||||
stack.enter_context(_team_admin_may_edit("max_budget"))
|
||||
stack.enter_context(_not_org_admin())
|
||||
stack.enter_context(
|
||||
|
|
@ -15222,6 +15216,7 @@ async def test_update_team_stops_a_team_admin_raising_an_org_team_budget_under_t
|
|||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
budget_after_raise = store.row["max_budget"]
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", max_budget=5.0),
|
||||
http_request=_update_request_stub(),
|
||||
|
|
@ -15230,8 +15225,65 @@ async def test_update_team_stops_a_team_admin_raising_an_org_team_budget_under_t
|
|||
|
||||
assert str(raised.value.code) == "403"
|
||||
assert "Only a proxy admin can raise a team's max_budget" in str(raised.value.message)
|
||||
assert prisma.db.litellm_teamtable.update.await_count == 1
|
||||
assert prisma.db.litellm_teamtable.update.call_args.kwargs["data"]["max_budget"] == 5.0
|
||||
assert budget_after_raise == 10.0
|
||||
assert store.row["max_budget"] == 5.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("organization_id", "budget_read", "requested"),
|
||||
[
|
||||
pytest.param(None, 100.0, 90.0, id="lowering"),
|
||||
pytest.param(None, None, 90.0, id="first-budget"),
|
||||
pytest.param("budgeted-org", 100.0, 90.0, id="org-team"),
|
||||
],
|
||||
)
|
||||
async def test_update_team_keeps_a_budget_cut_that_lands_while_a_team_admin_update_runs(
|
||||
disable_audit_logging_for_mocked_team, organization_id, budget_read, requested
|
||||
):
|
||||
"""The team admin's check passed against the budget it read, which no longer holds once a proxy admin
|
||||
cut it to 20, so writing 90 would grow the team's live ceiling."""
|
||||
import contextlib
|
||||
|
||||
with contextlib.ExitStack() as stack:
|
||||
prisma = _wire_update_team(stack, {})
|
||||
store = _TeamRowStore(
|
||||
prisma.db.litellm_teamtable,
|
||||
{
|
||||
"team_id": "test_team_id",
|
||||
"team_alias": "test_team",
|
||||
"organization_id": organization_id,
|
||||
"max_budget": budget_read,
|
||||
"members_with_roles": [{"user_id": "team-admin", "role": "admin"}],
|
||||
},
|
||||
budget_set_after_read=20.0,
|
||||
)
|
||||
stack.enter_context(_team_admin_may_edit("max_budget"))
|
||||
stack.enter_context(_not_org_admin())
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: update_team reads orgs through this module-level import; no seam to inject
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_org_object",
|
||||
AsyncMock(
|
||||
return_value=LiteLLM_OrganizationTable(
|
||||
organization_id="budgeted-org",
|
||||
budget_id="budgeted-org-budget",
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1000.0),
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
with pytest.raises(ProxyException) as raised:
|
||||
await update_team(
|
||||
data=UpdateTeamRequest(team_id="test_team_id", max_budget=requested),
|
||||
http_request=_update_request_stub(),
|
||||
user_api_key_dict=_TEAM_ADMIN_CALLER,
|
||||
)
|
||||
|
||||
assert str(raised.value.code) == "409"
|
||||
assert "max_budget changed" in str(raised.value.message)
|
||||
assert store.row["max_budget"] == 20.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue