feat(proxy): temporary budget increase for team members

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-17 17:33:50 +00:00
parent 4b368bf066
commit 7f3f8fae2d
10 changed files with 198 additions and 1 deletions

View file

@ -0,0 +1,5 @@
-- AlterTable
ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN IF NOT EXISTS "temp_budget_increase" DOUBLE PRECISION;
-- AlterTable
ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN IF NOT EXISTS "temp_budget_expiry" TIMESTAMP(3);

View file

@ -22,6 +22,8 @@ model LiteLLM_BudgetTable {
budget_duration String?
budget_reset_at DateTime?
allowed_models String[] @default([]) // per-member model scope; empty = inherit team models
temp_budget_increase Float?
temp_budget_expiry DateTime?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -30,6 +30,8 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
model_max_budget: dict | None = None
budget_duration: str | None = None
allowed_models: list[str] | None = None # per-member model scope; empty = inherit team models
temp_budget_increase: float | None = None
temp_budget_expiry: datetime | None = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -4397,6 +4397,21 @@ class TeamMemberUpdateRequest(TeamMemberDeleteRequest):
default=None,
description="List of models this team member can access. Pass an empty list to remove per-member model restrictions.",
)
temp_budget_increase: float | None = Field(
default=None,
description="Temporary additive budget increase for this team member, active until temp_budget_expiry",
)
temp_budget_expiry: datetime | None = Field(
default=None,
description="UTC expiry for temp_budget_increase",
)
@model_validator(mode="after")
def validate_temp_budget(self) -> "TeamMemberUpdateRequest":
if self.temp_budget_increase is not None or self.temp_budget_expiry is not None:
if self.temp_budget_increase is None or self.temp_budget_expiry is None:
raise ValueError("temp_budget_increase and temp_budget_expiry must be set together")
return self
class TeamMemberUpdateResponse(MemberUpdateResponse):
@ -4406,6 +4421,8 @@ class TeamMemberUpdateResponse(MemberUpdateResponse):
rpm_limit: int | None = None
budget_duration: str | None = None
allowed_models: list[str] | None = None
temp_budget_increase: float | None = None
temp_budget_expiry: datetime | None = None
class TeamModelAddRequest(BaseModel):

View file

@ -14,6 +14,7 @@ import math
import re
import time
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
@ -5295,6 +5296,24 @@ async def _virtual_key_max_budget_alert_check(
)
def _effective_team_member_budget(budget: LiteLLM_BudgetTable, now: datetime) -> float | None:
"""Per-member cap including an unexpired temp_budget_increase. Naive
temp_budget_expiry values are treated as UTC (same convention as
_get_temp_budget_increase for keys)."""
if budget.max_budget is None:
return None
if budget.temp_budget_increase is None or budget.temp_budget_expiry is None:
return budget.max_budget
expiry: Final = (
budget.temp_budget_expiry.replace(tzinfo=timezone.utc)
if budget.temp_budget_expiry.tzinfo is None
else budget.temp_budget_expiry
)
if expiry <= now:
return budget.max_budget
return budget.max_budget + budget.temp_budget_increase
async def _check_team_member_budget(
team_object: LiteLLM_TeamTable | None,
user_object: LiteLLM_UserTable | None,
@ -5330,7 +5349,10 @@ async def _check_team_member_budget(
and loaded_membership.litellm_budget_table is not None
and loaded_membership.litellm_budget_table.max_budget is not None
):
team_member_budget = loaded_membership.litellm_budget_table.max_budget
team_member_budget = _effective_team_member_budget(
loaded_membership.litellm_budget_table,
now=get_utc_datetime(),
)
else:
default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id")
if isinstance(default_budget_id, str):

View file

@ -3692,6 +3692,8 @@ _MEMBER_BUDGET_PATCH_FIELDS: Final = {
"rpm_limit": "rpm_limit",
"budget_duration": "budget_duration",
"allowed_models": "allowed_models",
"temp_budget_increase": "temp_budget_increase",
"temp_budget_expiry": "temp_budget_expiry",
}
@ -3862,6 +3864,8 @@ async def team_member_update(
rpm_limit=data.rpm_limit,
budget_duration=data.budget_duration,
allowed_models=data.allowed_models,
temp_budget_increase=data.temp_budget_increase,
temp_budget_expiry=data.temp_budget_expiry,
)

View file

@ -22,6 +22,8 @@ model LiteLLM_BudgetTable {
budget_duration String?
budget_reset_at DateTime?
allowed_models String[] @default([]) // per-member model scope; empty = inherit team models
temp_budget_increase Float?
temp_budget_expiry DateTime?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -22,6 +22,8 @@ model LiteLLM_BudgetTable {
budget_duration String?
budget_reset_at DateTime?
allowed_models String[] @default([]) // per-member model scope; empty = inherit team models
temp_budget_increase Float?
temp_budget_expiry DateTime?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -8461,3 +8461,121 @@ def test_route_skips_budget_checks_marks_only_spend_free_routes() -> None:
def test_request_skips_budget_checks_extends_route_rule_with_zero_cost_models() -> None:
assert request_skips_budget_checks(route="/v1/models", model=None, llm_router=None) is True
assert request_skips_budget_checks(route="/v1/chat/completions", model=None, llm_router=None) is False
def test_effective_team_member_budget_applies_unexpired_increase() -> None:
from litellm.proxy.auth.auth_checks import _effective_team_member_budget
budget: Final = LiteLLM_BudgetTable(
max_budget=100.0,
temp_budget_increase=50.0,
temp_budget_expiry=datetime(2100, 1, 1),
)
assert _effective_team_member_budget(budget, now=datetime(2026, 1, 1, tzinfo=timezone.utc)) == 150.0
def test_effective_team_member_budget_ignores_expired_increase() -> None:
from litellm.proxy.auth.auth_checks import _effective_team_member_budget
budget: Final = LiteLLM_BudgetTable(
max_budget=100.0,
temp_budget_increase=50.0,
temp_budget_expiry=datetime(2020, 1, 1, tzinfo=timezone.utc),
)
assert _effective_team_member_budget(budget, now=datetime(2026, 1, 1, tzinfo=timezone.utc)) == 100.0
def test_effective_team_member_budget_without_increase() -> None:
from litellm.proxy.auth.auth_checks import _effective_team_member_budget
now: Final = datetime(2026, 1, 1, tzinfo=timezone.utc)
assert _effective_team_member_budget(LiteLLM_BudgetTable(max_budget=100.0), now=now) == 100.0
assert _effective_team_member_budget(LiteLLM_BudgetTable(max_budget=None), now=now) is None
@pytest.mark.asyncio
async def test_team_member_budget_check_temp_budget_increase_extends_cap():
"""Spend above max_budget but below max_budget + active temp increase
must not raise; once the increase expires the same spend must raise."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import LiteLLM_TeamMembership
from litellm.proxy.utils import ProxyLogging
team_object = LiteLLM_TeamTable(team_id="test-team", metadata={})
user_object = LiteLLM_UserTable(user_id="test-user")
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
team_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=0.0,
budget_id="budget-1",
litellm_budget_table=LiteLLM_BudgetTable(
max_budget=100.0,
temp_budget_increase=100.0,
temp_budget_expiry=datetime.now(timezone.utc) + timedelta(hours=1),
),
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
prisma_client = MagicMock()
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
if counter_key == "spend:team_member:test-user:test-team":
return 150.0
return fallback_spend
# $150 spend is over the $100 cap but under the $200 temp-extended cap.
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=team_membership,
),
):
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=DualCache(),
proxy_logging_obj=proxy_logging_obj,
)
expired_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=0.0,
budget_id="budget-1",
litellm_budget_table=LiteLLM_BudgetTable(
max_budget=100.0,
temp_budget_increase=100.0,
temp_budget_expiry=datetime.now(timezone.utc) - timedelta(hours=1),
),
)
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=expired_membership,
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=DualCache(),
proxy_logging_obj=proxy_logging_obj,
)
assert exc_info.value.current_cost == 150.0
assert exc_info.value.max_budget == 100.0

View file

@ -15422,3 +15422,26 @@ async def test_team_info_reports_what_the_caller_may_edit(caller, org_admin, ena
)
assert response["team_info"].caller_edit_access.model_dump(mode="json") == expected
def test_build_member_budget_patch_maps_temp_budget_fields() -> None:
from litellm.proxy.management_endpoints.team_endpoints import _build_member_budget_patch
expiry: Final = datetime(2030, 1, 1, tzinfo=timezone.utc)
request: Final = TeamMemberUpdateRequest(
team_id="team-1",
user_id="user-1",
temp_budget_increase=50.0,
temp_budget_expiry=expiry,
)
assert _build_member_budget_patch(request) == {
"temp_budget_increase": 50.0,
"temp_budget_expiry": expiry,
}
def test_team_member_update_request_temp_budget_fields_must_be_set_together() -> None:
with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
TeamMemberUpdateRequest(team_id="team-1", user_id="user-1", temp_budget_increase=50.0)
with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
TeamMemberUpdateRequest(team_id="team-1", user_id="user-1", temp_budget_expiry="2030-01-01T00:00:00Z")