Merge pull request #41102 from BerriAI/litellm_team_membership_once_main

fix(auth): load team membership once per request and skip prisma on an L1 hit
This commit is contained in:
Yassin Kortam 2026-09-14 14:46:06 -07:00 committed by GitHub
commit 24bfd5fba1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 528 additions and 90 deletions

View file

@ -328,9 +328,23 @@ _safe_json_loads_obj: Final = _typed_json_loads(safe_json_loads)
last_db_access_time: Final = LimitedSizeOrderedDict(max_size=100)
db_cache_expiry: Final = DEFAULT_IN_MEMORY_TTL # refresh every 5s
_TEAM_MEMBERSHIP_INFLIGHT_MAX: Final = 10000
_team_membership_inflight: Final = LimitedSizeOrderedDict(max_size=_TEAM_MEMBERSHIP_INFLIGHT_MAX)
class _TeamMembershipCacheMiss:
__slots__ = ()
_TEAM_MEMBERSHIP_CACHE_MISS: Final = _TeamMembershipCacheMiss()
all_routes: Final = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
def _membership_from_shared_load(result: object) -> LiteLLM_TeamMembership | None:
return result if isinstance(result, LiteLLM_TeamMembership) else None
def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
"""
Log a warning when budget lookup fails; cache will not be populated.
@ -888,6 +902,22 @@ async def common_checks(
and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route))
)
membership_user_id: Final = (
valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None
)
team_membership_loaded: Final = team_object is not None and membership_user_id is not None
loaded_team_membership: Final = (
await get_team_membership(
user_id=membership_user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if team_object is not None and membership_user_id is not None
else None
)
unpriced_models: Final = (
_unpriced_models_in_request(model=_model, llm_router=llm_router)
if litellm.block_requests_for_models_without_pricing and RouteChecks.is_llm_api_route(route=route)
@ -937,6 +967,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
)
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
@ -988,6 +1020,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
)
# Run before apply_key_tags_pre_auth injects key metadata.tags into request_body.
@ -1097,6 +1131,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
),
_check_end_user_budget(end_user_obj=end_user_object, route=route)
if end_user_object is not None and end_user_object.litellm_budget_table is not None
@ -2142,7 +2178,76 @@ async def get_tag_object(
return tag_objects.get(tag_name)
def _membership_from_cached_payload(
cached: object,
) -> LiteLLM_TeamMembership | None | _TeamMembershipCacheMiss:
if cached is None:
return _TEAM_MEMBERSHIP_CACHE_MISS
if cached == NO_TEAM_MEMBERSHIP_SENTINEL:
return None
cached_membership: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership)
return cached_membership if cached_membership is not None else _TEAM_MEMBERSHIP_CACHE_MISS
@log_db_metrics
async def _fetch_team_membership_from_db(
user_id: str,
team_id: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
) -> LiteLLM_TeamMembership | None:
_ = parent_otel_span, proxy_logging_obj
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
)
membership: Final = None if response is None else LiteLLM_TeamMembership.model_validate(response.dict())
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
if membership is None:
await user_api_key_cache.async_set_cache(
key=_key,
value=NO_TEAM_MEMBERSHIP_SENTINEL,
ttl=get_management_object_ttl(user_api_key_cache),
)
else:
await user_api_key_cache.async_set_cache(
key=_key,
value=membership,
model_type=LiteLLM_TeamMembership,
)
return membership
async def _load_team_membership_on_cache_miss(
user_id: str,
team_id: str,
cache_key: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
) -> LiteLLM_TeamMembership | None:
try:
redis_cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
redis_membership: Final = _membership_from_cached_payload(redis_cached)
if not isinstance(redis_membership, _TeamMembershipCacheMiss):
return redis_membership
return await _fetch_team_membership_from_db(
user_id=user_id,
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
verbose_proxy_logger.exception("Error getting team membership")
return None
async def get_team_membership(
user_id: str,
team_id: str,
@ -2156,54 +2261,42 @@ async def get_team_membership(
Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership).
"""
from litellm.proxy._types import LiteLLM_TeamMembership
if prisma_client is None:
raise Exception("No db connected")
if user_id is None or team_id is None:
return None
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
# check if in cache
cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key)
if cached == NO_TEAM_MEMBERSHIP_SENTINEL:
return None
cached_membership_obj: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership)
if cached_membership_obj is not None:
return cached_membership_obj
l1_cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key, local_only=True)
l1_membership: Final = _membership_from_cached_payload(l1_cached)
if not isinstance(l1_membership, _TeamMembershipCacheMiss):
return l1_membership
# else, check db
try:
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
inflight: Final[object] = _team_membership_inflight.get(_key)
if isinstance(inflight, asyncio.Task):
return _membership_from_shared_load(await asyncio.shield(inflight))
if prisma_client is None:
raise Exception("No db connected")
task: Final = asyncio.ensure_future(
_load_team_membership_on_cache_miss(
user_id=user_id,
team_id=team_id,
cache_key=_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
)
_team_membership_inflight[_key] = task
if response is None:
await user_api_key_cache.async_set_cache(
key=_key,
value=NO_TEAM_MEMBERSHIP_SENTINEL,
ttl=get_management_object_ttl(user_api_key_cache),
)
return None
def _clear_inflight(_done: object) -> None:
if _team_membership_inflight.get(_key) is task:
_team_membership_inflight.pop(_key, None)
_response: Final = LiteLLM_TeamMembership.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=_key,
value=_response,
model_type=LiteLLM_TeamMembership,
)
return _response
except Exception:
verbose_proxy_logger.exception(
"Error getting team membership for user_id: %s, team_id: %s",
user_id,
team_id,
)
return None
task.add_done_callback(_clear_inflight)
return _membership_from_shared_load(await asyncio.shield(task))
def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool:
@ -2662,6 +2755,12 @@ async def invalidate_team_member_spend_state(
publish_auth_cache_invalidation,
)
inflight: Final[object] = _team_membership_inflight.pop(
team_membership_reservation_cache_key(user_id=user_id, team_id=team_id), None
)
if isinstance(inflight, asyncio.Task) and inflight is not asyncio.current_task():
await asyncio.wait((inflight,))
if new_spend is not None:
from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache
@ -4116,18 +4215,21 @@ async def _team_member_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> Sequence[str]:
"""The member's own ``allowed_models`` scope; empty when the member is not narrowed below the team."""
if team_object is None or valid_token.user_id is None:
return ()
team_membership: Final = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if not team_membership_loaded:
team_membership = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
return () if team_membership is None else _member_allowed_models(team_membership)
@ -4163,6 +4265,8 @@ async def _granted_model_lists(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[Sequence[str], ...]:
"""One model allowlist per level that participates in authorizing the request."""
return (
@ -4174,6 +4278,8 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
),
project_object.models if project_object is not None else (),
await _org_granted_models(
@ -4268,6 +4374,8 @@ async def collect_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[str, ...]:
"""
The budgeted model access groups that authorized this request, sorted and deduplicated.
@ -4313,6 +4421,8 @@ async def collect_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
)
for granted_model in granted_models
)
@ -4328,6 +4438,8 @@ async def stamp_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[str, ...]:
"""Record the groups that authorized this request on its auth object, for the post-call spend
writer and the reservation counters, and hand them back for the budget check."""
@ -4344,6 +4456,8 @@ async def stamp_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
)
except Exception as e: # noqa: BLE001 # fail-safe: attribution is spend telemetry, it must never break auth
verbose_proxy_logger.debug("model access group attribution failed: %s", e)
@ -5146,6 +5260,8 @@ async def _check_team_member_budget(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
):
"""Check if team member is over their max budget within the team."""
if (
@ -5154,23 +5270,25 @@ async def _check_team_member_budget(
and valid_token is not None
and valid_token.user_id is not None
):
team_membership: Final = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if not team_membership_loaded:
team_membership = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
loaded_membership = team_membership
# Per-member override wins; otherwise fall back to the team-level
# default configured via team.metadata["team_member_budget_id"].
team_member_budget: float | None = None
if (
team_membership is not None
and team_membership.litellm_budget_table is not None
and team_membership.litellm_budget_table.max_budget is not None
loaded_membership is not None
and loaded_membership.litellm_budget_table is not None
and loaded_membership.litellm_budget_table.max_budget is not None
):
team_member_budget = team_membership.litellm_budget_table.max_budget
team_member_budget = loaded_membership.litellm_budget_table.max_budget
else:
default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id")
if isinstance(default_budget_id, str):
@ -5189,7 +5307,7 @@ async def _check_team_member_budget(
team_member_budget = default_budget.max_budget
if team_member_budget is not None:
team_member_spend = (team_membership.spend if team_membership is not None else 0.0) or 0.0
team_member_spend = (loaded_membership.spend if loaded_membership is not None else 0.0) or 0.0
# Read from cross-pod counter (Redis-first) if available
from litellm.proxy.proxy_server import get_current_spend
@ -5218,6 +5336,8 @@ async def _check_team_member_model_access(
prisma_client: Optional["PrismaClient"],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> None:
"""
Check if a team member's per-member model scope allows access to the requested model.
@ -5228,22 +5348,24 @@ async def _check_team_member_model_access(
if valid_token.user_id is None or team_object.team_id is None:
return
team_membership: Final = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if not team_membership_loaded:
team_membership = await get_team_membership(
user_id=valid_token.user_id,
team_id=team_object.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
loaded_membership = team_membership
if (
team_membership is None
or team_membership.litellm_budget_table is None
or not team_membership.litellm_budget_table.allowed_models
loaded_membership is None
or loaded_membership.litellm_budget_table is None
or not loaded_membership.litellm_budget_table.allowed_models
):
return # no per-member restriction — inherit team-level check
member_allowed_models: Final[list[str]] = team_membership.litellm_budget_table.allowed_models
member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models
try:
_can_object_call_model(
model=model,

View file

@ -66,6 +66,7 @@ async def test_join_binds_the_membership_to_the_requested_team(prisma):
cache = _frozen_cache()
refs = AuthObjectRefs(user_id=user_id, team_id=team_a, membership_user_id=user_id, organization_id=org_id)
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
assert cache.in_memory_cache.get_cache(f"org_id:{org_id}") is not None
dead_db = _dead_db()
membership = await get_team_membership(

View file

@ -5527,7 +5527,9 @@ async def _run_internal_user_budget_alert(
with (
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam
patch("litellm.proxy.proxy_server.get_current_spend", _get_spend), # test-quality-ok: common_checks imports it locally
patch( # test-quality-ok: common_checks imports get_current_spend locally
"litellm.proxy.proxy_server.get_current_spend", _get_spend
),
patch.object(slack_alerting, "send_alert", send_alert),
):
error: Final = await _check_for_error()
@ -6419,9 +6421,7 @@ async def test_get_team_membership_negative_caches_a_missing_row():
assert first is None
assert second is None
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
cached = await cache.async_get_cache(
key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1")
)
cached = await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1"))
assert cached == NO_TEAM_MEMBERSHIP_SENTINEL
@ -6454,6 +6454,314 @@ async def test_get_team_membership_reads_sentinel_as_no_membership_not_a_model()
mock_prisma_client.db.litellm_teammembership.find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_team_membership_coalesces_parallel_db_fetches():
from litellm.proxy.auth.auth_checks import get_team_membership
started = asyncio.Event()
release = asyncio.Event()
membership_row = MagicMock()
membership_row.dict = lambda: {"user_id": "u-parallel", "team_id": "t-parallel", "spend": 1.0}
async def _slow_find_unique(*args, **kwargs):
started.set()
await release.wait()
return membership_row
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_slow_find_unique)
cache = UserApiKeyCache()
async def _load():
return await get_team_membership(
user_id="u-parallel",
team_id="t-parallel",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
first = asyncio.create_task(_load())
second = asyncio.create_task(_load())
await started.wait()
await asyncio.sleep(0)
release.set()
results = await asyncio.gather(first, second)
assert results[0] is not None and results[1] is not None
assert results[0].user_id == "u-parallel"
assert results[1].user_id == "u-parallel"
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_team_membership_invalidation_waits_for_in_flight_load_then_evicts_it():
from litellm.proxy._types import LiteLLM_TeamMembership
from litellm.proxy.auth.auth_checks import get_team_membership, invalidate_team_member_spend_state
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key
started = asyncio.Event()
release_stale = asyncio.Event()
rows = iter(("budget-old", "budget-new"))
async def _find_unique(*args, **kwargs):
budget_id = next(rows)
row = MagicMock()
row.dict = lambda: {"user_id": "u-inv", "team_id": "t-inv", "spend": 1.0, "budget_id": budget_id}
if budget_id == "budget-old":
started.set()
await release_stale.wait()
return row
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_find_unique)
cache = UserApiKeyCache()
_key = team_membership_reservation_cache_key(user_id="u-inv", team_id="t-inv")
async def _load():
return await get_team_membership(
user_id="u-inv", team_id="t-inv", prisma_client=mock_prisma_client, user_api_key_cache=cache
)
stale = asyncio.create_task(_load())
await asyncio.wait_for(started.wait(), timeout=2)
invalidation = asyncio.create_task(
invalidate_team_member_spend_state(user_id="u-inv", team_id="t-inv", user_api_key_cache=cache)
)
for _ in range(5):
await asyncio.sleep(0)
assert not invalidation.done()
release_stale.set()
await asyncio.wait_for(invalidation, timeout=2)
stale_result = await stale
assert stale_result is not None and stale_result.budget_id == "budget-old"
assert await cache.async_get_cache(key=_key) is None
fresh_result = await _load()
assert fresh_result is not None and fresh_result.budget_id == "budget-new"
assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2
cached = CacheCodec.deserialize(await cache.async_get_cache(key=_key), model_type=LiteLLM_TeamMembership)
assert cached is not None and cached.budget_id == "budget-new"
again = await _load()
assert again is not None and again.budget_id == "budget-new"
assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2
@pytest.mark.asyncio
async def test_get_team_membership_invalidation_during_cache_write_evicts_stale_entry():
from litellm.proxy.auth.auth_checks import get_team_membership, invalidate_team_member_spend_state
from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key
write_started = asyncio.Event()
release_write = asyncio.Event()
class _SlowWriteCache(UserApiKeyCache):
async def async_set_cache(self, key, value, local_only=False, **kwargs):
write_started.set()
await release_write.wait()
return await super().async_set_cache(key, value, local_only=local_only, **kwargs)
row = MagicMock()
row.dict = lambda: {"user_id": "u-w", "team_id": "t-w", "spend": 1.0, "budget_id": "budget-old"}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=row)
cache = _SlowWriteCache()
stale = asyncio.create_task(
get_team_membership(user_id="u-w", team_id="t-w", prisma_client=mock_prisma_client, user_api_key_cache=cache)
)
await asyncio.wait_for(write_started.wait(), timeout=2)
invalidation = asyncio.create_task(
invalidate_team_member_spend_state(user_id="u-w", team_id="t-w", user_api_key_cache=cache)
)
for _ in range(5):
await asyncio.sleep(0)
assert not invalidation.done()
release_write.set()
await asyncio.wait_for(invalidation, timeout=2)
stale_result = await stale
assert stale_result is not None and stale_result.budget_id == "budget-old"
assert await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-w", team_id="t-w")) is None
@pytest.mark.asyncio
async def test_common_checks_calls_get_team_membership_once_per_request():
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
team = LiteLLM_TeamTable(team_id="t-once")
token = UserAPIKeyAuth(token="k-once", user_id="u-once", team_id="t-once", models=["gpt-4o-mini"])
membership = MagicMock()
membership.litellm_budget_table = None
membership.spend = 0.0
with (
patch( # test-quality-ok: common_checks imports prisma_client from proxy_server
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
patch( # test-quality-ok: common_checks imports user_api_key_cache from proxy_server
"litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()
),
patch( # test-quality-ok: counts membership loads; common_checks has no membership seam
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=membership,
) as load_membership,
patch( # test-quality-ok: common_checks imports get_current_spend locally
"litellm.proxy.proxy_server.get_current_spend", new_callable=AsyncMock, return_value=0.0
),
):
result = await common_checks(
request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
team_object=team,
user_object=LiteLLM_UserTable(user_id="u-once"),
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=token,
request=MagicMock(spec=Request),
)
assert result is True
assert load_membership.await_count == 1
@pytest.mark.asyncio
async def test_common_checks_skips_membership_load_when_no_check_reads_it():
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
team = LiteLLM_TeamTable(team_id="t-lazy")
token = UserAPIKeyAuth(token="k-lazy", user_id="u-lazy", team_id="t-lazy")
with (
patch( # test-quality-ok: common_checks imports prisma_client from proxy_server
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
patch( # test-quality-ok: common_checks imports user_api_key_cache from proxy_server
"litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()
),
patch( # test-quality-ok: counts membership loads; common_checks has no membership seam
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
) as load_membership,
):
result = await common_checks(
request_body={},
team_object=team,
user_object=LiteLLM_UserTable(user_id="u-lazy"),
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/key/info",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=token,
request=MagicMock(spec=Request),
)
assert result is True
load_membership.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_team_membership_db_error_returns_none_and_retries_next_call():
from litellm.proxy.auth.auth_checks import get_team_membership
from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key
membership_row = MagicMock()
membership_row.dict = lambda: {"user_id": "u-fail", "team_id": "t-fail", "spend": 1.0}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(
side_effect=[RuntimeError("db down"), membership_row]
)
cache = UserApiKeyCache()
failed = await get_team_membership(
user_id="u-fail",
team_id="t-fail",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
cached_after_failure = await cache.async_get_cache(
key=team_membership_reservation_cache_key(user_id="u-fail", team_id="t-fail")
)
recovered = await get_team_membership(
user_id="u-fail",
team_id="t-fail",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
assert failed is None
assert cached_after_failure is None
assert recovered is not None
assert recovered.user_id == "u-fail"
assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2
@pytest.mark.asyncio
async def test_get_team_membership_string_prisma_client_returns_none():
from litellm.proxy.auth.auth_checks import get_team_membership
result = await get_team_membership(
user_id="u-str",
team_id="t-str",
prisma_client="hello-world",
user_api_key_cache=UserApiKeyCache(),
)
assert result is None
@pytest.mark.asyncio
async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load():
from litellm.proxy.auth.auth_checks import get_team_membership
started = asyncio.Event()
release = asyncio.Event()
membership_row = MagicMock()
membership_row.dict = lambda: {"user_id": "u-shield", "team_id": "t-shield", "spend": 1.0}
async def _slow_find_unique(*args, **kwargs):
started.set()
await release.wait()
return membership_row
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_slow_find_unique)
cache = UserApiKeyCache()
async def _load():
return await get_team_membership(
user_id="u-shield",
team_id="t-shield",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
owner = asyncio.create_task(_load())
await started.wait()
waiter = asyncio.create_task(_load())
await asyncio.sleep(0)
waiter.cancel()
with pytest.raises(asyncio.CancelledError):
await waiter
release.set()
result = await owner
assert result is not None
assert result.user_id == "u-shield"
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel():
"""
@ -6477,10 +6785,7 @@ async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sent
assert before is None
await invalidate_team_member_spend_state(user_id="u-1", team_id="t-1", user_api_key_cache=cache)
assert (
await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1"))
is None
)
assert await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1")) is None
after = await get_team_membership(
user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache

View file

@ -31,6 +31,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.common_utils.reset_budget_job import _model_access_group_counter_key
from litellm.proxy.common_utils.user_api_key_cache import (
NO_TEAM_MEMBERSHIP_SENTINEL,
UserApiKeyCache,
model_access_group_registry_cache_key,
model_access_group_spend_counter_key,
@ -95,6 +96,11 @@ async def _cache(
),
model_type=LiteLLM_TeamMembership,
)
else:
await cache.async_set_cache(
key=team_membership_reservation_cache_key(user_id=USER_ID, team_id=TEAM_ID),
value=NO_TEAM_MEMBERSHIP_SENTINEL,
)
if org_models:
await cache.async_set_cache(
key=f"org_id:{ORG_ID}",
@ -314,9 +320,7 @@ class _RecordingPrismaClient:
def __init__(self, *rows: _MagBudgetRow) -> None:
self.rows = {row.access_group_name: row for row in rows}
self.batches: list[list[str]] = []
self.db = SimpleNamespace(
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._find_many)
)
self.db = SimpleNamespace(litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._find_many))
async def _find_many(self, **kwargs):
requested = list(kwargs["where"]["access_group_name"]["in"])
@ -345,7 +349,9 @@ async def _enforce(
read, seen = _spend_reader(spend_by_counter_key or {})
# The check takes its client and cache as arguments, injected just below. get_current_spend is the
# one collaborator it reaches by a lazy `from litellm.proxy.proxy_server import`, with no parameter.
with patch("litellm.proxy.proxy_server.get_current_spend", read): # test-quality-ok: get_current_spend is lazily imported inside _model_access_group_max_budget_check and has no injection point
with patch( # test-quality-ok: get_current_spend is lazily imported inside the budget check
"litellm.proxy.proxy_server.get_current_spend", read
):
await _model_access_group_max_budget_check(
matched_model_access_groups=matched,
prisma_client=prisma_client if prisma_client is not None else _RecordingPrismaClient(*rows),
@ -491,9 +497,7 @@ async def test_a_second_request_serves_the_budget_row_from_cache():
async def test_a_database_error_does_not_block_the_request():
class _FailingPrismaClient:
def __init__(self) -> None:
self.db = SimpleNamespace(
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._boom)
)
self.db = SimpleNamespace(litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._boom))
async def _boom(self, **kwargs):
raise RuntimeError("database unavailable")
@ -507,11 +511,15 @@ async def _common_checks_with_over_budget_group(*, skip_budget_checks: bool) ->
read, _ = _spend_reader({MODEL_ACCESS_GROUP_COUNTER_KEY: 99.0})
with (
# common_checks resolves all three off the proxy_server module at call time; its signature
# has no client, cache or spend-reader parameter to pass them through instead.
patch("litellm.proxy.proxy_server.prisma_client", prisma_client), # test-quality-ok: common_checks lazily imports prisma_client from proxy_server and takes no client parameter
patch("litellm.proxy.proxy_server.user_api_key_cache", cache), # test-quality-ok: common_checks lazily imports user_api_key_cache from proxy_server and takes no cache parameter
patch("litellm.proxy.proxy_server.get_current_spend", read), # test-quality-ok: get_current_spend is lazily imported inside the budget check and has no injection point
patch( # test-quality-ok: common_checks lazily imports prisma_client from proxy_server
"litellm.proxy.proxy_server.prisma_client", prisma_client
),
patch( # test-quality-ok: common_checks lazily imports user_api_key_cache from proxy_server
"litellm.proxy.proxy_server.user_api_key_cache", cache
),
patch( # test-quality-ok: get_current_spend is lazily imported inside the budget check
"litellm.proxy.proxy_server.get_current_spend", read
),
):
return await common_checks(
request_body={"model": "gpt-4o", "messages": []},
@ -524,7 +532,9 @@ async def _common_checks_with_over_budget_group(*, skip_budget_checks: bool) ->
llm_router=Router(model_list=MODEL_LIST),
proxy_logging_obj=ProxyLogging(user_api_key_cache=cache),
valid_token=UserAPIKeyAuth(api_key="hashed", models=["tier-a"], user_id=USER_ID),
request=SimpleNamespace(method="POST", headers={}, query_params={}, url=SimpleNamespace(path="/v1/chat/completions")),
request=SimpleNamespace(
method="POST", headers={}, query_params={}, url=SimpleNamespace(path="/v1/chat/completions")
),
skip_budget_checks=skip_budget_checks,
)