mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix to ensure budget duration is being inherited from budget tier for keys
This commit is contained in:
parent
dd5c14baf8
commit
ba1b466480
4 changed files with 299 additions and 0 deletions
|
|
@ -69,6 +69,38 @@ 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 (e.g. keys created before the fix to inherit budget_duration from
|
||||
the linked budget tier).
|
||||
|
||||
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
|
||||
},
|
||||
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 +144,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:
|
||||
|
|
|
|||
|
|
@ -594,6 +594,13 @@ async def _common_key_generation_helper( # noqa: PLR0915
|
|||
|
||||
if "budget_duration" in data_json:
|
||||
data_json["key_budget_duration"] = data_json.pop("budget_duration", None)
|
||||
elif _budget_id is not None and prisma_client is not None:
|
||||
# Inherit budget_duration from linked budget tier if not explicitly set on the key
|
||||
budget_row = await prisma_client.db.litellm_budgettable.find_unique(
|
||||
where={"budget_id": _budget_id}
|
||||
)
|
||||
if budget_row is not None and budget_row.budget_duration is not None:
|
||||
data_json["key_budget_duration"] = budget_row.budget_duration
|
||||
|
||||
if user_api_key_dict.user_id is not None:
|
||||
data_json["created_by"] = user_api_key_dict.user_id
|
||||
|
|
|
|||
|
|
@ -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,107 @@ 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_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
|
||||
|
|
|
|||
|
|
@ -5559,3 +5559,143 @@ async def test_validate_key_list_check_key_hash_not_found():
|
|||
|
||||
assert exc_info.value.code == "403" or exc_info.value.code == 403
|
||||
assert "Key Hash not found" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_inherits_budget_duration_from_budget_tier():
|
||||
"""
|
||||
Test that when a key is created with budget_id pointing to a budget tier
|
||||
that has budget_duration, the key inherits budget_duration from the tier
|
||||
even when budget_duration is not explicitly set on the key request.
|
||||
|
||||
This verifies the fix for the bug where keys created with budget_id
|
||||
would have null budget_duration and budget_reset_at, causing the
|
||||
budget reset job to never reset their spend.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
# Mock the budget tier lookup to return a budget with budget_duration="7d"
|
||||
mock_budget_row = MagicMock()
|
||||
mock_budget_row.budget_duration = "7d"
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
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,
|
||||
# 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,
|
||||
)
|
||||
|
||||
# Verify generate_key_helper_fn was called
|
||||
mock_generate_key.assert_awaited_once()
|
||||
call_kwargs = mock_generate_key.call_args.kwargs
|
||||
|
||||
# The key should have inherited key_budget_duration from the budget tier
|
||||
assert call_kwargs.get("key_budget_duration") == "7d", (
|
||||
"key_budget_duration should be inherited from the linked budget tier "
|
||||
f"but got: {call_kwargs.get('key_budget_duration')}"
|
||||
)
|
||||
|
||||
# Verify the budget tier was looked up with the correct budget_id
|
||||
mock_prisma.db.litellm_budgettable.find_unique.assert_awaited_once_with(
|
||||
where={"budget_id": "7d-budget-tier"}
|
||||
)
|
||||
|
||||
|
||||
@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_awaited()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue