mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
ca1f69fb73
commit
d9c43d5e17
2 changed files with 51 additions and 208 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue