fix(proxy): reset a stuck team member's budget (#37971)

* fix(proxy): reset a stuck team member's budget

A per-team-member budget check reads a cross-pod spend counter that
nothing ever invalidates. Once a member exceeds their per-member
budget, resetting the key's spend, raising the user's or the team's
own budget, or issuing a new key all leave the member stuck, because
none of them touch this counter or its cached membership object.

Add POST /team/{team_id}/member/{user_id}/reset_spend to reset a
member's tracked spend, and invalidate the same cached state from
/team/member_update when it raises a member's own budget, so that
path also takes effect immediately instead of waiting on the
membership cache's TTL. Name the entity in the check's error message
so a stuck member is diagnosable from the 429 body alone.

* fix(proxy): close reset-vs-floor-read race and surface double Redis write failure on member spend reset

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): broadcast spend reset as a SET so the handler's self-delivered message cannot erase the reset guard

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): omit null fields from the invalidation message so plain evictions keep the old wire format

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yassin Kortam 2026-08-25 09:50:09 -07:00 committed by GitHub
parent bb27bfd9a7
commit 1d695a714b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1324 additions and 38 deletions

View file

@ -815,6 +815,7 @@ class LiteLLMRoutes(enum.Enum):
"/team/member_add",
"/team/member_delete",
"/team/member_update",
"/team/{team_id}/member/{user_id}/reset_spend",
"/team/permissions_list",
"/team/permissions_update",
"/team/daily/activity",
@ -1287,6 +1288,16 @@ class RegenerateKeyRequest(GenerateKeyRequest):
class ResetSpendRequest(LiteLLMPydanticObjectBase):
reset_to: float
@field_validator("reset_to", mode="before")
@classmethod
def reject_bool_reset_to(cls, v):
# bool is a subclass of int, so pydantic silently coerces True/False into
# 1.0/0.0 for a `float` field: a caller who accidentally sends a boolean
# would otherwise get an unintended spend reset instead of a 422.
if isinstance(v, bool):
raise ValueError("reset_to must be a number, not a boolean") # noqa: TRY004 # pydantic needs ValueError
return v
class KeyRequest(LiteLLMPydanticObjectBase):
keys: list[str] | None = None

View file

@ -87,6 +87,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
object_permission_cache_key,
tag_cache_key,
tag_registry_cache_key,
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
@ -1967,7 +1969,7 @@ async def get_team_membership(
if user_id is None or team_id is None:
return None
_key: Final = f"team_membership:{user_id}:{team_id}"
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
# check if in cache
cached_membership_obj: Final = await user_api_key_cache.async_get_cache(
@ -2402,6 +2404,116 @@ async def _cache_team_object(
)
async def invalidate_team_member_spend_state(
user_id: str,
team_id: str,
user_api_key_cache: UserApiKeyCache,
new_spend: float | None = None,
) -> None:
"""
Clear every cached read path for one team member's budget so a spend
reset or a raised cap takes effect on the next request instead of
waiting on the membership cache's TTL.
Two independently-keyed cache entries hold the same LiteLLM_TeamMembership
row: user_api_key_auth.py's admission check writes ``{team_id}_{user_id}``,
while budget_reservation.py's pre-call reservation and auth_checks.py's own
get_team_membership() (used by _check_team_member_budget) both write
``team_membership:{user_id}:{team_id}``. Both formats must be invalidated
explicitly; writing one does not refresh the other. All keys are also
broadcast (LIT-3803): each worker's own in-memory copy (membership object,
spend counter, or the counter's own short-TTL DB-floor marker) survives
eviction elsewhere until its TTL, so the handling worker alone clearing its
copy leaves every other worker still enforcing the pre-reset budget.
``new_spend`` is only passed by reset_team_member_spend_fn, which knows the
exact post-reset value: it is SET everywhere (matching /key/{key}/reset_spend's
own precedent) rather than deleted, so a worker's next read reflects it
directly instead of re-deriving it through a DB reseed. team_member_update
only changes the budget cap, not the tracked spend, so it passes no
new_spend; the live spend counter is untouched in that case (deleting it
would force a reseed from the DB's own spend column, which lags the live
counter via periodic batch writes, briefly under-enforcing the raised cap
against a spend value lower than what was actually tracked) and only the
membership caches carrying the new cap are invalidated.
The floor marker (``spend_db_floor:``, proxy_server.py's
_authoritative_floor_spend) caches the pre-reset DB spend for
SPEND_DB_FLOOR_CACHE_TTL_SECONDS; left stale after a real reset, a request
landing on the pod that cached it can read that higher floor and raise the
counter right back above the just-reset spend. It is overwritten here with
the post-reset floor (not merely deleted) and _authoritative_floor_spend
re-checks the marker after its DB read, so a floor read already in flight
on this pod when the reset commits cannot clobber it with the pre-reset
value. Both keys are broadcast as SETs carrying new_spend, not deletes:
every subscriber (remote pods AND this pod's own, which receives its own
message) writes the post-reset value, so the self-delivered message cannot
erase the guard just written here.
Raises HTTPException(503) if Redis still holds the stale pre-reset counter
after both the SET and the fallback DELETE fail: budget checks read Redis
first, so returning success would leave the old value authoritative for
every worker despite the DB write having committed.
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
evict_and_broadcast,
publish_auth_cache_invalidation,
)
if new_spend is not None:
from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache
spend_counter_key: Final = f"spend:team_member:{user_id}:{team_id}"
spend_db_floor_key: Final = f"spend_db_floor:{spend_counter_key}"
spend_counter_cache.in_memory_cache.set_cache(key=spend_counter_key, value=new_spend, ttl=60)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(key=spend_counter_key, value=new_spend, ttl=60)
except Exception as e: # noqa: BLE001 # fall back to deleting the stale entry before giving up
verbose_proxy_logger.warning(
"Failed to set spend counter %s in Redis after reset: %s; deleting it instead so the next "
"read reseeds from the DB rather than keeping the stale pre-reset value authoritative",
spend_counter_key,
e,
)
try:
await spend_counter_cache.redis_cache.async_delete_cache(key=spend_counter_key)
except Exception: # noqa: BLE001 # stale value now authoritative in Redis; surface instead of reporting success
verbose_proxy_logger.warning(
"Failed to delete stale spend counter %s in Redis after a failed reset write",
spend_counter_key,
exc_info=True,
)
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail={ # mutable-ok: HTTPException.detail takes a dict
"error": "Spend was reset in the database, but Redis is unreachable and still "
"holds the pre-reset counter. Retry once Redis is reachable."
},
) from e
spend_counter_cache.in_memory_cache.set_cache(
key=spend_db_floor_key,
value=new_spend,
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
)
await publish_auth_cache_invalidation(cache_key=spend_counter_key, new_value=new_spend, ttl=60)
await publish_auth_cache_invalidation(
cache_key=spend_db_floor_key,
new_value=new_spend,
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
)
await evict_and_broadcast(
cache_keys=(
team_membership_auth_cache_key(team_id=team_id, user_id=user_id),
team_membership_reservation_cache_key(user_id=user_id, team_id=team_id),
),
user_api_key_cache=user_api_key_cache,
)
async def delete_cache_team_object(
team_id: str,
team_alias: str | None,

View file

@ -87,7 +87,10 @@ from litellm.proxy.common_utils.http_parsing_utils import (
populate_request_with_path_params,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
team_membership_auth_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.utils import (
@ -1970,8 +1973,10 @@ async def _user_api_key_auth_builder(
# Check 3. Check if user is in their team budget
if not skip_budget_checks and valid_token.team_member_spend is not None:
if prisma_client is not None:
_cache_key: Final = f"{valid_token.team_id}_{valid_token.user_id}"
_user_id: Final = valid_token.user_id
_team_id: Final = valid_token.team_id
if prisma_client is not None and _user_id is not None and _team_id is not None:
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
team_member_info = await user_api_key_cache.async_get_cache(
key=_cache_key,
@ -1979,25 +1984,21 @@ async def _user_api_key_auth_builder(
)
if team_member_info is None:
# read from DB
_user_id: Final = valid_token.user_id
_team_id: Final = valid_token.team_id
if _user_id is not None and _team_id is not None:
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
where={
"user_id": _user_id,
"team_id": _team_id,
},
include={"litellm_budget_table": True},
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
where={
"user_id": _user_id,
"team_id": _team_id,
},
include={"litellm_budget_table": True},
)
if _db_member is not None:
team_member_info = LiteLLM_TeamMembership(**_db_member.dict())
await user_api_key_cache.async_set_cache(
key=_cache_key,
value=team_member_info,
model_type=LiteLLM_TeamMembership,
ttl=5,
)
if _db_member is not None:
team_member_info = LiteLLM_TeamMembership(**_db_member.dict())
await user_api_key_cache.async_set_cache(
key=_cache_key,
value=team_member_info,
model_type=LiteLLM_TeamMembership,
ttl=5,
)
if team_member_info is not None and team_member_info.litellm_budget_table is not None:
team_member_budget: Final = team_member_info.litellm_budget_table.max_budget
@ -2013,11 +2014,16 @@ async def _user_api_key_auth_builder(
max_budget=team_member_budget,
)
if team_member_spend > team_member_budget:
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
raise litellm.BudgetExceededError(
current_cost=team_member_spend,
max_budget=team_member_budget,
message=(
f"Budget has been exceeded! TeamMember={_entity_id} "
f"Current cost: {team_member_spend}, Max budget: {team_member_budget}"
),
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
entity_id=f"{valid_token.user_id}:{valid_token.team_id}",
entity_id=_entity_id,
)
# Check 3. If token is expired

View file

@ -12,6 +12,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
)
if TYPE_CHECKING:
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -30,15 +31,24 @@ def auth_cache_invalidation_channel(redis_cache: "RedisCache") -> str:
@dataclass(frozen=True, slots=True)
class _CacheInvalidationMessage:
cache_key: str
new_value: float | None = None
ttl: float | None = None
def _cache_invalidation_message_json(cache_key: str) -> str:
return json.dumps(asdict(_CacheInvalidationMessage(cache_key=cache_key)))
def _cache_invalidation_message_json(cache_key: str, new_value: float | None = None, ttl: float | None = None) -> str:
message: Final = asdict(_CacheInvalidationMessage(cache_key=cache_key, new_value=new_value, ttl=ttl))
return json.dumps({field: value for field, value in message.items() if value is not None})
def _cache_key_from_message_data(data: object) -> str | None:
def _finite_number_or_none(value: object) -> float | None:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return None
return float(value)
def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
if isinstance(data, bytes):
data = data.decode("utf-8", errors="replace")
data = data.decode("utf-8", errors="replace") # rebind-ok: normalizing the wire payload to str
if not isinstance(data, str):
return None
try:
@ -48,14 +58,28 @@ def _cache_key_from_message_data(data: object) -> str | None:
if not isinstance(parsed, dict):
return None
cache_key: Final = parsed.get("cache_key")
return cache_key if isinstance(cache_key, str) else None
if not isinstance(cache_key, str):
return None
return _CacheInvalidationMessage(
cache_key=cache_key,
new_value=_finite_number_or_none(parsed.get("new_value")),
ttl=_finite_number_or_none(parsed.get("ttl")),
)
async def publish_auth_cache_invalidation(cache_key: str) -> None:
async def publish_auth_cache_invalidation(
cache_key: str, new_value: float | None = None, ttl: float | None = None
) -> None:
"""
Best-effort broadcast so every worker drops its local in-memory copy of a
mutated management object; without this, only the handling worker and Redis
are evicted and other workers keep serving the stale object until its TTL.
Passing ``new_value`` broadcasts a SET instead of a delete: every subscriber
(including the publishing worker's own, which receives its own message)
writes the value into its additional in-memory caches rather than deleting
the key. A spend reset uses this so the handler's self-delivered message
cannot erase the freshly-written post-reset counter or floor marker.
"""
redis_cache: Final = coordination_redis_cache()
if redis_cache is None:
@ -68,7 +92,10 @@ async def publish_auth_cache_invalidation(cache_key: str) -> None:
cache_key,
)
return
await client.publish(auth_cache_invalidation_channel(redis_cache), _cache_invalidation_message_json(cache_key))
await client.publish(
auth_cache_invalidation_channel(redis_cache),
_cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl),
)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
@ -95,15 +122,17 @@ async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "Us
class AuthCacheInvalidationSubscriber:
__slots__ = ("_redis_cache", "_task", "_user_api_key_cache")
__slots__ = ("_additional_in_memory_caches", "_redis_cache", "_task", "_user_api_key_cache")
def __init__(
self,
redis_cache: "RedisCache",
user_api_key_cache: "UserApiKeyCache",
additional_in_memory_caches: Sequence["InMemoryCache"] = (),
) -> None:
self._redis_cache = redis_cache
self._user_api_key_cache = user_api_key_cache
self._additional_in_memory_caches = tuple(additional_in_memory_caches)
self._task: asyncio.Task[None] | None = None
def start(self) -> None:
@ -160,12 +189,18 @@ class AuthCacheInvalidationSubscriber:
def _apply_message(self, message: object) -> None:
data: Final = message.get("data") if isinstance(message, dict) else None
cache_key: Final = _cache_key_from_message_data(data)
if cache_key is None:
parsed: Final = _message_from_data(data)
if parsed is None:
return
if parsed.new_value is not None:
for additional_cache in self._additional_in_memory_caches:
additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl)
return
in_memory_cache: Final = self._user_api_key_cache.in_memory_cache
if in_memory_cache is not None:
in_memory_cache.delete_cache(cache_key)
in_memory_cache.delete_cache(parsed.cache_key)
for additional_cache in self._additional_in_memory_caches:
additional_cache.delete_cache(parsed.cache_key)
@staticmethod
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:

View file

@ -200,6 +200,21 @@ def end_user_restricted_registry_cache_key() -> str:
return "end_user_restricted_registry"
def team_membership_auth_cache_key(team_id: str, user_id: str) -> str:
"""Cache key one team member's ``LiteLLM_TeamMembership`` row is stored under for the admission check."""
return f"{team_id}_{user_id}"
def team_membership_reservation_cache_key(user_id: str, team_id: str) -> str:
"""Cache key the pre-call budget reservation stores the same ``LiteLLM_TeamMembership`` row under.
Deliberately not unified with ``team_membership_auth_cache_key``: the two readers wrote independent
keys before this file existed, so a fix that invalidates one must invalidate both explicitly rather
than assume a single write is visible to both.
"""
return f"team_membership:{user_id}:{team_id}"
def get_management_object_ttl(cache: DualCache) -> float:
"""
In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...).

View file

@ -16,7 +16,7 @@ import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Annotated, Final, NamedTuple, Protocol, TypedDict, TypeVar, cast
from typing import Annotated, Final, NamedTuple, NoReturn, Protocol, TypedDict, TypeVar, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -56,6 +56,7 @@ from litellm.proxy._types import (
PatchTeamRequest,
ProxyErrorTypes,
ProxyException,
ResetSpendRequest,
SpecialManagementEndpointEnums,
SpecialModelNames,
SpecialProxyStrings,
@ -84,6 +85,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_membership,
get_team_object,
get_user_object,
invalidate_team_member_spend_state,
)
from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
@ -3392,7 +3394,7 @@ async def team_member_update(
Update team member budgets and team member role
"""
from litellm.proxy.proxy_server import premium_user, prisma_client
from litellm.proxy.proxy_server import premium_user, prisma_client, user_api_key_cache
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
@ -3491,6 +3493,12 @@ async def team_member_update(
budget_patch=budget_patch,
team_default_budget_id=team_default_budget_id,
)
if budget_patch:
await invalidate_team_member_spend_state(
user_id=received_user_id,
team_id=data.team_id,
user_api_key_cache=user_api_key_cache,
)
### update team member role
if data.role is not None:
@ -3527,6 +3535,125 @@ async def team_member_update(
)
def _check_not_resetting_own_spend(user_id: str, user_api_key_dict: UserAPIKeyAuth) -> None:
"""
_verify_team_access authorizes a team admin (or org admin) over their own
team, with no check that the target user_id differs from the caller. Left
unchecked, that admin could target their own LiteLLM_TeamMembership row and
repeatedly reset it to 0 right before it crosses their per-member cap,
consuming the shared team budget without the configured limit ever binding.
Only a proxy admin may reset an admin's own spend.
"""
if user_id == user_api_key_dict.user_id and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
_raise_reset_spend_error(status.HTTP_403_FORBIDDEN, "Cannot reset your own spend. Ask a proxy admin.")
def _raise_reset_spend_error(status_code: int, message: str) -> NoReturn:
detail: Final = {"error": message} # mutable-ok: HTTPException.detail takes a dict
raise HTTPException(status_code=status_code, detail=detail)
def _validate_team_member_reset_spend_value(
reset_to: object,
membership: LiteLLM_TeamMembership,
) -> float:
if not isinstance(reset_to, (int, float)):
_raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a float")
reset_to_float: Final = float(reset_to)
if not math.isfinite(reset_to_float) or reset_to_float < 0:
_raise_reset_spend_error(status.HTTP_400_BAD_REQUEST, "reset_to must be a finite number >= 0")
current_spend: Final = membership.spend or 0.0
if reset_to_float > current_spend:
_raise_reset_spend_error(
status.HTTP_400_BAD_REQUEST,
f"reset_to ({reset_to_float}) must be <= current spend ({current_spend})",
)
max_budget: Final = membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None
if max_budget is not None and reset_to_float > max_budget:
_raise_reset_spend_error(
status.HTTP_400_BAD_REQUEST,
f"reset_to ({reset_to_float}) must be <= budget ({max_budget})",
)
return reset_to_float
@router.post(
"/team/{team_id}/member/{user_id}/reset_spend",
tags=["team management"], # mutable-ok: FastAPI's `tags` param is typed as list[str], not Sequence
dependencies=(Depends(user_api_key_auth),),
)
@management_endpoint_wrapper
async def reset_team_member_spend_fn(
team_id: str,
user_id: str,
data: ResetSpendRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
Reset a team member's tracked spend against their per-member budget.
A member's spend is tracked separately from both their own personal
budget and the team's own budget (LiteLLM_TeamMembership.spend), so
neither /user/update nor /team/update can clear it: this is the only
endpoint that does. The cross-pod spend counter and cached membership
reads are invalidated so the reset takes effect on the member's next
request rather than waiting on the membership cache's TTL.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
_raise_reset_spend_error(status.HTTP_500_INTERNAL_SERVER_ERROR, "DB not connected. prisma_client is None")
team_obj: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
check_db_only=True,
)
await _verify_team_access(team_obj=team_obj, user_api_key_dict=user_api_key_dict)
_check_not_resetting_own_spend(user_id=user_id, user_api_key_dict=user_api_key_dict)
membership_where: Final = { # mutable-ok: prisma client requires a plain dict where= argument
"user_id_team_id": {"user_id": user_id, "team_id": team_id} # mutable-ok: same prisma where= argument
}
_membership_row: Final = await _team_membership_db(prisma_client).find_unique(
where=membership_where,
include={"litellm_budget_table": True}, # mutable-ok: prisma client requires a plain dict include= argument
)
if _membership_row is None:
_raise_reset_spend_error(status.HTTP_404_NOT_FOUND, f"User {user_id} is not a member of team {team_id}.")
membership: Final = LiteLLM_TeamMembership.model_validate(_membership_row.model_dump())
current_spend: Final = membership.spend or 0.0
reset_to: Final = _validate_team_member_reset_spend_value(data.reset_to, membership)
await _team_membership_db(prisma_client).update(
where=membership_where,
data={"spend": reset_to}, # mutable-ok: prisma client requires a plain dict data= argument
)
await invalidate_team_member_spend_state(
user_id=user_id,
team_id=team_id,
user_api_key_cache=user_api_key_cache,
new_spend=reset_to,
)
return { # mutable-ok: matches this router's established untyped-response-dict convention
"team_id": team_id,
"user_id": user_id,
"spend": reset_to,
"previous_spend": current_spend,
"max_budget": membership.litellm_budget_table.max_budget if membership.litellm_budget_table else None,
}
def _create_results_from_response(
members: list[Member],
response: TeamAddMemberResponse,

View file

@ -2555,6 +2555,12 @@ async def _authoritative_floor_spend(
if db_spend is None:
return None
# a spend reset that committed during the DB read above wrote the post-reset
# floor to the marker; keep it over this read's now-stale pre-commit value
rechecked: Final = spend_counter_cache.in_memory_cache.get_cache(key=marker_key)
if rechecked is not None:
return float(rechecked)
spend_counter_cache.in_memory_cache.set_cache(
key=marker_key,
value=db_spend,
@ -6798,6 +6804,7 @@ class ProxyConfig:
subscriber: Final = AuthCacheInvalidationSubscriber(
redis_cache=redis_cache,
user_api_key_cache=user_api_key_cache,
additional_in_memory_caches=(spend_counter_cache.in_memory_cache,),
)
self.auth_cache_invalidation_subscriber = subscriber
subscriber.start()

View file

@ -25,7 +25,11 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_utils import get_model_from_request
from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.user_api_key_cache import end_user_cache_key, tag_cache_key
from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
tag_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router import Router
@ -546,7 +550,9 @@ async def _get_team_member_budget_counter(
if team_object is None or team_object.team_id is None or user_object is None or valid_token.user_id is None:
return None
membership_cache_key: Final = f"team_membership:{valid_token.user_id}:{team_object.team_id}"
membership_cache_key: Final = team_membership_reservation_cache_key(
user_id=valid_token.user_id, team_id=team_object.team_id
)
cached_team_membership: Final = await user_api_key_cache.async_get_cache(key=membership_cache_key)
team_membership: LiteLLM_TeamMembership | None = None
if isinstance(cached_team_membership, LiteLLM_TeamMembership):

View file

@ -0,0 +1,152 @@
import uuid
import pytest
from .actors import Actor
from .conftest import create_scratch_team
pytestmark = pytest.mark.asyncio(loop_scope="session")
_SEED_SPEND = 5.0
_RESET_TO = 2.0
# POST /team/{team_id}/member/{user_id}/reset_spend. The handler gate is
# _verify_team_access (proxy admin / team admin of this team / org admin of
# the team's org) — the same gate /team/member_update uses, so this mirrors
# that file's matrix exactly.
_MATRIX = [
("alpha/proxy_admin", Actor.PROXY_ADMIN, "alpha", 200),
("alpha/org_admin", Actor.ORG_ADMIN, "alpha", 200),
("alpha/team_admin", Actor.TEAM_ADMIN, "alpha", 200),
("alpha/internal_user", Actor.INTERNAL_USER, "alpha", 403),
("alpha/owner", Actor.OWNER, "alpha", 403),
("alpha/unrelated_same_org", Actor.UNRELATED_SAME_ORG, "alpha", 403),
("alpha/cross_org_user", Actor.CROSS_ORG_USER, "alpha", 403),
("alpha/service_account", Actor.SERVICE_ACCOUNT, "alpha", 403),
("alpha/org_b_admin", Actor.ORG_B_ADMIN, "alpha", 403),
("beta/proxy_admin", Actor.PROXY_ADMIN, "beta", 200),
("beta/org_admin", Actor.ORG_ADMIN, "beta", 403),
("beta/team_admin", Actor.TEAM_ADMIN, "beta", 403),
("beta/org_b_admin", Actor.ORG_B_ADMIN, "beta", 200),
]
async def _seed_target(prisma, world, shape: str, team_id: str, member_id: str) -> None:
if shape == "alpha":
await create_scratch_team(
prisma,
team_id,
organization_id=world.org_a_id,
admin_user_ids=[world.keys[Actor.TEAM_ADMIN].user_id],
)
elif shape == "beta":
await create_scratch_team(prisma, team_id, organization_id=world.org_b_id)
else: # pragma: no cover - guard
pytest.fail(f"unknown shape={shape}")
await prisma.db.litellm_teammembership.create(
data={"user_id": member_id, "team_id": team_id, "spend": _SEED_SPEND}
)
@pytest.mark.parametrize(
"actor,shape,expected_status",
[(a, sh, s) for (_id, a, sh, s) in _MATRIX],
ids=[s[0] for s in _MATRIX],
)
async def test_team_member_reset_spend_authz_matrix(
actor: Actor,
shape: str,
expected_status: int,
proxy_client,
prisma,
scratch,
world,
):
member_id = scratch.tag("member")
await _seed_target(prisma, world, shape, scratch.prefix, member_id)
caller = world.keys[actor]
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{member_id}/reset_spend",
headers={"Authorization": f"Bearer {caller.cleartext}"},
json={"reset_to": _RESET_TO},
)
assert (
resp.status_code == expected_status
), f"{actor.value} {shape}: {resp.status_code} {resp.text}"
row = await prisma.db.litellm_teammembership.find_unique(
where={"user_id_team_id": {"user_id": member_id, "team_id": scratch.prefix}}
)
assert row is not None
if expected_status == 200:
assert row.spend == _RESET_TO
else:
assert row.spend == _SEED_SPEND, "denied but spend reset"
async def test_team_member_reset_spend_missing_team_is_404(proxy_client, world):
resp = await proxy_client.post(
f"/team/behavior-pin-no-such-team/member/{uuid.uuid4().hex}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 404, resp.text
async def test_team_member_reset_spend_missing_membership_is_404(
proxy_client, prisma, scratch, world
):
"""A well-formed team but a user_id with no LiteLLM_TeamMembership row is 404."""
await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{uuid.uuid4().hex}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 404, resp.text
async def test_team_member_reset_spend_above_current_spend_is_400(
proxy_client, prisma, scratch, world
):
member_id = scratch.tag("member")
await create_scratch_team(prisma, scratch.prefix, organization_id=world.org_a_id)
await prisma.db.litellm_teammembership.create(
data={"user_id": member_id, "team_id": scratch.prefix, "spend": 1.0}
)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{member_id}/reset_spend",
headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"},
json={"reset_to": 5.0},
)
assert resp.status_code == 400, resp.text
async def test_team_member_reset_spend_team_admin_cannot_reset_own_spend(
proxy_client, prisma, scratch, world
):
"""A team admin targeting their own LiteLLM_TeamMembership row is 403: unchecked, an
admin could repeatedly zero their own spend right before it crosses their per-member
cap, consuming the shared team budget without the configured limit ever binding."""
team_admin = world.keys[Actor.TEAM_ADMIN]
await create_scratch_team(
prisma,
scratch.prefix,
organization_id=world.org_a_id,
admin_user_ids=[team_admin.user_id],
)
await prisma.db.litellm_teammembership.create(
data={"user_id": team_admin.user_id, "team_id": scratch.prefix, "spend": _SEED_SPEND}
)
resp = await proxy_client.post(
f"/team/{scratch.prefix}/member/{team_admin.user_id}/reset_spend",
headers={"Authorization": f"Bearer {team_admin.cleartext}"},
json={"reset_to": 0.0},
)
assert resp.status_code == 403, resp.text
row = await prisma.db.litellm_teammembership.find_unique(
where={"user_id_team_id": {"user_id": team_admin.user_id, "team_id": scratch.prefix}}
)
assert row is not None and row.spend == _SEED_SPEND, "denied but spend reset"

View file

@ -47,6 +47,7 @@ from litellm.proxy.auth.auth_checks import (
_virtual_key_soft_budget_check,
get_key_object,
get_user_object,
invalidate_team_member_spend_state,
vector_store_access_check,
)
from litellm.caching.in_memory_cache import InMemoryCache
@ -6939,3 +6940,307 @@ def test_model_has_no_cost_mapping_alias_to_a_group_priced_through_model_info_is
router = _router_with_a_group_priced_through_model_info()
assert model_has_no_cost_mapping(model="model-info-priced-alias", llm_router=router) is False
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_sets_the_spend_counter_and_clears_both_membership_cache_keys():
"""A team-member budget reset (new_spend passed) must SET the spend counter to the reset
value, clear its DB-floor marker, AND invalidate both independently-keyed membership caches
(user_api_key_auth.py's admission check writes one key format, budget_reservation.py and
auth_checks.py's own get_team_membership() write the other) or a stale read keeps 429ing
after the reset. Asserted against real cache reads, not mock call args, so a change that
keeps the call but drops its effect still fails."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
real_cache = UserApiKeyCache()
await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership")
await real_cache.async_set_cache(key="team_membership:user-1:team-1", value="stale-membership")
real_spend_counter_cache = DualCache()
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0)
real_spend_counter_cache.in_memory_cache.set_cache(
key="spend_db_floor:spend:team_member:user-1:team-1", value=999.0
)
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=real_cache,
new_spend=0.0,
)
assert await real_cache.async_get_cache(key="team-1_user-1") is None
assert await real_cache.async_get_cache(key="team_membership:user-1:team-1") is None
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 0.0
assert (
real_spend_counter_cache.in_memory_cache.get_cache(key="spend_db_floor:spend:team_member:user-1:team-1")
== 0.0
), "the DB-floor marker kept the pre-reset value; a stale-floor read can raise the counter right back up"
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_leaves_the_live_spend_counter_alone_without_new_spend():
"""team_member_update only changes the budget cap, not the tracked spend, so it calls
invalidate_team_member_spend_state with no new_spend. Deleting the live spend counter in that
case would force the next read to reseed from the DB's own spend column, which lags the live
counter via periodic batch writes, briefly UNDER-enforcing the raised cap against a spend
value lower than what was actually tracked (regression: PR #37971 Bugbot finding)."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
real_cache = UserApiKeyCache()
await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership")
real_spend_counter_cache = DualCache()
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0)
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=real_cache,
)
assert await real_cache.async_get_cache(key="team-1_user-1") is None
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 999.0
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_sets_new_spend_instead_of_deleting():
"""/key/{key}/reset_spend SETs its counter to the reset value rather than deleting it, so a
worker's next read reflects it directly instead of falling back through a DB reseed. A reset
caller passing new_spend must match that precedent, not merely delete the counter."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
real_cache = UserApiKeyCache()
real_spend_counter_cache = DualCache()
fake_redis_cache = MagicMock()
fake_redis_cache.async_set_cache = AsyncMock()
real_spend_counter_cache.redis_cache = fake_redis_cache
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=real_cache,
new_spend=2.5,
)
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 2.5
fake_redis_cache.async_set_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1", value=2.5, ttl=60)
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_deletes_redis_counter_when_set_fails(): # test-quality-ok: only observable effect is the fallback call on the same fake client
"""Redis reads take priority over the local in-memory copy (get_current_spend reads Redis
first), so a failed Redis SET would otherwise leave the OLD pre-reset value authoritative
for every worker even though the reset reported success. On a failed SET, the stale Redis
entry must be deleted instead, so the next read clean-misses and reseeds from the DB."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
real_cache = UserApiKeyCache()
real_spend_counter_cache = DualCache()
fake_redis_cache = MagicMock()
fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down"))
fake_redis_cache.async_delete_cache = AsyncMock()
real_spend_counter_cache.redis_cache = fake_redis_cache
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=real_cache,
new_spend=2.5,
)
fake_redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1")
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_raises_503_when_both_redis_writes_fail():
"""If the Redis SET fails AND the fallback DELETE fails, the stale pre-reset counter is still
authoritative in Redis for every worker. Reporting success would silently keep 429ing the
member, so the reset must surface a 503 instead (regression: PR #37971 Greptile finding)."""
from fastapi import HTTPException
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
real_cache = UserApiKeyCache()
real_spend_counter_cache = DualCache()
fake_redis_cache = MagicMock()
fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down"))
fake_redis_cache.async_delete_cache = AsyncMock(side_effect=ConnectionError("redis still down"))
real_spend_counter_cache.redis_cache = fake_redis_cache
with (
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
),
pytest.raises(HTTPException) as exc_info,
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=real_cache,
new_spend=2.5,
)
assert exc_info.value.status_code == 503
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_to_remote_workers():
"""The test above only proves the handling worker's own spend counter is
cleared. A remote worker's spend counter is a separate DualCache instance;
if the reset never reaches it, that worker keeps enforcing the pre-reset
spend the moment its own Redis read for the counter fails and it falls
back to its own (now-stale) in-memory copy. Drives the actual message
published onto the invalidation channel through a second, independent
AuthCacheInvalidationSubscriber standing in for that remote worker, rather
than asserting on the publish call args."""
from redis.asyncio import Redis
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
published: list[tuple[str, str]] = []
class _RecordingRedisClient(Redis):
def __init__(self) -> None:
pass
async def publish(self, channel: str, message: str) -> int:
published.append((channel, message))
return 1
class _FakeRedisCache:
def __init__(self) -> None:
self.namespace = None
def init_async_client(self) -> object:
return _RecordingRedisClient()
local_spend_counter_cache = DualCache()
remote_user_api_key_cache = UserApiKeyCache()
remote_spend_counter_in_memory_cache = InMemoryCache()
remote_spend_counter_in_memory_cache.set_cache("spend:team_member:user-1:team-1", 999.0)
remote_spend_counter_in_memory_cache.set_cache("spend_db_floor:spend:team_member:user-1:team-1", 999.0)
with (
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache
),
patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
return_value=_FakeRedisCache(),
),
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=UserApiKeyCache(),
new_spend=0.0,
)
def _published_message_for(cache_key: str) -> str:
matches = [message for _, message in published if json.loads(message)["cache_key"] == cache_key]
assert matches, f"{cache_key} never reached the cross-worker invalidation channel"
return matches[-1]
remote_subscriber = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(),
user_api_key_cache=remote_user_api_key_cache,
additional_in_memory_caches=(remote_spend_counter_in_memory_cache,),
)
for cache_key in ("spend:team_member:user-1:team-1", "spend_db_floor:spend:team_member:user-1:team-1"):
remote_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API
{"type": "message", "data": _published_message_for(cache_key)}
)
assert remote_spend_counter_in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0
assert (
remote_spend_counter_in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0
), "the DB-floor marker was not broadcast; a remote worker can re-raise the counter off its stale floor"
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_not_erase_the_reset():
"""The handling worker subscribes to the same invalidation channel it publishes on, so it
receives its own reset message. A delete-style broadcast would erase the post-reset counter
and floor marker the handler just wrote, reopening the stale-floor race the reset closed
(regression: PR #37971 Greptile finding). The broadcast carries the reset value as a SET, so
applying the self-delivered message must leave both keys at the post-reset value."""
from redis.asyncio import Redis
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
published: list[tuple[str, str]] = []
class _RecordingRedisClient(Redis):
def __init__(self) -> None:
pass
async def publish(self, channel: str, message: str) -> int:
published.append((channel, message))
return 1
class _FakeRedisCache:
def __init__(self) -> None:
self.namespace = None
def init_async_client(self) -> object:
return _RecordingRedisClient()
local_spend_counter_cache = DualCache()
local_user_api_key_cache = UserApiKeyCache()
with (
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache
),
patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
return_value=_FakeRedisCache(),
),
):
await invalidate_team_member_spend_state(
user_id="user-1",
team_id="team-1",
user_api_key_cache=local_user_api_key_cache,
new_spend=0.0,
)
own_subscriber = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(),
user_api_key_cache=local_user_api_key_cache,
additional_in_memory_caches=(local_spend_counter_cache.in_memory_cache,),
)
for _, message in published:
own_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API
{"type": "message", "data": message}
)
assert local_spend_counter_cache.in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0, (
"the handler's self-delivered broadcast erased the post-reset spend counter"
)
assert (
local_spend_counter_cache.in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0
), "the handler's self-delivered broadcast erased the post-reset floor marker, reopening the stale-floor race"

View file

@ -6,6 +6,7 @@ from unittest.mock import patch
import pytest
from redis.asyncio import Redis
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
AUTH_CACHE_INVALIDATION_CHANNEL,
AuthCacheInvalidationSubscriber,
@ -144,6 +145,37 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None:
assert pubsub.subscribed_channels == [AUTH_CACHE_INVALIDATION_CHANNEL]
@pytest.mark.asyncio
async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message() -> None:
"""
The spend-counter half of the same cross-worker gap: a remote worker's own
spend counter can hold a stale value (its fallback path when that worker's
own Redis read for the counter fails), and only clearing user_api_key_cache
on message would leave that separate DualCache's in-memory copy untouched.
"""
cache = UserApiKeyCache()
spend_counter_in_memory_cache = InMemoryCache()
spend_counter_in_memory_cache.set_cache("spend:team_member:u-1:t-1", 999.0)
assert spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is not None
pubsub = _QueuePubSub(initial_messages=[_invalidation_message("spend:team_member:u-1:t-1")])
subscriber = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])),
user_api_key_cache=cache,
additional_in_memory_caches=(spend_counter_in_memory_cache,),
)
subscriber.start()
try:
for _ in range(200):
if spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is None:
break
await asyncio.sleep(0.01)
finally:
await subscriber.stop()
assert spend_counter_in_memory_cache.get_cache("spend:team_member:u-1:t-1") is None
@pytest.mark.asyncio
async def test_subscriber_ignores_malformed_messages() -> None:
cache = UserApiKeyCache()

View file

@ -9,11 +9,13 @@ from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from pydantic import ValidationError
from litellm._uuid import uuid
from litellm.proxy._types import UserAPIKeyAuth # Import UserAPIKeyAuth
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_BudgetTableFull,
LiteLLM_ModelTable,
LiteLLM_OrganizationMembershipTable,
@ -27,7 +29,9 @@ from litellm.proxy._types import (
Member,
ProxyErrorTypes,
ProxyException,
ResetSpendRequest,
TeamMemberAddRequest,
TeamMemberUpdateRequest,
UpdateTeamRequest,
)
from litellm.proxy.management_endpoints.team_endpoints import (
@ -42,12 +46,15 @@ from litellm.proxy.management_endpoints.team_endpoints import (
_transform_teams_to_deleted_records,
_update_model_table,
_validate_and_populate_member_user_info,
_validate_team_member_reset_spend_value,
_verify_team_access,
delete_team,
list_available_teams,
reset_team_member_spend_fn,
router,
team_member_add_duplication_check,
team_member_delete,
team_member_update,
update_team,
validate_team_org_change,
)
@ -12603,3 +12610,376 @@ async def test_invalidate_access_group_cache_deletes_the_cached_object():
"user_api_key_cache": cache,
"proxy_logging_obj": logging_obj,
}
def test_validate_team_member_reset_spend_value_rejects_non_numeric():
with pytest.raises(HTTPException) as exc:
_validate_team_member_reset_spend_value(
reset_to="not-a-number",
membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0),
)
assert exc.value.status_code == 400
def test_validate_team_member_reset_spend_value_rejects_negative():
with pytest.raises(HTTPException) as exc:
_validate_team_member_reset_spend_value(
reset_to=-1.0,
membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0),
)
assert exc.value.status_code == 400
@pytest.mark.parametrize("reset_to", [float("nan"), float("inf"), float("-inf")])
def test_validate_team_member_reset_spend_value_rejects_non_finite(reset_to):
"""NaN and +/-inf are instances of float and compare False against every bound
below (`nan < 0`, `nan > current_spend` are both False), so an isinstance-and-range
check alone lets them through to persist as the member's spend and silently
disable every later budget comparison against it."""
with pytest.raises(HTTPException) as exc:
_validate_team_member_reset_spend_value(
reset_to=reset_to,
membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0),
)
assert exc.value.status_code == 400
@pytest.mark.parametrize("reset_to", [True, False])
def test_reset_spend_request_rejects_bool_reset_to(reset_to):
"""bool is a subclass of int, so pydantic silently coerces True/False into 1.0/0.0 for a
``float`` field: {"reset_to": true} would otherwise reach _validate_team_member_reset_spend_value
as an indistinguishable 1.0 and reset the member's spend instead of failing the request."""
with pytest.raises(ValidationError):
ResetSpendRequest(reset_to=reset_to)
def test_validate_team_member_reset_spend_value_rejects_above_current_spend():
with pytest.raises(HTTPException) as exc:
_validate_team_member_reset_spend_value(
reset_to=20.0,
membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0),
)
assert exc.value.status_code == 400
def test_validate_team_member_reset_spend_value_rejects_above_max_budget():
with pytest.raises(HTTPException) as exc:
_validate_team_member_reset_spend_value(
reset_to=10.0,
membership=LiteLLM_TeamMembership(
user_id="u1",
team_id="t1",
spend=10.0,
litellm_budget_table=LiteLLM_BudgetTable(budget_id="b1", max_budget=5.0),
),
)
assert exc.value.status_code == 400
def test_validate_team_member_reset_spend_value_accepts_valid_reset():
result = _validate_team_member_reset_spend_value(
reset_to=0.0,
membership=LiteLLM_TeamMembership(user_id="u1", team_id="t1", spend=10.0),
)
assert result == 0.0
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_success(monkeypatch):
"""A proxy admin resetting a stuck team member's spend must write the DB
row to reset_to AND invalidate the cached spend/membership state, or the
429 the endpoint exists to clear keeps firing off the stale cache.
Asserted against real cache reads, not mock call args, so a change that
keeps the call but drops its effect still fails."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
mock_prisma_client = MagicMock()
mock_proxy_logging_obj = MagicMock()
real_cache = UserApiKeyCache()
await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership")
await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership")
real_spend_counter_cache = DualCache()
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=999.0)
membership_row = LiteLLM_TeamMembership(
user_id="member-1",
team_id="team-1",
spend=10.0,
litellm_budget_table=LiteLLM_BudgetTable(budget_id="b1", max_budget=50.0),
)
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
mock_prisma_client.db.litellm_teammembership.update = AsyncMock(return_value=membership_row)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj)
monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache)
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")),
):
response = await reset_team_member_spend_fn(
team_id="team-1",
user_id="member-1",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert response["spend"] == 0.0
assert response["previous_spend"] == 10.0
assert response["max_budget"] == 50.0
mock_prisma_client.db.litellm_teammembership.update.assert_awaited_once_with(
where={"user_id_team_id": {"user_id": "member-1", "team_id": "team-1"}},
data={"spend": 0.0},
)
assert await real_cache.async_get_cache(key="team-1_member-1") is None
assert await real_cache.async_get_cache(key="team_membership:member-1:team-1") is None
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 0.0
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_membership_not_found(monkeypatch):
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")),
):
with pytest.raises(HTTPException) as exc:
await reset_team_member_spend_fn(
team_id="team-1",
user_id="ghost-user",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert exc.value.status_code == 404
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_team_not_found(monkeypatch):
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "Team doesn't exist in db."})),
):
with pytest.raises(HTTPException) as exc:
await reset_team_member_spend_fn(
team_id="ghost-team",
user_id="member-1",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert exc.value.status_code == 404
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_forbidden_for_non_admin(monkeypatch):
"""A caller who is neither proxy admin, org admin, nor this team's admin must be refused,
matching every other team-mutating endpoint's authorization."""
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1", members_with_roles=[])),
):
with pytest.raises(HTTPException) as exc:
await reset_team_member_spend_fn(
team_id="team-1",
user_id="member-1",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="plain-user"
),
)
assert exc.value.status_code == 403
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_team_admin_cannot_reset_own_spend(monkeypatch):
"""_verify_team_access authorizes a team admin over their own team with no check that the
target differs from the caller. Unchecked, that admin could target their own membership row
and repeatedly zero it right before it crosses their per-member cap, consuming the shared
team budget without the configured limit ever binding (Veria finding on PR #37971)."""
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
team_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-admin", user_id="team-admin-1")
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(
return_value=LiteLLM_TeamTable(
team_id="team-1",
members_with_roles=[Member(user_id="team-admin-1", role="admin")],
)
),
):
with pytest.raises(HTTPException) as exc:
await reset_team_member_spend_fn(
team_id="team-1",
user_id="team-admin-1",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=team_admin,
)
assert exc.value.status_code == 403
mock_prisma_client.db.litellm_teammembership.update.assert_not_called()
@pytest.mark.asyncio
async def test_reset_team_member_spend_fn_proxy_admin_can_reset_own_spend(monkeypatch):
"""The self-reset guard is scoped to non-proxy-admin roles: a proxy admin resetting their
own membership spend is the platform-wide trust boundary, not a team-scoped one."""
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
membership_row = LiteLLM_TeamMembership(user_id="admin-user", team_id="team-1", spend=10.0)
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
mock_prisma_client.db.litellm_teammembership.update = AsyncMock(return_value=membership_row)
with patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")),
):
response = await reset_team_member_spend_fn(
team_id="team-1",
user_id="admin-user",
data=ResetSpendRequest(reset_to=0.0),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert response["spend"] == 0.0
@pytest.mark.asyncio
async def test_team_member_update_invalidates_team_member_spend_state_when_budget_patch_applied(monkeypatch):
"""Raising a stuck member's max_budget_in_team via the documented /team/member_update
endpoint must invalidate the cached membership state, or the raised cap never reaches the
admission check and the member stays 429ing. The live spend counter itself must be left
untouched: only the cap changed, and deleting the counter would force a reseed from the
DB's own spend column, which lags the live counter via periodic batch writes, briefly
UNDER-enforcing the raised cap against a spend value lower than what was actually tracked.
Asserted against real cache reads, not mock call args, so a change that keeps the call but
drops its effect still fails."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
mock_prisma_client = MagicMock()
real_cache = UserApiKeyCache()
await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership")
await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership")
real_spend_counter_cache = DualCache()
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=999.0)
team_row = LiteLLM_TeamTable(team_id="team-1", metadata={}, members_with_roles=[])
team_info_response = {
"team_info": team_row,
"team_memberships": [LiteLLM_TeamMembership(user_id="member-1", team_id="team-1", budget_id=None)],
}
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache)
mock_tx = AsyncMock()
mock_prisma_client.tx.return_value.__aenter__ = AsyncMock(return_value=mock_tx)
mock_prisma_client.tx.return_value.__aexit__ = AsyncMock(return_value=None)
with (
patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.team_info",
AsyncMock(return_value=team_info_response),
),
patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership",
AsyncMock(),
),
):
await team_member_update(
data=TeamMemberUpdateRequest(team_id="team-1", user_id="member-1", max_budget_in_team=999999.0),
http_request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert await real_cache.async_get_cache(key="team-1_member-1") is None
assert await real_cache.async_get_cache(key="team_membership:member-1:team-1") is None
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 999.0
@pytest.mark.asyncio
async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent(monkeypatch):
"""A role-only update carries an empty budget_patch and touches no budget state,
so the member's cached spend/membership state must be left untouched."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
mock_prisma_client = MagicMock()
real_cache = UserApiKeyCache()
await real_cache.async_set_cache(key="team-1_member-1", value="still-fresh-membership")
real_spend_counter_cache = DualCache()
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:member-1:team-1", value=1.5)
team_row = LiteLLM_TeamTable(team_id="team-1", metadata={}, members_with_roles=[])
team_info_response = {
"team_info": team_row,
"team_memberships": [LiteLLM_TeamMembership(user_id="member-1", team_id="team-1", budget_id=None)],
}
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", real_cache)
monkeypatch.setattr("litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache)
mock_tx = AsyncMock()
mock_prisma_client.tx.return_value.__aenter__ = AsyncMock(return_value=mock_tx)
mock_prisma_client.tx.return_value.__aexit__ = AsyncMock(return_value=None)
with (
patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints.team_info",
AsyncMock(return_value=team_info_response),
),
patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests
"litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership",
AsyncMock(),
),
):
await team_member_update(
data=TeamMemberUpdateRequest(team_id="team-1", user_id="member-1"),
http_request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
)
assert await real_cache.async_get_cache(key="team-1_member-1") == "still-fresh-membership"
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 1.5

View file

@ -11347,3 +11347,38 @@ class TestRouterModelNameOnStreamingChunks:
assert len(frames) >= 3
assert '"router_model_name":"deep-model"' in frames[0]
assert all('"router_model_name":"backup-tier"' in frame for frame in frames[1:])
@pytest.mark.asyncio
async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the_db_read():
"""A team-member spend reset writes the post-reset floor to the spend_db_floor marker
(auth_checks.invalidate_team_member_spend_state). A floor read already in flight when the
reset commits would otherwise cache its stale pre-reset DB value over the fresh marker,
letting a budget check raise the counter right back above the just-reset spend
(regression: PR #37971 Greptile finding)."""
from litellm.proxy.proxy_server import _authoritative_floor_spend
real_spend_counter_cache = DualCache()
counter_key = "spend:team_member:user-1:team-1"
marker_key = f"spend_db_floor:{counter_key}"
async def db_read_racing_with_a_reset(prisma_client, counter_key):
real_spend_counter_cache.in_memory_cache.set_cache(key=marker_key, value=0.0)
return 999.0
with (
patch.object( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
proxy_server_module, "spend_counter_cache", real_spend_counter_cache
),
patch.object( # test-quality-ok: the DB read must race the reset; no injectable seam for module-global prisma reads
proxy_server_module.SpendCounterReseed,
"from_db",
AsyncMock(side_effect=db_read_racing_with_a_reset),
),
):
result = await _authoritative_floor_spend(counter_key=counter_key)
assert result == 0.0
assert real_spend_counter_cache.in_memory_cache.get_cache(key=marker_key) == 0.0, (
"the in-flight DB read clobbered the post-reset floor marker with the stale pre-reset value"
)

View file

@ -14936,6 +14936,33 @@ export interface paths {
patch?: never;
trace?: never;
};
"/team/{team_id}/member/{user_id}/reset_spend": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
/**
* Reset Team Member Spend Fn
* @description Reset a team member's tracked spend against their per-member budget.
*
* A member's spend is tracked separately from both their own personal
* budget and the team's own budget (LiteLLM_TeamMembership.spend), so
* neither /user/update nor /team/update can clear it: this is the only
* endpoint that does. The cross-pod spend counter and cached membership
* reads are invalidated so the reset takes effect on the member's next
* request rather than waiting on the membership cache's TTL.
*/
post: operations["reset_team_member_spend_fn_team__team_id__member__user_id__reset_spend_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/team/{team_id}/members/me": {
parameters: {
query?: never;
@ -55083,6 +55110,42 @@ export interface operations {
};
};
};
reset_team_member_spend_fn_team__team_id__member__user_id__reset_spend_post: {
parameters: {
query?: never;
header?: never;
path: {
team_id: string;
user_id: string;
};
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["ResetSpendRequest"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
team_member_me_team__team_id__members_me_get: {
parameters: {
query?: never;