fix: preserve inherited member limits and resets after top-ups

This commit is contained in:
moe-berri 2026-09-24 15:45:02 -07:00
parent 4aa3ff47fe
commit f71885c246
10 changed files with 577 additions and 5 deletions

View file

@ -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

View file

@ -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,

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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;

View file

@ -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"]> = [

View file

@ -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);
});
});

View 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;
};