fix(live): keep budget reads loud and complete in the delegation gate

The cached auth helpers cover the same rows this gate needs, but two of them answer a
question this gate asks differently. get_team_object reports every failed database read as
an HTTP 404, and the 404 mapping turned that into "no team"; get_team_member_default_budget
returns None when its read raises. During an outage both convert a configured limit into an
absent one, and the gate answers "no limits" by allowing managed Live delegation, so
unsupervised spend becomes reachable through a failing database. The group-budget batch
helper also flattens its row to spend and max_budget, dropping a linked group's rpm and tpm
limits, and no other path in the proxy enforces those, so this gate was their only consumer.

Read the team and default-budget rows through the proxy cache under the keys and TTL the rest
of the proxy already uses, storing the entry on a miss and letting read errors reach the 503
handler, and fetch each group's linked budget in one batched query cached per group. The
cached team entry keeps last_refreshed_at, so it is the same object the chat path writes.
This commit is contained in:
jibanez-staticduo 2026-09-26 01:36:18 +02:00
parent 8e30dd1be5
commit 1a08de0af3
No known key found for this signature in database
2 changed files with 219 additions and 27 deletions

View file

@ -3,10 +3,10 @@ import base64
import hashlib
import json
import time
from collections.abc import AsyncGenerator, Mapping
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
from contextlib import asynccontextmanager, nullcontext
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, TypeVar
import httpx
from fastapi import APIRouter, HTTPException, Request, Response, WebSocket, WebSocketDisconnect
@ -23,6 +23,7 @@ from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.team import LiteLLM_TeamTable
from litellm.proxy._types import (
LiteLLM_ProjectTableCachedObj,
LiteLLM_TeamTableCachedObj,
LitellmUserRoles,
UserAPIKeyAuth,
)
@ -31,16 +32,15 @@ from litellm.proxy.auth.auth_checks import (
can_org_access_model,
can_user_call_model,
collect_matched_model_access_groups,
get_model_access_group_budgets_batch,
get_org_object,
get_project_object,
get_team_member_default_budget,
get_team_membership,
get_team_object,
get_user_object,
)
from litellm.proxy.auth.user_api_key_auth import get_websocket_api_key, user_api_key_auth
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity
)
@ -54,10 +54,15 @@ from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS,
from litellm.proxy.spend_tracking.budget_reservation import (
release_or_invalidate_budget_reservation, # pyright: ignore[reportUnknownVariableType] # budget helper accepts legacy reservation dicts
)
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.table_repositories import ModelAccessGroupBudgetRepository
from litellm.repositories.team_repository import TeamRepository
_routes: Final = APIRouter()
_JSON: Final = TypeAdapter[JsonValue](JsonValue)
_EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({})
_CACHEABLE_MODEL = TypeVar("_CACHEABLE_MODEL", bound=BaseModel)
_LIVE_GROUP_LIMITS_CACHE_PREFIX: Final = "live:model_access_group_limits:"
_MAPPING: Final = TypeAdapter(Mapping[str, object])
_OBJECT: Final = TypeAdapter(Mapping[str, JsonValue])
_DEPLOYMENT: Final = TypeAdapter(LiveDeployment)
@ -685,22 +690,57 @@ async def _live_team_membership(auth: UserAPIKeyAuth) -> object | None:
)
async def _live_cached_object(
*,
key: str,
model_type: type[_CACHEABLE_MODEL],
load: Callable[[], Awaitable[_CACHEABLE_MODEL | None]],
) -> _CACHEABLE_MODEL | None:
"""Read one management object through the proxy cache, storing the row when the read misses.
``auth_checks`` already caches these rows for the chat path, but the two getters this gate
would use are unusable there: ``get_team_object`` reports every failed database read as an
HTTP 404, and ``get_team_member_default_budget`` returns ``None`` when its read raises. Both
turn an outage into "no limit configured", and this gate answers that question by allowing
managed delegation, so an unreadable limit has to stay an error. The cache key and TTL stay
the shared ones, so the entry is still written, read, and invalidated like any other.
"""
from litellm.proxy import proxy_server as server
cached: Final = await server.user_api_key_cache.async_get_cache(key=key, model_type=model_type)
if cached is not None:
return cached
loaded: Final = await load()
if loaded is not None:
await server.user_api_key_cache.async_set_cache(
key=key,
value=loaded,
model_type=model_type,
ttl=get_management_object_ttl(server.user_api_key_cache),
)
return loaded
async def _live_team(auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:
from litellm.proxy import proxy_server as server
if auth.team_id is None:
return None
try:
return await get_team_object(
team_id=auth.team_id,
prisma_client=server.prisma_client,
user_api_key_cache=server.user_api_key_cache,
proxy_logging_obj=server.proxy_logging_obj,
)
except HTTPException as exc:
if exc.status_code == 404:
team_id: Final = auth.team_id
async def load() -> LiteLLM_TeamTableCachedObj | None:
row: Final = await TeamRepository(server.prisma_client).find_by_id(team_id, id_field="team_id")
if row is None:
return None
raise
team: Final = LiteLLM_TeamTableCachedObj.model_validate(row.model_dump())
team.last_refreshed_at = time.time()
return team
return await _live_cached_object(
key=f"team_id:{team_id}",
model_type=LiteLLM_TeamTableCachedObj,
load=load,
)
def _live_team_budget_configured(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> bool:
@ -731,11 +771,13 @@ async def _live_default_budget(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | N
return None
from litellm.proxy import proxy_server as server
# Like chat auth, a failed default-budget read returns None; membership errors still fail closed above.
return await get_team_member_default_budget(
default_id,
server.prisma_client,
server.user_api_key_cache,
async def load() -> LiteLLM_BudgetTable | None:
return await BudgetRepository(server.prisma_client).find_by_id(default_id, id_field="budget_id")
return await _live_cached_object(
key=f"team_member_default_budget:{default_id}",
model_type=LiteLLM_BudgetTable,
load=load,
)
@ -769,6 +811,72 @@ async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM
return _managed_constraints(auth.model_copy(update=MappingProxyType({"project_metadata": project_metadata})))
def _live_group_limits(row: object) -> LiteLLM_BudgetTable:
"""The limit fields of the budget linked to one model access group row.
The row arrives as a Prisma join, so the limits are read by name. A group with no linked
budget yields an empty budget table: it reads as no limit, which is what the gate needs, and
it stays cacheable so the group is not re-read on every request.
"""
budget: Final = getattr(row, "litellm_budget_table", None)
if budget is None:
return LiteLLM_BudgetTable()
return LiteLLM_BudgetTable.model_validate(
{ # mutable-ok: field values are read from the joined row into a fresh validation mapping
field: getattr(budget, field, None)
for field in ("max_budget", "rpm_limit", "tpm_limit", "model_max_budget", "max_parallel_requests")
}
)
async def _live_fetch_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_BudgetTable, ...]:
"""Fetch the linked budget of each group in one query and cache one entry per group."""
if not groups:
return ()
from litellm.proxy import proxy_server as server
rows: Final = await ModelAccessGroupBudgetRepository(server.prisma_client).table.find_many(
where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries
"access_group_name": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries
"in": list(groups),
}
},
include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries
)
linked: Final = MappingProxyType({getattr(row, "access_group_name", None): _live_group_limits(row) for row in rows})
limits: Final = tuple(linked.get(group) or LiteLLM_BudgetTable() for group in groups)
await asyncio.gather(
*(
server.user_api_key_cache.async_set_cache(
key=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{group}",
value=limit,
model_type=LiteLLM_BudgetTable,
ttl=get_management_object_ttl(server.user_api_key_cache),
)
for group, limit in zip(groups, limits)
)
)
return limits
async def _live_model_group_limits(groups: tuple[str, ...]) -> tuple[LiteLLM_BudgetTable, ...]:
"""One cached budget entry per group, served from a single row batch on a cold miss."""
from litellm.proxy import proxy_server as server
cached: Final = await asyncio.gather(
*(
server.user_api_key_cache.async_get_cache(
key=f"{_LIVE_GROUP_LIMITS_CACHE_PREFIX}{group}",
model_type=LiteLLM_BudgetTable,
)
for group in groups
)
)
uncached: Final = tuple(group for group, entry in zip(groups, cached) if entry is None)
fetched: Final = MappingProxyType(dict(zip(uncached, await _live_fetch_group_limits(uncached))))
return tuple(entry if entry is not None else fetched[group] for group, entry in zip(groups, cached))
async def _live_model_group_budget_configured(
auth: UserAPIKeyAuth,
model: str | None,
@ -792,13 +900,10 @@ async def _live_model_group_budget_configured(
)
if not matched_groups:
return False
budgets: Final = await get_model_access_group_budgets_batch(
matched_groups,
server.prisma_client,
server.user_api_key_cache,
)
# Match chat auth: group budget rows contribute max_budget, not rpm/tpm, to this gate.
return any(_live_budget_configured(budget, zero_is_limit=False) for budget in budgets.values())
# The shared group-budget helper flattens the row down to spend and max_budget, which would
# drop the rpm and tpm limits this gate exists to refuse, so the linked row is read in full.
limits: Final = await _live_model_group_limits(matched_groups)
return any(_live_budget_configured(limit, zero_is_limit=False) for limit in limits)
async def _managed_member_budget(auth: UserAPIKeyAuth, model: str | None = None) -> bool:

View file

@ -1440,7 +1440,94 @@ async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backen
assert await live._managed_member_budget(auth, model="backend") is blocked
assert await live._managed_member_budget(auth, model="backend") is blocked
db.litellm_modelaccessgroupbudgettable.find_many.assert_awaited_once()
assert cache.async_set_cache.await_args.kwargs["key"] == "model_access_group:backend-group"
assert cache.async_set_cache.await_args.kwargs["key"] == "live:model_access_group_limits:backend-group"
@pytest.mark.asyncio
@pytest.mark.parametrize("limit_field", ["rpm_limit", "tpm_limit"])
async def test_managed_budget_blocks_a_delegated_group_rate_limit_without_a_budget(monkeypatch, limit_field):
from litellm.proxy import proxy_server
group_budget = SimpleNamespace(max_budget=None, rpm_limit=None, tpm_limit=None)
setattr(group_budget, limit_field, 100)
db = SimpleNamespace(
litellm_modelaccessgroupbudgettable=SimpleNamespace(
find_many=AsyncMock(
return_value=[SimpleNamespace(access_group_name="voice-group", litellm_budget_table=group_budget)]
)
),
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache())
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace())
monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("voice-group",)))
assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", models=["voice"]), model="backend") is True
@pytest.mark.asyncio
async def test_managed_budget_fails_closed_when_the_team_row_is_unreadable(monkeypatch):
from litellm.proxy import proxy_server
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable"))),
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache())
monkeypatch.setattr(proxy_server, "llm_router", None)
with pytest.raises(HTTPException) as rejected:
await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member"))
assert rejected.value.status_code == 503
@pytest.mark.asyncio
async def test_managed_budget_fails_closed_when_the_default_budget_is_unreadable(monkeypatch):
from litellm.proxy import proxy_server
team = LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "budget-1"})
db = SimpleNamespace(
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)),
litellm_budgettable=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable"))),
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", _auth_cache())
monkeypatch.setattr(proxy_server, "llm_router", None)
with pytest.raises(HTTPException) as rejected:
await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member"))
assert rejected.value.status_code == 503
@pytest.mark.asyncio
async def test_managed_budget_caches_the_team_and_its_default_budget(monkeypatch):
from litellm.proxy import proxy_server
team = LiteLLM_TeamTable(team_id="team", metadata={"team_member_budget_id": "budget-1"})
budget = LiteLLM_BudgetTable(max_budget=5)
team_lookup = AsyncMock(return_value=team)
budget_lookup = AsyncMock(return_value=budget)
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
litellm_teamtable=SimpleNamespace(find_unique=team_lookup),
litellm_budgettable=SimpleNamespace(find_unique=budget_lookup),
)
cache = _auth_cache()
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "llm_router", None)
auth = UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")
assert await live._managed_member_budget(auth) is True
assert await live._managed_member_budget(auth) is True
team_lookup.assert_awaited_once()
budget_lookup.assert_awaited_once()
cached_keys: set[str] = {call.kwargs["key"] for call in cache.async_set_cache.await_args_list}
assert {"team_id:team", "team_member_default_budget:budget-1"} <= cached_keys
@pytest.mark.asyncio