Merge pull request #20688 from BerriAI/litellm_budget_tier_enforcement_for_keys

[Fix] Budget-linked keys never had spend reset
This commit is contained in:
yuneng-jiang 2026-03-06 20:44:58 -08:00 • committed by GitHub
commit b314e8d20a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 350 additions and 1 deletions

View file

@ -69,6 +69,39 @@ class ResetBudgetJob:
},
)
async def reset_budget_for_keys_linked_to_budgets(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
"""
Resets the spend for keys linked to budget tiers that are being reset.
This handles keys that have budget_id but no budget_duration set on the key
itself. Keys with budget_id rely on their linked budget tier's reset schedule
rather than having their own budget_duration.
Keys that have their own budget_duration are already handled by
reset_budget_for_litellm_keys() and are excluded here to avoid
double-resetting.
"""
budget_ids = [
budget.budget_id
for budget in budgets_to_reset
if budget.budget_id is not None
]
if not budget_ids:
return
return await self.prisma_client.db.litellm_verificationtoken.update_many(
where={
"budget_id": {"in": budget_ids},
"budget_duration": None, # only keys without their own reset schedule
"spend": {"gt": 0}, # only reset keys that have accumulated spend
},
data={
"spend": 0,
},
)
async def reset_budget_for_litellm_budget_table(self):
"""
Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired
@ -112,6 +145,10 @@ class ResetBudgetJob:
budgets_to_reset=budgets_to_reset
)
await self.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
if endusers_to_reset is not None and len(endusers_to_reset) > 0:
for enduser in endusers_to_reset:
try:
@ -579,4 +616,4 @@ class ResetBudgetJob:
await ResetBudgetJob._reset_budget_common(
item=key, current_time=current_time, item_type="key"
)
return key
return key

View file

@ -610,6 +610,10 @@ async def _common_key_generation_helper( # noqa: PLR0915
if _budget_id is not None:
data_json["budget_id"] = _budget_id
# Only set budget_duration on key when explicitly provided. Keys with budget_id
# but no explicit budget_duration follow their linked budget tier's schedule;
# reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
# This avoids duplicating budget_duration on keys so tier updates apply automatically.
if "budget_duration" in data_json:
data_json["key_budget_duration"] = data_json.pop("budget_duration", None)

View file

@ -229,6 +229,10 @@ async def test_reset_budget_endusers_partial_failure():
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -389,6 +393,10 @@ async def test_reset_budget_continues_other_categories_on_failure():
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -863,6 +871,10 @@ async def test_service_logger_endusers_success():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -938,6 +950,10 @@ async def test_service_logger_endusers_failure():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1026,6 +1042,9 @@ async def test_reset_budget_for_litellm_team_members_called():
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 2}
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()

View file

@ -25,9 +25,21 @@ class MockLiteLLMTeamMembership:
return {"count": 1}
class MockLiteLLMVerificationToken:
def __init__(self):
self.update_many_calls: List[Dict[str, Any]] = []
async def update_many(
self, where: Dict[str, Any], data: Dict[str, Any]
) -> Dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
class MockDB:
def __init__(self):
self.litellm_teammembership = MockLiteLLMTeamMembership()
self.litellm_verificationtoken = MockLiteLLMVerificationToken()
class MockPrismaClient:
@ -320,3 +332,150 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
assert mock_prisma_client.updated_data["user"][0].spend == 0.0
assert mock_prisma_client.updated_data["team"][0].spend == 0.0
assert mock_prisma_client.updated_data["enduser"][0].spend == 0.0
def test_reset_budget_for_keys_linked_to_budgets(reset_budget_job, mock_prisma_client):
"""
Test that when a budget tier is reset, keys linked to that budget
(via budget_id) that don't have their own budget_duration also get
their spend reset.
This covers the case where keys were created with budget_id but
budget_duration was not inherited to the key (pre-fix keys).
"""
from litellm.proxy._types import LiteLLM_BudgetTableFull
now = datetime.now(timezone.utc)
# Create a budget tier that is due for reset
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 10.0,
"budget_duration": "7d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "7d-budget-tier",
"created_at": now - timedelta(days=7),
},
)
budgets_to_reset = [test_budget]
# Run the method
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
)
# Verify that update_many was called on litellm_verificationtoken
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
assert len(calls) == 1, f"Expected 1 update_many call, got {len(calls)}"
# Verify the where clause filters by budget_id and null budget_duration
call = calls[0]
assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]}
assert call["where"]["budget_duration"] is None
# Verify spend is reset to 0
assert call["data"]["spend"] == 0
def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_duration(
reset_budget_job, mock_prisma_client
):
"""
Test that keys with BOTH budget_id AND budget_duration are excluded from
reset_budget_for_keys_linked_to_budgets. Such keys have their own reset
schedule and are handled only by reset_budget_for_litellm_keys(). The
budget_duration=None filter ensures they are NOT double-reset when the
linked budget tier expires.
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 10.0,
"budget_duration": "7d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "7d-budget-tier",
"created_at": now - timedelta(days=7),
},
)
budgets_to_reset = [test_budget]
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=budgets_to_reset
)
)
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
assert len(calls) == 1
call = calls[0]
# Critical: budget_duration must be None so keys with their own budget_duration
# (e.g. key has budget_id="X" AND budget_duration=60) are excluded.
# Those keys are reset only by reset_budget_for_litellm_keys() - no double-reset.
assert call["where"]["budget_duration"] is None
assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]}
def test_reset_budget_for_keys_linked_to_budgets_empty(
reset_budget_job, mock_prisma_client
):
"""
Test that when there are no budgets to reset, no update is performed
on the verification token table.
"""
# Run with empty list
asyncio.run(
reset_budget_job.reset_budget_for_keys_linked_to_budgets(
budgets_to_reset=[]
)
)
# Verify no update_many calls were made
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
assert len(calls) == 0
def test_budget_table_reset_also_resets_linked_keys(
reset_budget_job, mock_prisma_client
):
"""
Integration-style test: when reset_budget_for_litellm_budget_table runs,
it should also reset spend for keys linked to the expiring budget tiers
(in addition to end-users and team members).
"""
now = datetime.now(timezone.utc)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 10.0,
"budget_duration": "7d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "7d-budget-tier",
"created_at": now - timedelta(days=7),
},
)
mock_prisma_client.data["budget"] = [test_budget]
# Run the full budget table reset
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Verify that keys linked to the budget were also reset
calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls
assert len(calls) == 1, (
"Expected reset_budget_for_litellm_budget_table to also reset keys "
f"linked to expiring budgets, but got {len(calls)} update_many calls"
)
assert calls[0]["where"]["budget_id"] == {"in": ["7d-budget-tier"]}
assert calls[0]["data"]["spend"] == 0

View file

@ -5637,6 +5637,136 @@ async def test_validate_key_list_check_key_hash_not_found():
assert "Key Hash not found" in exc_info.value.message
@pytest.mark.asyncio
async def test_key_with_budget_id_does_not_store_budget_duration():
"""
Test that when a key is created with budget_id but without explicit
budget_duration, the key does NOT get budget_duration stored on it.
Keys with budget_id follow their linked budget tier's reset schedule;
reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
This avoids duplicating budget_duration on keys so tier updates apply
automatically to all linked keys.
"""
from unittest.mock import AsyncMock, MagicMock, patch
mock_prisma = MagicMock()
mock_generate_key = AsyncMock(
return_value={
"key": "sk-test-key",
"expires": None,
"user_id": "test-user",
"team_id": None,
}
)
with patch(
"litellm.proxy.proxy_server.prisma_client", mock_prisma
), patch(
"litellm.proxy.proxy_server.llm_router", None
), patch(
"litellm.proxy.proxy_server.premium_user", False
), patch(
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
), patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
mock_generate_key,
):
await _common_key_generation_helper(
data=GenerateKeyRequest(
budget_id="7d-budget-tier",
max_budget=10.0,
# NOTE: budget_duration is intentionally NOT set here
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="admin-user",
),
litellm_changed_by=None,
team_table=None,
)
mock_generate_key.assert_awaited_once()
call_kwargs = mock_generate_key.call_args.kwargs
# Key should NOT have key_budget_duration - it follows the budget tier's schedule
assert call_kwargs.get("key_budget_duration") is None, (
"key_budget_duration should be None for budget-linked keys without explicit "
f"budget_duration; got: {call_kwargs.get('key_budget_duration')}"
)
# No budget tier lookup - we don't copy budget_duration onto the key
mock_prisma.db.litellm_budgettable.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_key_does_not_override_explicit_budget_duration():
"""
Test that when a key is created with both budget_id and an explicit
budget_duration, the explicit budget_duration takes precedence over
the budget tier's budget_duration.
"""
from unittest.mock import AsyncMock, MagicMock, patch
mock_prisma = MagicMock()
# The budget tier has budget_duration="7d"
mock_budget_row = MagicMock()
mock_budget_row.budget_duration = "7d"
mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(
return_value=mock_budget_row
)
mock_generate_key = AsyncMock(
return_value={
"key": "sk-test-key",
"expires": None,
"user_id": "test-user",
"team_id": None,
}
)
with patch(
"litellm.proxy.proxy_server.prisma_client", mock_prisma
), patch(
"litellm.proxy.proxy_server.llm_router", None
), patch(
"litellm.proxy.proxy_server.premium_user", False
), patch(
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
), patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
mock_generate_key,
):
await _common_key_generation_helper(
data=GenerateKeyRequest(
budget_id="7d-budget-tier",
max_budget=10.0,
budget_duration="30d", # explicit budget_duration should take precedence
),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="admin-user",
),
litellm_changed_by=None,
team_table=None,
)
mock_generate_key.assert_awaited_once()
call_kwargs = mock_generate_key.call_args.kwargs
# The explicit budget_duration should take precedence
assert call_kwargs.get("key_budget_duration") == "30d", (
"Explicit budget_duration should take precedence over the budget tier's value "
f"but got: {call_kwargs.get('key_budget_duration')}"
)
# The budget tier should NOT have been looked up since budget_duration was explicit
mock_prisma.db.litellm_budgettable.find_unique.assert_not_called()
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"