fix(proxy): reserve project budget so concurrent requests cannot race past it

Addresses review feedback: dropping the team member reservation for
project-scoped keys left no reservation in its place, so overlapping
requests could all pass against the same cached project spend and
collectively exceed the project budget before the batched spend flush.

Adds a spend:project:{id} counter: reserved pre-call in
budget_reservation, incremented in increment_spend_counters, reseeded
from LiteLLM_ProjectTable in SpendCounterReseed, and read
(Redis-first) by _project_max_budget_check so enforcement is
near-real-time instead of waiting on the DB flush
This commit is contained in:
Shivam Rawat 2026-08-03 16:59:34 -07:00
parent 3d2932edb6
commit ae5d90d486
8 changed files with 161 additions and 9 deletions

View file

@ -4176,16 +4176,22 @@ async def _project_max_budget_check(
if project_object.litellm_budget_table is not None:
max_budget = project_object.litellm_budget_table.max_budget
if (
max_budget is not None
and project_object.spend is not None
and math.isfinite(max_budget)
and project_object.spend > max_budget
):
if max_budget is None or not math.isfinite(max_budget):
return
from litellm.proxy.proxy_server import get_current_spend
project_spend = await get_current_spend(
counter_key=f"spend:project:{project_object.project_id}",
fallback_spend=project_object.spend or 0.0,
max_budget=max_budget,
)
if project_spend > max_budget:
if valid_token:
call_info = CallInfo(
token=valid_token.token,
spend=project_object.spend,
spend=project_spend,
max_budget=max_budget,
user_id=valid_token.user_id,
team_id=valid_token.team_id,
@ -4201,9 +4207,9 @@ async def _project_max_budget_check(
)
raise litellm.BudgetExceededError(
current_cost=project_object.spend,
current_cost=project_spend,
max_budget=max_budget,
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_object.spend}, Max budget: {max_budget}",
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_spend}, Max budget: {max_budget}",
entity_type=Litellm_EntityType.PROJECT.value,
entity_id=project_object.project_id,
)

View file

@ -21,6 +21,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
SpendLogsRepository,
TeamMembershipRepository,
@ -47,6 +48,7 @@ class SpendCounterReseed:
spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend
spend:user:{user_id} -> LiteLLM_UserTable.spend
spend:org:{org_id} -> LiteLLM_OrganizationTable.spend
spend:project:{project_id} -> LiteLLM_ProjectTable.spend
End-user and tag spend counters intentionally do not reseed here. Their
auth paths already load the corresponding objects via get_end_user_object()
@ -111,6 +113,9 @@ class SpendCounterReseed:
elif counter_key.startswith("spend:org:"):
org_id = counter_key[len("spend:org:") :]
row = await OrganizationRepository(prisma_client).table.find_unique(where={"organization_id": org_id})
elif counter_key.startswith("spend:project:"):
project_id = counter_key[len("spend:project:") :]
row = await ProjectRepository(prisma_client).table.find_unique(where={"project_id": project_id})
else:
return None
except Exception:

View file

@ -533,6 +533,7 @@ async def _update_database_and_spend_counters(
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
project_id=project_id,
)
except Exception:
if budget_reservation is not None:

View file

@ -2370,6 +2370,7 @@ async def increment_spend_counters(
budget_reservation: dict | None = None,
end_user_id: str | None = None,
tags: list[str] | None = None,
project_id: str | None = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -2484,6 +2485,16 @@ async def increment_spend_counters(
increment=cost,
)
async def _project_scope(scope_project_id: str) -> None:
project_counter_key = f"spend:project:{scope_project_id}"
if project_counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=project_counter_key,
source_cache_key=f"project_id:{scope_project_id}",
increment=cost,
)
scope_coros = tuple(
coro
for coro in (
@ -2491,6 +2502,7 @@ async def increment_spend_counters(
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
_user_scope(user_id) if user_id is not None else None,
_project_scope(project_id) if project_id is not None else None,
_increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,

View file

@ -48,6 +48,7 @@ _COUNTER_ENTITY_TYPES: Mapping[str, str] = {
"EndUser": Litellm_EntityType.END_USER.value,
"Tag": Litellm_EntityType.TAG.value,
"Organization": Litellm_EntityType.ORGANIZATION.value,
"Project": Litellm_EntityType.PROJECT.value,
}
@ -432,6 +433,13 @@ async def _get_budget_counters(
if org_counter is not None:
counters.append(org_counter)
project_counter = await _get_project_budget_counter(
valid_token=valid_token,
user_api_key_cache=user_api_key_cache,
)
if project_counter is not None:
counters.append(project_counter)
return counters
@ -564,6 +572,33 @@ async def _get_team_member_budget_counter(
)
async def _get_project_budget_counter(
valid_token: UserAPIKeyAuth,
user_api_key_cache: DualCache,
) -> _BudgetCounter | None:
if valid_token.project_id is None:
return None
source_cache_key = f"project_id:{valid_token.project_id}"
project_object = await user_api_key_cache.async_get_cache(key=source_cache_key)
if project_object is None:
return None
budget_table = _get_value(project_object, "litellm_budget_table")
max_budget = _to_float(_get_value(budget_table, "max_budget"))
if max_budget is None or max_budget <= 0:
return None
return _BudgetCounter(
counter_key=f"spend:project:{valid_token.project_id}",
source_cache_key=source_cache_key,
max_budget=max_budget,
fallback_spend=_to_float(_get_value(project_object, "spend")) or 0.0,
entity_type="Project",
entity_id=valid_token.project_id,
)
async def _get_org_budget_counter(
valid_token: UserAPIKeyAuth,
team_object: LiteLLM_TeamTable | None,

View file

@ -474,6 +474,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
budget_reservation=budget_reservation,
end_user_id="test_end_user_id",
tags=["tag-a"],
project_id=None,
)

View file

@ -1270,3 +1270,50 @@ async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkey
)
assert result is None
@pytest.mark.asyncio
async def test_increment_spend_counters_increments_project_counter(monkeypatch):
"""A request attributed to a project must increment spend:project:{id} so
concurrent project-scoped requests see near-real-time project spend."""
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=5.0
)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
async def _fake_coalesced(**kwargs):
return None
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced)
)
await ps.increment_spend_counters(
token="hashed-tok",
team_id="t1",
user_id="u1",
response_cost=5.0,
project_id="p1",
)
incremented_keys = {
call.kwargs["key"]
for call in fake_cache.redis_cache.async_increment.call_args_list
}
assert "spend:project:p1" in incremented_keys
fake_cache.redis_cache.async_increment.reset_mock()
await ps.increment_spend_counters(
token="hashed-tok",
team_id="t1",
user_id="u1",
response_cost=5.0,
)
incremented_keys = {
call.kwargs["key"]
for call in fake_cache.redis_cache.async_increment.call_args_list
}
assert not any(key.startswith("spend:project:") for key in incremented_keys)

View file

@ -2585,3 +2585,48 @@ async def test_team_member_budget_counter_skipped_for_project_scoped_key():
user_api_key_cache=cache,
)
assert counter is None
@pytest.mark.asyncio
async def test_project_budget_counter_reserved_for_project_scoped_key():
"""Project-scoped keys must reserve against the project budget so
concurrent requests cannot collectively race past it."""
from litellm.proxy.spend_tracking.budget_reservation import (
_get_project_budget_counter,
)
cache = MagicMock()
cache.async_get_cache = AsyncMock(
return_value={
"project_id": "project-1",
"spend": 0.4,
"litellm_budget_table": {"max_budget": 10.0},
}
)
counter = await _get_project_budget_counter(
valid_token=UserAPIKeyAuth(token="hashed", project_id="project-1"),
user_api_key_cache=cache,
)
assert counter is not None
assert counter.counter_key == "spend:project:project-1"
assert counter.source_cache_key == "project_id:project-1"
assert counter.max_budget == 10.0
assert counter.fallback_spend == 0.4
assert counter.entity_type == "Project"
counter = await _get_project_budget_counter(
valid_token=UserAPIKeyAuth(token="hashed"),
user_api_key_cache=cache,
)
assert counter is None
cache.async_get_cache = AsyncMock(
return_value={"project_id": "project-1", "spend": 0.4, "litellm_budget_table": None}
)
counter = await _get_project_budget_counter(
valid_token=UserAPIKeyAuth(token="hashed", project_id="project-1"),
user_api_key_cache=cache,
)
assert counter is None