fix to ensure budget duration is being inherited from budget tier for keys

This commit is contained in:
shivam 2026-02-07 18:16:01 -08:00
parent dd5c14baf8
commit ba1b466480
4 changed files with 299 additions and 0 deletions

View file

@ -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:

View file

@ -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

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,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

View file

@ -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()