fix(key management): read budget window usage from the window spend table

Pass window_duration to get_current_spend so /key/info re-checks a stale-low counter against the LiteLLM_BudgetWindowSpend row instead of aggregating LiteLLM_SpendLogs, and reuse _budget_limit_windows for the stored-column coercion. Drop the /v2/key/info batch cap (a new 422 for callers that work today) and the unrelated CI timeout bump and soft_budget test
This commit is contained in:
ryan-crabbe-berri 2026-08-31 20:07:28 -07:00
parent ca1f69fb73
commit d9c43d5e17
2 changed files with 51 additions and 208 deletions

View file

@ -3579,95 +3579,58 @@ async def _build_model_max_budget_usage(
)
def _budget_window_to_dict(window: object) -> Mapping[str, object] | None:
"""Coerce a budget_limits entry to a dict; None when the entry is unusable."""
if isinstance(window, dict):
return window
model_dump: Final = getattr(window, "model_dump", None)
if not callable(model_dump):
def _window_max_budget(window: Mapping[str, object]) -> float | None:
"""A window's max_budget as a float; None when absent or unparseable."""
value: Final = window.get("max_budget")
if not isinstance(value, (int, float, str)):
return None
try:
dumped: Final = model_dump()
except Exception: # noqa: BLE001 # model_dump implementations can raise arbitrary errors
return float(value)
except ValueError:
return None
return dumped if isinstance(dumped, dict) else None
def _coerce_budget_limits(budget_limits: object) -> Sequence[object] | None:
"""Coerce budget_limits to a sequence of windows, parsing JSON strings; None when unusable."""
if isinstance(budget_limits, str):
try:
parsed: Final = json.loads(budget_limits)
except (TypeError, ValueError):
return None
return parsed if isinstance(parsed, list) else None
return budget_limits if isinstance(budget_limits, list) else None
def _parse_window_max_budget(value: object) -> float | None:
"""Coerce a window's max_budget to float; None when absent or unparseable."""
if isinstance(value, (int, float, str)):
try:
return float(value)
except (TypeError, ValueError):
return None
return None
async def _budget_window_with_usage(window: Mapping[str, object], api_key_hash: str) -> Mapping[str, object]:
"""
Return a copy of a budget window with current-window spend attached.
Copy of a budget window with current-window spend attached.
Per-window spend is not persisted in the DB; it lives in the cross-pod spend
counters (spend:key:{hashed_token}:window:{budget_duration}) that
_virtual_key_multi_budget_check enforces against, so we read the same
counters via get_current_spend. Passing max_budget + window_start makes the
read re-check against the authoritative spend-log aggregate when the counter
is stale-low (e.g. after a Redis flush), same as the enforcement path.
Reads the same cross-pod counter (spend:key:{hashed_token}:window:{budget_duration})
that _virtual_key_multi_budget_check enforces against, passing the same
window_duration + window_start so a stale-low counter is re-checked against
the LiteLLM_BudgetWindowSpend row instead of a spend-log aggregate.
"""
from litellm.proxy.proxy_server import get_current_spend
duration: Final = window.get("budget_duration")
if not duration:
return dict(window) # mutable-ok: per-window response copy, built once per window
if not isinstance(duration, str) or not duration:
return window
spend: Final = await get_current_spend(
counter_key=f"spend:key:{api_key_hash}:window:{duration}",
fallback_spend=0.0,
max_budget=_parse_window_max_budget(window.get("max_budget")),
max_budget=_window_max_budget(window),
window_entity_type="Key",
window_entity_id=api_key_hash,
window_duration=duration,
window_start=get_budget_window_start(window),
)
return {**window, "current_spend": round(spend, 4)} # mutable-ok: per-window response copy, built once per window
async def _budget_limits_entry_with_usage(window: object, api_key_hash: str) -> object:
"""Return the window as an enriched dict when dict-coercible; the original entry otherwise."""
coerced: Final = _budget_window_to_dict(window)
if not coerced:
return window
return await _budget_window_with_usage(window=coerced, api_key_hash=api_key_hash)
async def _budget_limits_with_usage(budget_limits: object, api_key_hash: str) -> Sequence[object] | None:
async def _budget_limits_with_usage(
budget_limits: Sequence[object] | str | None, api_key_hash: str
) -> tuple[Mapping[str, object], ...] | None:
"""
Return budget_limits as window dicts with current-window spend attached.
None when budget_limits is not a usable (possibly JSON-encoded) list; the
caller keeps the original value then. Entries that are not dict-coercible
are preserved as-is.
budget_limits as window dicts with current-window spend attached; None when
the key has no windows so the caller keeps the stored value.
"""
windows: Final = _coerce_budget_limits(budget_limits)
if windows is None:
windows: Final = _budget_limit_windows(budget_limits)
if not windows:
return None
return [ # mutable-ok: entries are awaited, so they cannot be built inside a frozen wrapper
await _budget_limits_entry_with_usage(window=window, api_key_hash=api_key_hash) for window in windows
]
# Caps per-request fan-out: each key with budget windows costs one spend-counter
# read (worst case a SpendLogs aggregation) per window.
MAX_KEY_INFO_KEYS_PER_REQUEST: Final = 100
return tuple(
await asyncio.gather(
*(_budget_window_with_usage(window=window, api_key_hash=api_key_hash) for window in windows)
)
)
@router.post(
@ -3712,18 +3675,6 @@ async def info_key_fn_v2(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail={"message": "Malformed request. No keys passed in."},
)
requested_key_count: Final = len(data.keys or ()) + len(data.key_aliases or ())
if requested_key_count > MAX_KEY_INFO_KEYS_PER_REQUEST:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail={ # mutable-ok: one-shot HTTPException payload matching the sibling detail dict above; never mutated after construction
"message": (
f"Too many keys requested: {requested_key_count}. "
f"At most {MAX_KEY_INFO_KEYS_PER_REQUEST} keys and key_aliases combined per request."
)
},
)
# Resolve key_aliases to tokens so we never pass token=None (unbounded query)
tokens_to_query: Final = list(data.keys) if data.keys else []
if data.key_aliases:

View file

@ -539,57 +539,6 @@ async def test_key_generation_with_object_permission(monkeypatch):
assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123"
@pytest.mark.asyncio
async def test_generate_key_with_soft_budget_creates_budget_row(monkeypatch):
"""soft_budget on /key/generate must create a budget table row and link its budget_id to the key."""
mock_prisma_client = AsyncMock()
mock_prisma_client.jsonify_object = lambda data: data
mock_prisma_client.db = MagicMock()
mock_budget_create = AsyncMock(return_value=MagicMock(budget_id="budget-soft-123"))
mock_prisma_client.db.litellm_budgettable = MagicMock()
mock_prisma_client.db.litellm_budgettable.create = mock_budget_create
async def _insert_data_side_effect(*args, **kwargs):
if kwargs.get("table_name") == "user":
return MagicMock(models=[], spend=0)
return MagicMock(
token="hashed_token_soft",
litellm_budget_table=None,
object_permission=None,
)
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_fn,
)
await generate_key_fn(
data=GenerateKeyRequest(soft_budget=5.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="admin-1",
),
)
mock_budget_create.assert_awaited_once()
created_budget = mock_budget_create.call_args.kwargs["data"]
assert created_budget["soft_budget"] == 5.0
assert created_budget["created_by"] == "admin-1"
key_insert_calls = [
call.kwargs
for call in mock_prisma_client.insert_data.call_args_list
if call.kwargs.get("table_name") == "key"
]
assert len(key_insert_calls) == 1
assert key_insert_calls[0]["data"].get("budget_id") == "budget-soft-123"
@pytest.mark.asyncio
async def test_generate_key_debug_log_never_contains_raw_token(monkeypatch, caplog):
"""Regression for LIT-4356: /key/generate must never emit the raw virtual key
@ -14267,6 +14216,7 @@ async def test_info_key_fn_budget_limits_includes_current_spend(monkeypatch):
assert call_kwargs["max_budget"] == 2.0
assert call_kwargs["window_entity_type"] == "Key"
assert call_kwargs["window_entity_id"] == test_key_token
assert call_kwargs["window_duration"] == "1h"
assert call_kwargs["window_start"] is not None
@ -14397,39 +14347,9 @@ async def test_info_key_fn_v2_budget_limits_includes_current_spend(monkeypatch):
f"spend:key:{test_key_token}:window:1h",
f"spend:key:{test_key_token}:window:1d",
}
@pytest.mark.asyncio
async def test_info_key_fn_v2_rejects_oversized_batch(monkeypatch):
"""/v2/key/info must reject over-cap batches before doing any DB work."""
from unittest.mock import AsyncMock
from litellm.proxy._types import KeyRequest, ProxyException
from litellm.proxy.management_endpoints.key_management_endpoints import (
MAX_KEY_INFO_KEYS_PER_REQUEST,
info_key_fn_v2,
)
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock())
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin-batch-cap",
)
with pytest.raises(ProxyException) as exc_info:
await info_key_fn_v2(
data=KeyRequest(
keys=[f"hash-{i}" for i in range(MAX_KEY_INFO_KEYS_PER_REQUEST)],
key_aliases=["alias-over-cap"],
),
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.code == "422"
mock_prisma_client.get_data.assert_not_awaited()
assert {
call.kwargs["window_duration"] for call in mock_get_current_spend.await_args_list
} == {"1h", "1d"}
@pytest.mark.asyncio
@ -14452,7 +14372,7 @@ async def test_budget_limits_with_usage_json_string_input(monkeypatch):
)
result = await _budget_limits_with_usage(budget_limits=raw, api_key_hash="hash-1")
assert result == [
assert list(result) == [
{
"budget_duration": "1h",
"max_budget": 2.0,
@ -14464,8 +14384,8 @@ async def test_budget_limits_with_usage_json_string_input(monkeypatch):
@pytest.mark.asyncio
async def test_budget_limits_with_usage_skips_unusable_inputs(monkeypatch):
"""Invalid JSON strings, non-list values, and malformed windows are skipped."""
async def test_budget_limits_with_usage_empty_windows_keep_stored_value(monkeypatch):
"""A key with no windows (None, [], or "[]") returns None so /key/info keeps the stored value; no spend lookup runs."""
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -14477,32 +14397,9 @@ async def test_budget_limits_with_usage_skips_unusable_inputs(monkeypatch):
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
)
# invalid JSON string and non-list values return None: callers keep the original
assert await _budget_limits_with_usage(budget_limits="{not json", api_key_hash="hash-1") is None
assert await _budget_limits_with_usage(budget_limits={"budget_duration": "1h"}, api_key_hash="hash-1") is None
# windows that are falsy, missing budget_duration, or not dict-like
windows = [
{},
{"max_budget": 2.0},
{"budget_duration": "1h", "max_budget": "not-a-number"},
42,
]
result = await _budget_limits_with_usage(budget_limits=windows, api_key_hash="hash-1")
# only the well-formed window (with unparseable max_budget coerced to None)
# triggers a spend lookup
mock_get_current_spend.assert_awaited_once()
call_kwargs = mock_get_current_spend.await_args.kwargs
assert call_kwargs["counter_key"] == "spend:key:hash-1:window:1h"
assert call_kwargs["max_budget"] is None
assert result is not None
assert result[0] == {}
assert result[1] == {"max_budget": 2.0}
assert result[2] == {"budget_duration": "1h", "max_budget": "not-a-number", "current_spend": 0.0}
assert result[3] == 42
# input is not mutated
assert windows[2] == {"budget_duration": "1h", "max_budget": "not-a-number"}
for stored in (None, [], "[]"):
assert await _budget_limits_with_usage(budget_limits=stored, api_key_hash="hash-1") is None
mock_get_current_spend.assert_not_awaited()
@pytest.mark.asyncio
@ -14523,17 +14420,19 @@ async def test_budget_limits_with_usage_window_without_max_budget(monkeypatch):
budget_limits=[{"budget_duration": "2d"}], api_key_hash="hash-no-max"
)
assert result == [{"budget_duration": "2d", "current_spend": 0.75}]
assert list(result) == [{"budget_duration": "2d", "current_spend": 0.75}]
call_kwargs = mock_get_current_spend.await_args.kwargs
assert call_kwargs["counter_key"] == "spend:key:hash-no-max:window:2d"
assert call_kwargs["window_duration"] == "2d"
assert call_kwargs["max_budget"] is None
@pytest.mark.asyncio
async def test_budget_limits_with_usage_pydantic_windows(monkeypatch):
"""Window objects with model_dump() are converted to dicts; failing windows pass through."""
from unittest.mock import AsyncMock, MagicMock
"""BudgetLimitEntry windows (the shape UserAPIKeyAuth carries) are dumped to dicts and annotated."""
from unittest.mock import AsyncMock
from litellm.models.team import BudgetLimitEntry
from litellm.proxy.management_endpoints.key_management_endpoints import (
_budget_limits_with_usage,
)
@ -14543,25 +14442,18 @@ async def test_budget_limits_with_usage_pydantic_windows(monkeypatch):
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
)
good_window = MagicMock()
good_window.model_dump.return_value = {
"budget_duration": "7d",
"max_budget": 10.0,
"reset_at": None,
}
bad_window = MagicMock()
bad_window.model_dump.side_effect = ValueError("boom")
result = await _budget_limits_with_usage(
budget_limits=[good_window, bad_window], api_key_hash="hash-2"
budget_limits=[BudgetLimitEntry(budget_duration="7d", max_budget=10.0)],
api_key_hash="hash-2",
)
# good window converted to dict and annotated; failing window left as-is
assert result is not None
assert isinstance(result[0], dict)
assert result[0]["current_spend"] == 1.0
assert result[1] is bad_window
mock_get_current_spend.assert_awaited_once()
assert list(result) == [
{"budget_duration": "7d", "max_budget": 10.0, "reset_at": None, "current_spend": 1.0}
]
call_kwargs = mock_get_current_spend.await_args.kwargs
assert call_kwargs["counter_key"] == "spend:key:hash-2:window:7d"
assert call_kwargs["window_duration"] == "7d"
assert call_kwargs["max_budget"] == 10.0
@pytest.mark.asyncio