mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix: preserve inherited member limits and resets after top-ups
This commit is contained in:
parent
4aa3ff47fe
commit
f71885c246
10 changed files with 577 additions and 5 deletions
|
|
@ -12,6 +12,17 @@ from pydantic import ConfigDict
|
|||
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
|
||||
TEMPORARY_BUDGET_INHERITED_FIELDS: Final = (
|
||||
"soft_budget",
|
||||
"max_budget",
|
||||
"max_parallel_requests",
|
||||
"tpm_limit",
|
||||
"rpm_limit",
|
||||
"tpd_limit",
|
||||
"model_max_budget",
|
||||
"budget_duration",
|
||||
)
|
||||
|
||||
|
||||
class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
|
||||
"""Represents user-controllable params for a LiteLLM_BudgetTable record.
|
||||
|
|
@ -36,6 +47,14 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
|
|||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
def is_temporary_only(self) -> bool:
|
||||
return (
|
||||
self.temp_budget_increase is not None
|
||||
and self.temp_budget_expiry is not None
|
||||
and not self.allowed_models
|
||||
and all(getattr(self, field) is None for field in TEMPORARY_BUDGET_INHERITED_FIELDS)
|
||||
)
|
||||
|
||||
def active_temp_budget_increase(self, now: datetime) -> float:
|
||||
if self.temp_budget_increase is None or self.temp_budget_expiry is None:
|
||||
return 0.0
|
||||
|
|
|
|||
|
|
@ -991,6 +991,14 @@ async def common_checks(
|
|||
if team_object is not None and membership_user_id is not None
|
||||
else None
|
||||
)
|
||||
if team_membership_loaded:
|
||||
await _inherit_team_member_rate_limits(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
membership=loaded_team_membership,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
unpriced_models: Final = (
|
||||
_unpriced_models_in_request(model=_model, llm_router=llm_router)
|
||||
|
|
@ -5597,6 +5605,34 @@ async def _virtual_key_max_budget_alert_check(
|
|||
)
|
||||
|
||||
|
||||
async def _inherit_team_member_rate_limits(
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
membership: LiteLLM_TeamMembership | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> None:
|
||||
if valid_token is None or valid_token.user_id is None or team_object is None:
|
||||
return
|
||||
member_budget: Final = membership.litellm_budget_table if membership is not None else None
|
||||
if member_budget is not None and not member_budget.is_temporary_only():
|
||||
return
|
||||
default_budget_id: Final = (
|
||||
team_object.metadata.get("team_member_budget_id") if team_object.metadata is not None else None
|
||||
)
|
||||
if not isinstance(default_budget_id, str):
|
||||
return
|
||||
default_budget: Final = await get_team_member_default_budget(
|
||||
budget_id=default_budget_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
if default_budget is None:
|
||||
return
|
||||
valid_token.team_member_rpm_limit = default_budget.rpm_limit # rebind-ok: pin resolved limits on the request
|
||||
valid_token.team_member_tpm_limit = default_budget.tpm_limit # rebind-ok: pin resolved limits on the request
|
||||
|
||||
|
||||
async def _check_team_member_budget(
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
user_object: LiteLLM_UserTable | None,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,9 @@ from enum import Enum
|
|||
from types import MappingProxyType
|
||||
from typing import Final, Generic, Literal, Protocol, TypeVar
|
||||
|
||||
from typing_extensions import assert_never
|
||||
from prisma import Json
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -23,6 +25,7 @@ from litellm.constants import (
|
|||
RESET_BUDGET_JOB_NAME,
|
||||
)
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.models.budget import TEMPORARY_BUDGET_INHERITED_FIELDS
|
||||
from litellm.proxy._types import (
|
||||
DB_RETRY_SAFE_ERROR_TYPES,
|
||||
LiteLLM_BudgetTableFull,
|
||||
|
|
@ -279,6 +282,80 @@ class _BudgetCascade:
|
|||
counter_resets: tuple[tuple[str, float], ...] = ()
|
||||
cache_keys: tuple[str, ...] = ()
|
||||
rollover_caps: Mapping[str, float] = field(default_factory=lambda: MappingProxyType({}))
|
||||
inherited_members: tuple["_InheritedMemberReset", ...] = ()
|
||||
|
||||
|
||||
class _TeamDefaultMetadata(BaseModel):
|
||||
team_member_budget_id: str | None = None
|
||||
|
||||
|
||||
class _TeamDefaultLink(BaseModel):
|
||||
team_id: str
|
||||
metadata: _TeamDefaultMetadata | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _InheritedMemberReset:
|
||||
budget_id: str
|
||||
where: "_InheritedMemberWhere"
|
||||
members: tuple[_TeamMembershipRow, ...]
|
||||
|
||||
|
||||
class _InheritedMemberWhere(TypedDict):
|
||||
team_id: ReadOnly[Mapping[str, Sequence[str]]]
|
||||
OR: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
spend: NotRequired[ReadOnly[Mapping[str, float]]]
|
||||
|
||||
|
||||
class _TeamDefaultWhere(TypedDict):
|
||||
metadata: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
def _team_default_where(budget_id: str) -> _TeamDefaultWhere:
|
||||
where: Final[_TeamDefaultWhere] = {"metadata": {"path": ("team_member_budget_id",), "equals": Json(budget_id)}}
|
||||
return where
|
||||
|
||||
|
||||
class _TeamDefaultsWhere(TypedDict):
|
||||
OR: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
|
||||
|
||||
def _inherited_member_where(team_ids: Sequence[str]) -> _InheritedMemberWhere:
|
||||
where: Final[_InheritedMemberWhere] = {
|
||||
"team_id": {"in": team_ids},
|
||||
"OR": (
|
||||
{"budget_id": None},
|
||||
{
|
||||
"litellm_budget_table": {
|
||||
"is": {
|
||||
**MappingProxyType({field: None for field in TEMPORARY_BUDGET_INHERITED_FIELDS}),
|
||||
"model_max_budget": {"equals": "AnyNull"},
|
||||
"temp_budget_increase": {"not": None},
|
||||
"temp_budget_expiry": {"not": None},
|
||||
"allowed_models": {"equals": ()},
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
}
|
||||
return where
|
||||
|
||||
|
||||
def _queue_inherited_member_reset(
|
||||
writes: LinkedSpendResetWrites, inherited: _InheritedMemberReset, cap: float | None
|
||||
) -> None:
|
||||
if cap is None:
|
||||
writes.queue_spend_zero(where=inherited.where)
|
||||
return
|
||||
under_cap: Final[_InheritedMemberWhere] = {**inherited.where, "spend": {"gt": 0, "lte": cap}}
|
||||
over_cap: Final[_InheritedMemberWhere] = {**inherited.where, "spend": {"gt": cap}}
|
||||
writes.queue_spend_zero(where=under_cap)
|
||||
writes.queue_spend_decrement(where=over_cap, amount=cap)
|
||||
|
||||
|
||||
def _queue_inherited_member_resets(writes: LinkedSpendResetWrites, cascade: _BudgetCascade) -> None:
|
||||
for inherited in cascade.inherited_members:
|
||||
_queue_inherited_member_reset(writes, inherited, cascade.rollover_caps.get(inherited.budget_id))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -729,6 +806,42 @@ class ResetBudgetJob:
|
|||
)
|
||||
)
|
||||
|
||||
async def _collect_inherited_member_resets(self, budget_ids: Sequence[str]) -> tuple[_InheritedMemberReset, ...]:
|
||||
filters: Final = tuple(_team_default_where(budget_id) for budget_id in budget_ids)
|
||||
where: Final[_TeamDefaultsWhere] = {"OR": filters}
|
||||
teams: Final = await self._with_db_retry(
|
||||
lambda: TeamRepository(self.prisma_client).table.find_many(where=where),
|
||||
reason="reset_budget_read_team_defaults_failure",
|
||||
)
|
||||
links: Final = tuple(_TeamDefaultLink.model_validate(team, from_attributes=True) for team in teams)
|
||||
groups: Final = tuple(
|
||||
(
|
||||
budget_id,
|
||||
tuple(
|
||||
link.team_id
|
||||
for link in links
|
||||
if link.metadata is not None and link.metadata.team_member_budget_id == budget_id
|
||||
),
|
||||
)
|
||||
for budget_id in budget_ids
|
||||
)
|
||||
return tuple(
|
||||
[
|
||||
_InheritedMemberReset(
|
||||
budget_id=budget_id,
|
||||
where=member_where,
|
||||
members=await self._fetch_linked_rows(
|
||||
table=TeamMembershipRepository(self.prisma_client).table,
|
||||
where=member_where,
|
||||
log_subject="inherited team memberships",
|
||||
),
|
||||
)
|
||||
for budget_id, team_ids in groups
|
||||
for offset in range(0, len(team_ids), RESET_BUDGET_JOB_BATCH_SIZE)
|
||||
for member_where in (_inherited_member_where(team_ids[offset : offset + RESET_BUDGET_JOB_BATCH_SIZE]),)
|
||||
]
|
||||
)
|
||||
|
||||
async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade:
|
||||
"""Resolve every row the expiring budget tiers gate, before any write.
|
||||
|
||||
|
|
@ -740,6 +853,7 @@ class ResetBudgetJob:
|
|||
if not budget_ids:
|
||||
return _EMPTY_CASCADE
|
||||
|
||||
inherited_members: Final = await self._collect_inherited_member_resets(budget_ids)
|
||||
team_memberships: Final[tuple[_TeamMembershipRow, ...]] = await self._fetch_linked_rows(
|
||||
table=TeamMembershipRepository(self.prisma_client).table,
|
||||
where=_budget_link_where(budget_ids),
|
||||
|
|
@ -791,6 +905,11 @@ class ResetBudgetJob:
|
|||
if b.budget_id is not None and b.budget_duration is not None
|
||||
),
|
||||
counter_resets=(
|
||||
*(
|
||||
(_team_membership_counter_key(row), _carried_spend(row.spend, rollover_caps.get(group.budget_id)))
|
||||
for group in inherited_members
|
||||
for row in group.members
|
||||
),
|
||||
*(
|
||||
(_team_membership_counter_key(row), _row_carried_spend(row, rollover_caps))
|
||||
for row in team_memberships
|
||||
|
|
@ -805,7 +924,14 @@ class ResetBudgetJob:
|
|||
*((_project_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in projects),
|
||||
),
|
||||
rollover_caps=rollover_caps,
|
||||
inherited_members=inherited_members,
|
||||
cache_keys=(
|
||||
*(
|
||||
key
|
||||
for group in inherited_members
|
||||
for row in group.members
|
||||
for key in _team_membership_cache_keys(row)
|
||||
),
|
||||
*(key for row in team_memberships for key in _team_membership_cache_keys(row)),
|
||||
*(key for row in keys for key in _key_cache_keys(row)),
|
||||
*(key for row in orgs for key in _org_cache_keys(row)),
|
||||
|
|
@ -834,6 +960,7 @@ class ResetBudgetJob:
|
|||
async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None:
|
||||
async with budget_cascade_unit_of_work(self._new_batch) as uow:
|
||||
_queue_budget_linked_resets(uow.team_memberships, cascade)
|
||||
_queue_inherited_member_resets(uow.team_memberships, cascade)
|
||||
_queue_budget_linked_resets(uow.keys, cascade, extra=_LINKED_KEYS_WHERE)
|
||||
_queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE)
|
||||
_queue_budget_linked_resets(uow.tags, cascade, extra=_SPENT_ROWS_WHERE)
|
||||
|
|
|
|||
|
|
@ -97,3 +97,115 @@ def test_member_over_budget_is_blocked_when_redis_counter_reads_stale_low(gatewa
|
|||
assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text
|
||||
assert upstream.get("/__observations").json()["requests"] == []
|
||||
assert float(cache.get(counter_key)) == pytest.approx(0.06), denied.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize("grant_state", ["active", "expired", "cleared"])
|
||||
def test_temporary_topup_resets_with_team_default_and_preserves_private_budget(
|
||||
gateway: Gateway, grant_state: str
|
||||
) -> None:
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import psycopg
|
||||
|
||||
from integration._support.client import string_value
|
||||
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
team: Final = scenario.team(team_member_budget=0.05, team_member_budget_duration="1d", models=[model])
|
||||
users: Final = tuple(scenario.user() for _ in range(4))
|
||||
for user in users:
|
||||
gateway.post("/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}})
|
||||
keys: Final = tuple(scenario.key(team_id=team, user_id=user) for user in users)
|
||||
for key in keys:
|
||||
gateway.chat(model, key=key)
|
||||
spent: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT user_id, spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=ANY(%s)',
|
||||
(team, list(users)),
|
||||
),
|
||||
lambda rows: len(rows) == 4 and all(float(row["spend"]) >= 0.06 for row in rows),
|
||||
seconds=70,
|
||||
)
|
||||
assert all(float(row["spend"]) == pytest.approx(0.06) for row in spent), spent
|
||||
for user, cap, duration in ((users[1], 0.05, None), (users[2], 0, None), (users[3], 0.05, "1d")):
|
||||
gateway.post(
|
||||
"/team/member_update",
|
||||
{
|
||||
"team_id": team,
|
||||
"user_id": user,
|
||||
"max_budget_in_team": cap,
|
||||
"budget_duration": duration,
|
||||
},
|
||||
)
|
||||
gateway.post(
|
||||
"/team/member_update",
|
||||
{
|
||||
"team_id": team,
|
||||
"user_id": users[0],
|
||||
"temp_budget_increase": 0.001,
|
||||
"temp_budget_expiry": (datetime.now(timezone.utc) + timedelta(hours=1)).isoformat(),
|
||||
},
|
||||
)
|
||||
if grant_state == "cleared":
|
||||
gateway.post(
|
||||
"/team/member_update",
|
||||
{
|
||||
"team_id": team,
|
||||
"user_id": users[0],
|
||||
"temp_budget_increase": None,
|
||||
"temp_budget_expiry": None,
|
||||
},
|
||||
)
|
||||
team_info: Final = object_value(gateway.get("/team/info", {"team_id": team})["team_info"])
|
||||
default_id: Final = string_value(object_value(team_info["metadata"])["team_member_budget_id"])
|
||||
denied_before: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "before reset"}],
|
||||
},
|
||||
key=keys[0],
|
||||
)
|
||||
assert denied_before.status_code == 422, denied_before.text
|
||||
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
|
||||
if grant_state == "expired":
|
||||
connection.execute(
|
||||
"UPDATE \"LiteLLM_BudgetTable\" SET temp_budget_expiry=now() - interval '1 day' "
|
||||
'WHERE budget_id=(SELECT budget_id FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s)',
|
||||
(team, users[0]),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE \"LiteLLM_BudgetTable\" SET budget_reset_at=now() - interval '1 day' WHERE budget_id=%s",
|
||||
(default_id,),
|
||||
)
|
||||
reset: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend, total_spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s',
|
||||
(team, users[0]),
|
||||
),
|
||||
lambda rows: len(rows) == 1 and float(rows[0]["spend"]) == 0,
|
||||
seconds=70,
|
||||
)
|
||||
assert float(reset[0]["total_spend"]) == pytest.approx(0.06), reset
|
||||
allowed: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "after reset"}],
|
||||
},
|
||||
key=keys[0],
|
||||
)
|
||||
assert allowed.status_code == 200, allowed.text
|
||||
for key in keys[1:]:
|
||||
denied: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "private budget did not reset"}],
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text
|
||||
|
|
|
|||
|
|
@ -9456,3 +9456,132 @@ def test_can_object_call_model_allows_listed_model_for_key():
|
|||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("expired", [False, True])
|
||||
@pytest.mark.parametrize("has_membership", [False, True])
|
||||
async def test_temporary_topup_keeps_live_default_rate_limits(expired: bool, has_membership: bool) -> None:
|
||||
from litellm.proxy.auth.auth_checks import _inherit_team_member_rate_limits
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
|
||||
cache: Final = UserApiKeyCache()
|
||||
team: Final = LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "default"})
|
||||
grant: Final = LiteLLM_BudgetTable(
|
||||
temp_budget_increase=200,
|
||||
temp_budget_expiry=datetime(2020 if expired else 2100, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
membership: Final = (
|
||||
LiteLLM_TeamMembership(user_id="member", team_id="team", litellm_budget_table=grant) if has_membership else None
|
||||
)
|
||||
for rpm, tpm in ((2, 100), (5, 200)):
|
||||
await cache.async_set_cache(
|
||||
key="team_member_default_budget:default",
|
||||
value=LiteLLM_BudgetTable(rpm_limit=rpm, tpm_limit=tpm),
|
||||
model_type=LiteLLM_BudgetTable,
|
||||
)
|
||||
token: Final = UserAPIKeyAuth(user_id="member", team_id="team")
|
||||
await _inherit_team_member_rate_limits(token, team, membership, MagicMock(), cache)
|
||||
assert (token.team_member_rpm_limit, token.team_member_tpm_limit) == (rpm, tpm)
|
||||
assert (grant.rpm_limit, grant.tpm_limit) == (None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"permanent",
|
||||
[
|
||||
{"max_budget": 0},
|
||||
{"max_budget": 500},
|
||||
{"rpm_limit": 1},
|
||||
{"tpm_limit": 100},
|
||||
{"tpd_limit": 100},
|
||||
{"soft_budget": 50},
|
||||
{"max_parallel_requests": 1},
|
||||
{"budget_duration": "1d"},
|
||||
{"model_max_budget": {"model": 1}},
|
||||
{"allowed_models": ["model"]},
|
||||
],
|
||||
)
|
||||
async def test_topup_on_explicit_private_budget_does_not_inherit_limits(permanent: dict[str, object]) -> None:
|
||||
from litellm.proxy.auth.auth_checks import _inherit_team_member_rate_limits
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
|
||||
budget: Final = LiteLLM_BudgetTable.model_validate(
|
||||
{
|
||||
"temp_budget_increase": 200,
|
||||
"temp_budget_expiry": datetime(2100, 1, 1, tzinfo=timezone.utc),
|
||||
**permanent,
|
||||
}
|
||||
)
|
||||
token: Final = UserAPIKeyAuth(user_id="member", team_id="team", team_member_rpm_limit=7)
|
||||
membership: Final = LiteLLM_TeamMembership(user_id="member", team_id="team", litellm_budget_table=budget)
|
||||
cache: Final = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="team_member_default_budget:default",
|
||||
value=LiteLLM_BudgetTable(rpm_limit=2, tpm_limit=100),
|
||||
model_type=LiteLLM_BudgetTable,
|
||||
)
|
||||
await _inherit_team_member_rate_limits(
|
||||
token,
|
||||
LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "default"}),
|
||||
membership,
|
||||
MagicMock(),
|
||||
cache,
|
||||
)
|
||||
assert (token.team_member_rpm_limit, token.team_member_tpm_limit) == (7, None)
|
||||
assert not budget.is_temporary_only()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/responses", "/v1/messages"])
|
||||
@pytest.mark.parametrize("token_value", [None, "member-key"])
|
||||
async def test_common_checks_resolves_temporary_member_limits_for_key_and_jwt_context(
|
||||
monkeypatch: pytest.MonkeyPatch, route: str, token_value: str | None
|
||||
) -> None:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key
|
||||
|
||||
cache: Final = UserApiKeyCache()
|
||||
membership: Final = LiteLLM_TeamMembership(
|
||||
user_id="member",
|
||||
team_id="team",
|
||||
budget_id="temporary",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(
|
||||
temp_budget_increase=200,
|
||||
temp_budget_expiry=datetime(2100, 1, 1, tzinfo=timezone.utc),
|
||||
),
|
||||
)
|
||||
await cache.async_set_cache(
|
||||
key=team_membership_reservation_cache_key(user_id="member", team_id="team"),
|
||||
value=membership,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
await cache.async_set_cache(
|
||||
key="team_member_default_budget:default",
|
||||
value=LiteLLM_BudgetTable(rpm_limit=2, tpm_limit=100),
|
||||
model_type=LiteLLM_BudgetTable,
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
|
||||
token: Final = UserAPIKeyAuth(token=token_value, user_id="member", team_id="team")
|
||||
result: Final = await common_checks(
|
||||
request_body={"model": "test-model"},
|
||||
team_object=LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "default"}),
|
||||
user_object=None,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings={},
|
||||
route=route,
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=token,
|
||||
request=Request({"type": "http", "headers": [], "path": route}),
|
||||
skip_budget_checks=True,
|
||||
)
|
||||
assert result is True
|
||||
assert (token.team_member_rpm_limit, token.team_member_tpm_limit) == (2, 100)
|
||||
assert (membership.litellm_budget_table.rpm_limit, membership.litellm_budget_table.tpm_limit) == (None, None)
|
||||
|
|
|
|||
|
|
@ -57,6 +57,8 @@ class MockTable:
|
|||
self.find_many_calls.append({"where": where, **paging})
|
||||
rows = list(self._find_many_results)
|
||||
for field, condition in where.items():
|
||||
if field == "budget_id" and isinstance(condition, dict) and "in" in condition:
|
||||
rows = [row for row in rows if not hasattr(row, "budget_id") or row.budget_id in condition["in"]]
|
||||
if isinstance(condition, dict) and "gt" in condition and field != "spend":
|
||||
rows = [row for row in rows if getattr(row, field, "") > condition["gt"]]
|
||||
for field, direction in (order or {}).items():
|
||||
|
|
@ -112,6 +114,7 @@ class MockBatcher:
|
|||
|
||||
class MockDB:
|
||||
def __init__(self):
|
||||
self.litellm_teamtable = MockTable()
|
||||
self.litellm_teammembership = MockTable()
|
||||
self.litellm_verificationtoken = MockTable()
|
||||
self.litellm_endusertable = MockTable()
|
||||
|
|
@ -3578,3 +3581,51 @@ def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_p
|
|||
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol")
|
||||
counter_cache.in_memory_cache.set_cache.assert_not_called()
|
||||
counter_cache.redis_cache.async_set_cache.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inherited_member_rollover_uses_default_cap_and_invalidates_private_link(
|
||||
rollover_enabled: None, reset_budget_job: ResetBudgetJob, mock_prisma_client: MockPrismaClient
|
||||
) -> None:
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
from litellm.models.team import LiteLLM_TeamTable
|
||||
|
||||
budget: Final = _budget_row(budget_id="default", max_budget=10.0)
|
||||
mock_prisma_client.db.litellm_teamtable.set_find_many_results(
|
||||
[LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "default"})]
|
||||
)
|
||||
member: Final = LiteLLM_TeamMembership(user_id="member", team_id="team", budget_id="overlay", spend=15.0)
|
||||
mock_prisma_client.db.litellm_teammembership.set_find_many_results([member])
|
||||
cascade: Final = await reset_budget_job._collect_budget_cascade([budget])
|
||||
assert ("spend:team_member:member:team", 5.0) in cascade.counter_resets
|
||||
assert f"{member.team_id}_{member.user_id}" in cascade.cache_keys
|
||||
batch: Final = MockBatcher()
|
||||
reset_budget_job_module._queue_inherited_member_resets(
|
||||
reset_budget_job_module.LinkedSpendResetWrites(batch.litellm_teammembership), cascade
|
||||
)
|
||||
assert _replay_spend_writes(batch.calls, 15.0) == 5.0
|
||||
assert _replay_spend_writes(batch.calls, 8.0) == 0.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("db_factory,commits", [(MockDB, True), (FailingCommitDB, False)])
|
||||
async def test_inherited_member_cache_eviction_requires_successful_reset(
|
||||
db_factory: type[MockDB], commits: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
from litellm.models.team import LiteLLM_TeamTable
|
||||
|
||||
cache: Final = _make_counter_invalidation_job(monkeypatch)
|
||||
job, client = _job_with_expired_budget(db_factory())
|
||||
client.db.litellm_teamtable.set_find_many_results([
|
||||
LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "budget-1"})
|
||||
])
|
||||
client.db.litellm_teammembership.set_find_many_results([
|
||||
LiteLLM_TeamMembership(user_id="member", team_id="team", budget_id="overlay", spend=15.0)
|
||||
])
|
||||
await job.reset_budget_for_litellm_budget_table()
|
||||
deleted: Final = {call.kwargs["key"] for call in cache.in_memory_cache.delete_cache.call_args_list}
|
||||
evicted: Final = {call.kwargs["key"] for call in cache.user_api_key_cache.async_delete_cache.await_args_list}
|
||||
assert ("spend:team_member:member:team" in deleted) == commits
|
||||
assert ("team_member" in evicted) == commits
|
||||
assert client.db.batchers[0].committed == commits
|
||||
|
|
|
|||
|
|
@ -264,6 +264,7 @@ export interface TeamMembership {
|
|||
max_parallel_requests: number | null;
|
||||
tpm_limit: number | null;
|
||||
rpm_limit: number | null;
|
||||
tpd_limit?: number | null;
|
||||
model_max_budget: Record<string, number> | null;
|
||||
budget_duration: string | null;
|
||||
budget_reset_at: string | null;
|
||||
|
|
@ -315,6 +316,7 @@ export interface TeamData {
|
|||
team_member_budget_table: {
|
||||
max_budget: number;
|
||||
budget_duration: string | null;
|
||||
budget_reset_at?: string | null;
|
||||
tpm_limit: number | null;
|
||||
rpm_limit: number | null;
|
||||
} | null;
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ import { isProxyAdminRole, isUserTeamAdminForSingleTeam } from "@/utils/roles";
|
|||
import { CircleHelp } from "lucide-react";
|
||||
import { useState, type ComponentProps } from "react";
|
||||
import { TeamData, TeamMemberBudgetSource, TeamMembership } from "./TeamInfo";
|
||||
import { displayedMemberBudget } from "./memberBudget";
|
||||
|
||||
const BUDGET_SOURCE_LABELS: Record<Exclude<TeamMemberBudgetSource, "none">, string> = {
|
||||
team_default: "Team default",
|
||||
|
|
@ -100,6 +101,8 @@ export default function TeamMemberTab({
|
|||
const getUserBudgetSource = (userId: string | null): TeamMemberBudgetSource => {
|
||||
if (!userId) return "none";
|
||||
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
|
||||
const teamDefault = teamData.team_info.team_member_budget_table;
|
||||
if (teamDefault && displayedMemberBudget(membership, teamDefault) === teamDefault) return "team_default";
|
||||
return membership?.budget_source ?? "none";
|
||||
};
|
||||
|
||||
|
|
@ -107,7 +110,7 @@ export default function TeamMemberTab({
|
|||
if (!userId) return null;
|
||||
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
|
||||
return (
|
||||
membership?.litellm_budget_table?.max_budget ??
|
||||
displayedMemberBudget(membership, teamData.team_info.team_member_budget_table)?.max_budget ??
|
||||
(membership?.budget_source === "team_default" ? teamDefaultBudget : null)
|
||||
);
|
||||
};
|
||||
|
|
@ -116,8 +119,9 @@ export default function TeamMemberTab({
|
|||
const getUserRateLimits = (userId: string | null): string => {
|
||||
if (!userId) return "No Limits";
|
||||
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
|
||||
const rpmLimit = membership?.litellm_budget_table?.rpm_limit;
|
||||
const tpmLimit = membership?.litellm_budget_table?.tpm_limit;
|
||||
const budget = displayedMemberBudget(membership, teamData.team_info.team_member_budget_table);
|
||||
const rpmLimit = budget?.rpm_limit;
|
||||
const tpmLimit = budget?.tpm_limit;
|
||||
|
||||
const rpmText = rpmLimit != null ? `${formatNumber(rpmLimit)} RPM` : null;
|
||||
const tpmText = tpmLimit != null ? `${formatNumber(tpmLimit)} TPM` : null;
|
||||
|
|
@ -142,7 +146,7 @@ export default function TeamMemberTab({
|
|||
const getUserBudgetReset = (userId: string | null): string | null => {
|
||||
if (!userId) return null;
|
||||
const membership = teamData.team_memberships.find((tm) => tm.user_id === userId);
|
||||
return membership?.litellm_budget_table?.budget_reset_at ?? null;
|
||||
return displayedMemberBudget(membership, teamData.team_info.team_member_budget_table)?.budget_reset_at ?? null;
|
||||
};
|
||||
|
||||
const extraColumns: NonNullable<ComponentProps<typeof MemberTable>["extraColumns"]> = [
|
||||
|
|
|
|||
|
|
@ -0,0 +1,66 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import type { TeamMembership } from "./TeamInfo";
|
||||
import { displayedMemberBudget } from "./memberBudget";
|
||||
|
||||
const teamDefault = {
|
||||
max_budget: 1000,
|
||||
budget_duration: "1mo",
|
||||
budget_reset_at: "2100-02-01T00:00:00Z",
|
||||
rpm_limit: 2,
|
||||
tpm_limit: 100,
|
||||
};
|
||||
const membership: TeamMembership = {
|
||||
user_id: "member",
|
||||
team_id: "team",
|
||||
budget_id: "temporary",
|
||||
budget_source: "custom",
|
||||
spend: 1100,
|
||||
total_spend: 1100,
|
||||
litellm_budget_table: {
|
||||
budget_id: "temporary",
|
||||
max_budget: null,
|
||||
soft_budget: null,
|
||||
max_parallel_requests: null,
|
||||
rpm_limit: null,
|
||||
tpm_limit: null,
|
||||
model_max_budget: null,
|
||||
budget_duration: null,
|
||||
budget_reset_at: null,
|
||||
temp_budget_increase: 200,
|
||||
temp_budget_expiry: "2100-02-01T00:00:00Z",
|
||||
},
|
||||
};
|
||||
|
||||
describe("displayedMemberBudget", () => {
|
||||
it.each(["2020-02-01T00:00:00Z", "2100-02-01T00:00:00Z"])(
|
||||
"shows inherited rates, cap and reset for a grant expiring at %s without changing edit values",
|
||||
(expiry) => {
|
||||
const member = {
|
||||
...membership,
|
||||
litellm_budget_table: { ...membership.litellm_budget_table, temp_budget_expiry: expiry },
|
||||
};
|
||||
expect(displayedMemberBudget(member, teamDefault)).toEqual(teamDefault);
|
||||
expect(member.litellm_budget_table.max_budget).toBeNull();
|
||||
expect(member.litellm_budget_table.budget_duration).toBeNull();
|
||||
},
|
||||
);
|
||||
|
||||
it.each([
|
||||
{ max_budget: 0 },
|
||||
{ rpm_limit: 1 },
|
||||
{ tpm_limit: 50 },
|
||||
{ tpd_limit: 10 },
|
||||
{ soft_budget: 50 },
|
||||
{ max_parallel_requests: 1 },
|
||||
{ budget_duration: "1d" },
|
||||
{ model_max_budget: { model: 1 } },
|
||||
{ allowed_models: ["model"] },
|
||||
{ temp_budget_increase: null },
|
||||
])("keeps an explicit member budget independent: %j", (override) => {
|
||||
const member = {
|
||||
...membership,
|
||||
litellm_budget_table: { ...membership.litellm_budget_table, ...override },
|
||||
};
|
||||
expect(displayedMemberBudget(member, teamDefault)).toBe(member.litellm_budget_table);
|
||||
});
|
||||
});
|
||||
26
ui/litellm-dashboard/src/components/team/memberBudget.ts
Normal file
26
ui/litellm-dashboard/src/components/team/memberBudget.ts
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
import type { TeamData, TeamMembership } from "./TeamInfo";
|
||||
|
||||
export const displayedMemberBudget = (
|
||||
membership: TeamMembership | undefined,
|
||||
teamDefault: TeamData["team_info"]["team_member_budget_table"],
|
||||
) => {
|
||||
const budget = membership?.litellm_budget_table;
|
||||
const hasTemporaryGrant = budget?.temp_budget_increase != null && budget.temp_budget_expiry != null;
|
||||
const temporaryOnly =
|
||||
hasTemporaryGrant &&
|
||||
!budget.allowed_models?.length &&
|
||||
[
|
||||
budget.max_budget,
|
||||
budget.soft_budget,
|
||||
budget.max_parallel_requests,
|
||||
budget.rpm_limit,
|
||||
budget.tpm_limit,
|
||||
budget.tpd_limit,
|
||||
budget.model_max_budget,
|
||||
budget.budget_duration,
|
||||
].every((value) => value == null);
|
||||
if (temporaryOnly || (!budget && membership?.budget_source === "team_default")) {
|
||||
return teamDefault;
|
||||
}
|
||||
return budget;
|
||||
};
|
||||
Loading…
Add table
Reference in a new issue