perf(proxy): read all user budget windows in one redis batch during auth

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
jesus 2026-09-18 22:42:50 +00:00
parent 6e9e40d608
commit 14ae5bd8d8
2 changed files with 40 additions and 3 deletions

View file

@ -108,6 +108,7 @@ from litellm.proxy.guardrails.tool_name_extraction import (
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.carried_budget_state import carry_organization_budget_state
from litellm.proxy.spend_tracking.spend_counter_batch import bind_spend_counter_keys
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
@ -5616,8 +5617,13 @@ async def _user_multi_budget_check(
from litellm.proxy.proxy_server import get_current_spend
for window in valid_token.user_budget_limits:
w: dict = window if isinstance(window, dict) else window.model_dump()
windows: Final[tuple[dict, ...]] = tuple(
window if isinstance(window, dict) else window.model_dump() for window in valid_token.user_budget_limits
)
bind_spend_counter_keys(
frozenset(f"spend:user:{valid_token.user_id}:window:{w['budget_duration']}" for w in windows)
)
for w in windows:
counter_key = f"spend:user:{valid_token.user_id}:window:{w['budget_duration']}"
window_spend = await get_current_spend(
counter_key=counter_key,

View file

@ -2,7 +2,7 @@
Unit tests for multi-budget-window enforcement on API keys.
"""
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -178,6 +178,37 @@ async def test_user_under_all_windows_passes():
assert all(call.kwargs["window_entity_type"] == "User" for call in spend_mock.await_args_list)
@pytest.mark.asyncio
async def test_user_windows_share_one_redis_batch_read(monkeypatch):
import litellm.proxy.proxy_server as ps
from litellm.caching.dual_cache import DualCache
from litellm.proxy.spend_tracking.spend_counter_batch import spend_counter_batch_scope
token = _make_user_token(
user_budget_limits=[
{"budget_duration": "24h", "max_budget": 10.0, "reset_at": None},
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": None},
]
)
redis = MagicMock()
redis.async_batch_get_cache = AsyncMock(
return_value={"spend:user:user-1:window:24h": 1.0, "spend:user:user-1:window:30d": 2.0}
)
redis.async_get_cache = AsyncMock(return_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", DualCache(redis_cache=redis))
monkeypatch.setattr(ps, "prisma_client", None)
with spend_counter_batch_scope(redis):
await _user_multi_budget_check(valid_token=token, team_object=None, general_settings={})
assert redis.async_batch_get_cache.await_count == 1
assert sorted(redis.async_batch_get_cache.await_args.kwargs["key_list"]) == [
"spend:user:user-1:window:24h",
"spend:user:user-1:window:30d",
]
assert redis.async_get_cache.await_count == 0
@pytest.mark.asyncio
async def test_user_over_any_window_raises():
token = _make_user_token(