diff --git a/litellm/models/budget.py b/litellm/models/budget.py index 61123810fd1..2324479a13f 100644 --- a/litellm/models/budget.py +++ b/litellm/models/budget.py @@ -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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 67950e603c0..d43a18c795a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index b35b876b475..5dcd5b04684 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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) diff --git a/tests/integration/spend/test_team_member_spend.py b/tests/integration/spend/test_team_member_spend.py index cf1dac25793..bafe8ab777b 100644 --- a/tests/integration/spend/test_team_member_spend.py +++ b/tests/integration/spend/test_team_member_spend.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 30f5abdbb98..d372c931bcd 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 131db55ee01..1d2b0d6d584 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index d9e308e9d6f..f27c8493142 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -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 | 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; diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx index 660416504fe..e6c449908d7 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx @@ -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, 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["extraColumns"]> = [ diff --git a/ui/litellm-dashboard/src/components/team/memberBudget.test.ts b/ui/litellm-dashboard/src/components/team/memberBudget.test.ts new file mode 100644 index 00000000000..28a1fd757a0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/memberBudget.test.ts @@ -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); + }); +}); diff --git a/ui/litellm-dashboard/src/components/team/memberBudget.ts b/ui/litellm-dashboard/src/components/team/memberBudget.ts new file mode 100644 index 00000000000..7b599764954 --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/memberBudget.ts @@ -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; +};