mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
4b368bf066
commit
7f3f8fae2d
10 changed files with 198 additions and 1 deletions
|
|
@ -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);
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue