From e93245131261ad4c4a7b05feba8d1325925accc7 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 19:56:35 +0000 Subject: [PATCH] refactor(proxy): move effective member budget onto the budget model and reject negative temp increases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/models/budget.py | 17 ++++++++++- litellm/proxy/_types.py | 1 + litellm/proxy/auth/auth_checks.py | 21 +------------ litellm/proxy/auth/user_api_key_auth.py | 4 +-- .../spend_tracking/budget_reservation.py | 7 +---- tests/test_litellm/models/test_models.py | 23 +++++++++++++- .../proxy/auth/test_auth_checks.py | 30 ------------------- .../test_team_endpoints.py | 7 +++++ 8 files changed, 49 insertions(+), 61 deletions(-) diff --git a/litellm/models/budget.py b/litellm/models/budget.py index ddc694743c4..2bef5d279d6 100644 --- a/litellm/models/budget.py +++ b/litellm/models/budget.py @@ -5,7 +5,8 @@ Canonical definition for ``litellm_budgettable``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ -from datetime import datetime +from datetime import datetime, timezone +from typing import Final from pydantic import ConfigDict @@ -35,6 +36,20 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase): model_config = ConfigDict(protected_namespaces=()) + def effective_max_budget(self, now: datetime) -> float | None: + if self.max_budget is None: + return None + if self.temp_budget_increase is None or self.temp_budget_expiry is None: + return self.max_budget + expiry: Final = ( + self.temp_budget_expiry.replace(tzinfo=timezone.utc) + if self.temp_budget_expiry.tzinfo is None + else self.temp_budget_expiry + ) + if expiry <= now: + return self.max_budget + return self.max_budget + self.temp_budget_increase + class LiteLLM_BudgetTableFull(LiteLLM_BudgetTable): """LiteLLM_BudgetTable + server-managed fields returned on API responses.""" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index df891bea5e0..7728c7ee34f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4399,6 +4399,7 @@ class TeamMemberUpdateRequest(TeamMemberDeleteRequest): ) temp_budget_increase: float | None = Field( default=None, + ge=0, description="Temporary additive budget increase for this team member, active until temp_budget_expiry", ) temp_budget_expiry: datetime | None = Field( diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index a33cda43758..a36ef548e2f 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -14,7 +14,6 @@ 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 @@ -5296,21 +5295,6 @@ async def _virtual_key_max_budget_alert_check( ) -def _effective_team_member_budget(budget: LiteLLM_BudgetTable, now: datetime) -> float | None: - 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, @@ -5346,10 +5330,7 @@ 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 = _effective_team_member_budget( - loaded_membership.litellm_budget_table, - now=get_utc_datetime(), - ) + team_member_budget = loaded_membership.litellm_budget_table.effective_max_budget(now=get_utc_datetime()) else: default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id") if isinstance(default_budget_id, str): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index c2dc9b16735..cffb745fff2 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -46,7 +46,6 @@ from litellm.proxy.auth.auth_checks import ( _can_object_call_model, _check_end_user_budget, _delete_cache_key_object, - _effective_team_member_budget, _get_user_role, _is_model_cost_zero, _is_user_proxy_admin, @@ -2249,8 +2248,7 @@ async def _user_api_key_auth_builder( ) if team_member_info is not None and team_member_info.litellm_budget_table is not None: - team_member_budget: Final = _effective_team_member_budget( - team_member_info.litellm_budget_table, + team_member_budget: Final = team_member_info.litellm_budget_table.effective_max_budget( now=datetime.now(timezone.utc), ) if team_member_budget is not None and team_member_budget > 0: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 1c6f20e515b..4028fcaf2ae 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -690,12 +690,7 @@ async def _get_team_member_budget_counter( team_member_budget: float | None = None if team_membership is not None and team_membership.litellm_budget_table is not None: - from litellm.proxy.auth.auth_checks import _effective_team_member_budget - - team_member_budget = _effective_team_member_budget( - team_membership.litellm_budget_table, - now=datetime.now(timezone.utc), - ) + team_member_budget = team_membership.litellm_budget_table.effective_max_budget(now=datetime.now(timezone.utc)) else: default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id") if isinstance(default_budget_id, str): diff --git a/tests/test_litellm/models/test_models.py b/tests/test_litellm/models/test_models.py index 9b803c14062..46648efdfff 100644 --- a/tests/test_litellm/models/test_models.py +++ b/tests/test_litellm/models/test_models.py @@ -2,7 +2,7 @@ Tests for backend domain models. """ -from datetime import datetime +from datetime import datetime, timezone import pytest from pydantic import BaseModel, TypeAdapter @@ -71,6 +71,27 @@ class TestBudget: assert budget.max_budget is None assert budget.allowed_models is None + def test_effective_max_budget_applies_unexpired_increase(self): + budget = LiteLLM_BudgetTable( + max_budget=100.0, + temp_budget_increase=50.0, + temp_budget_expiry=datetime(2100, 1, 1), + ) + assert budget.effective_max_budget(now=datetime(2026, 1, 1, tzinfo=timezone.utc)) == 150.0 + + def test_effective_max_budget_ignores_expired_increase(self): + budget = LiteLLM_BudgetTable( + max_budget=100.0, + temp_budget_increase=50.0, + temp_budget_expiry=datetime(2020, 1, 1, tzinfo=timezone.utc), + ) + assert budget.effective_max_budget(now=datetime(2026, 1, 1, tzinfo=timezone.utc)) == 100.0 + + def test_effective_max_budget_without_increase(self): + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + assert LiteLLM_BudgetTable(max_budget=100.0).effective_max_budget(now=now) == 100.0 + assert LiteLLM_BudgetTable(max_budget=None, temp_budget_increase=50.0).effective_max_budget(now=now) is None + class TestCredentials: def test_credentials_creation(self): diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 518dad8c48f..3048d4a6a02 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -8463,36 +8463,6 @@ def test_request_skips_budget_checks_extends_route_rule_with_zero_cost_models() 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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 0b9597da057..b458b85602d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -15445,3 +15445,10 @@ def test_team_member_update_request_temp_budget_fields_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") + + +def test_team_member_update_request_rejects_negative_temp_budget_increase() -> None: + with pytest.raises(ValidationError, match="greater than or equal to 0"): + TeamMemberUpdateRequest( + team_id="team-1", user_id="user-1", temp_budget_increase=-1.0, temp_budget_expiry="2030-01-01T00:00:00Z" + )