mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
b314e8d20a
5 changed files with 350 additions and 1 deletions
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue