mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(auth): drop the membership write-epoch and background Redis replicate, load membership lazily
Move the cache-miss marker out of litellm.constants into auth_checks (CodeQL cyclic import) and stop logging user_id/team_id in the lookup failure (CodeQL log injection). Write the membership row through DualCache synchronously again instead of a background Redis task guarded by a bounded write-epoch map: the epoch was sampled after the Prisma read, so an invalidate that raced the read could be cached as current, and eviction of the epoch entry could let an old Redis write land. The synchronous write keeps invalidate_team_member_spend_state authoritative. Lookup failures return None again (fail-open like main) instead of 503, and the load is skipped on routes that neither resolve a model nor run budget checks. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f4e21430c0
commit
b00bb35563
3 changed files with 106 additions and 306 deletions
|
|
@ -2039,10 +2039,3 @@ BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit"
|
|||
# Shared read-only empty mapping, for defaulting optional Mapping parameters without
|
||||
# constructing a fresh mutable dict at each call site.
|
||||
EMPTY_MAPPING: Final = MappingProxyType({})
|
||||
|
||||
|
||||
class TeamMembershipCacheMiss:
|
||||
__slots__ = ()
|
||||
|
||||
|
||||
TEAM_MEMBERSHIP_CACHE_MISS: Final = TeamMembershipCacheMiss()
|
||||
|
|
|
|||
|
|
@ -35,8 +35,6 @@ from litellm.constants import (
|
|||
MODEL_ACCESS_GROUP_REGISTRY_MAX_SIZE,
|
||||
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
|
||||
TAG_REGISTRY_MAX_SIZE,
|
||||
TEAM_MEMBERSHIP_CACHE_MISS,
|
||||
TeamMembershipCacheMiss,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
|
@ -331,27 +329,19 @@ 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)
|
||||
_team_membership_write_epoch: 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_write_epoch(key: str) -> int:
|
||||
cached: Final[object] = _team_membership_write_epoch.get(key, 0)
|
||||
return cached if isinstance(cached, int) else 0
|
||||
|
||||
|
||||
def _bump_membership_write_epoch(key: str) -> None:
|
||||
_team_membership_write_epoch[key] = _membership_write_epoch(key) + 1
|
||||
|
||||
|
||||
def _membership_from_shared_load(result: object) -> LiteLLM_TeamMembership | None:
|
||||
if result is None or isinstance(result, LiteLLM_TeamMembership):
|
||||
return result
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Failed to load team membership",
|
||||
)
|
||||
return result if isinstance(result, LiteLLM_TeamMembership) else None
|
||||
|
||||
|
||||
def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
|
||||
|
|
@ -897,21 +887,6 @@ async def common_checks(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
# One membership read for model-access, access-group attribution, and
|
||||
# member-budget. Each used to call get_team_membership independently;
|
||||
# DualCache Redis SET/GET on every miss made those look like two Postgres spans.
|
||||
loaded_team_membership: LiteLLM_TeamMembership | None = None
|
||||
team_membership_loaded = False
|
||||
if team_object is not None and valid_token is not None and valid_token.user_id is not None:
|
||||
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,
|
||||
)
|
||||
team_membership_loaded = True
|
||||
|
||||
_model: Final[str | list[str] | None] = get_model_from_request(
|
||||
request_data=request_body,
|
||||
route=route,
|
||||
|
|
@ -926,6 +901,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)
|
||||
|
|
@ -2188,90 +2179,13 @@ async def get_tag_object(
|
|||
|
||||
def _membership_from_cached_payload(
|
||||
cached: object,
|
||||
) -> LiteLLM_TeamMembership | None | TeamMembershipCacheMiss:
|
||||
) -> LiteLLM_TeamMembership | None | _TeamMembershipCacheMiss:
|
||||
if cached is None:
|
||||
return TEAM_MEMBERSHIP_CACHE_MISS
|
||||
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
|
||||
|
||||
|
||||
async def _set_team_membership_l1(
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
key: str,
|
||||
value: object,
|
||||
*,
|
||||
model_type: type[LiteLLM_TeamMembership] | None,
|
||||
ttl: float | None,
|
||||
) -> None:
|
||||
match (model_type is not None, ttl is not None):
|
||||
case (False, False):
|
||||
await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True)
|
||||
case (False, True):
|
||||
await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, ttl=ttl)
|
||||
case (True, False):
|
||||
await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, model_type=model_type)
|
||||
case (True, True):
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=key, value=value, local_only=True, model_type=model_type, ttl=ttl
|
||||
)
|
||||
|
||||
|
||||
async def _replicate_team_membership_to_redis(
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
key: str,
|
||||
value: object,
|
||||
*,
|
||||
model_type: type[LiteLLM_TeamMembership] | None,
|
||||
ttl: float | None,
|
||||
write_epoch: int,
|
||||
) -> None:
|
||||
redis_cache: Final = user_api_key_cache.redis_cache
|
||||
if redis_cache is None or _membership_write_epoch(key) != write_epoch:
|
||||
return
|
||||
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
|
||||
try:
|
||||
if ttl is None:
|
||||
await redis_cache.async_set_cache(key, payload)
|
||||
else:
|
||||
await redis_cache.async_set_cache(key, payload, ttl=ttl)
|
||||
if _membership_write_epoch(key) != write_epoch:
|
||||
await redis_cache.async_delete_cache(key)
|
||||
except Exception:
|
||||
return
|
||||
|
||||
|
||||
async def _populate_team_membership_cache(
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
key: str,
|
||||
value: object,
|
||||
*,
|
||||
model_type: type[LiteLLM_TeamMembership] | None = None,
|
||||
ttl: float | None = None,
|
||||
) -> None:
|
||||
write_epoch: Final = _membership_write_epoch(key)
|
||||
await _set_team_membership_l1(
|
||||
user_api_key_cache,
|
||||
key,
|
||||
value,
|
||||
model_type=model_type,
|
||||
ttl=ttl,
|
||||
)
|
||||
if _membership_write_epoch(key) != write_epoch:
|
||||
user_api_key_cache.in_memory_cache_for(key).delete_cache(key)
|
||||
return
|
||||
|
||||
asyncio.create_task(
|
||||
_replicate_team_membership_to_redis(
|
||||
user_api_key_cache,
|
||||
key,
|
||||
value,
|
||||
model_type=model_type,
|
||||
ttl=ttl,
|
||||
write_epoch=write_epoch,
|
||||
)
|
||||
)
|
||||
return cached_membership if cached_membership is not None else _TEAM_MEMBERSHIP_CACHE_MISS
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
|
|
@ -2283,7 +2197,7 @@ async def _fetch_team_membership_from_db(
|
|||
parent_otel_span: Span | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> LiteLLM_TeamMembership | None:
|
||||
"""Prisma read + L1 populate. Decorated so cache hits on ``get_team_membership`` are not postgres spans."""
|
||||
"""Prisma read + cache populate. Decorated so cache hits on ``get_team_membership`` are not postgres spans."""
|
||||
_ = 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}},
|
||||
|
|
@ -2291,19 +2205,17 @@ async def _fetch_team_membership_from_db(
|
|||
)
|
||||
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
|
||||
if response is None:
|
||||
await _populate_team_membership_cache(
|
||||
user_api_key_cache,
|
||||
_key,
|
||||
NO_TEAM_MEMBERSHIP_SENTINEL,
|
||||
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
|
||||
|
||||
membership: Final = LiteLLM_TeamMembership.model_validate(response.dict())
|
||||
await _populate_team_membership_cache(
|
||||
user_api_key_cache,
|
||||
_key,
|
||||
membership,
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_key,
|
||||
value=membership,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
return membership
|
||||
|
|
@ -2321,7 +2233,7 @@ async def _load_team_membership_on_cache_miss(
|
|||
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):
|
||||
if not isinstance(redis_membership, _TeamMembershipCacheMiss):
|
||||
return redis_membership
|
||||
|
||||
return await _fetch_team_membership_from_db(
|
||||
|
|
@ -2332,18 +2244,9 @@ async def _load_team_membership_on_cache_miss(
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error getting team membership for user_id: %s, team_id: %s",
|
||||
user_id,
|
||||
team_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Failed to load team membership",
|
||||
) from e
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Error getting team membership")
|
||||
return None
|
||||
|
||||
|
||||
async def get_team_membership(
|
||||
|
|
@ -2366,7 +2269,7 @@ async def get_team_membership(
|
|||
|
||||
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):
|
||||
if not isinstance(l1_membership, _TeamMembershipCacheMiss):
|
||||
return l1_membership
|
||||
|
||||
inflight: Final[object] = _team_membership_inflight.get(_key)
|
||||
|
|
@ -2375,9 +2278,6 @@ async def get_team_membership(
|
|||
|
||||
if prisma_client is None:
|
||||
raise Exception("No db connected")
|
||||
prisma: Final[object] = prisma_client
|
||||
if isinstance(prisma, str):
|
||||
return None
|
||||
|
||||
task: Final = asyncio.ensure_future(
|
||||
_load_team_membership_on_cache_miss(
|
||||
|
|
@ -2908,7 +2808,6 @@ async def invalidate_team_member_spend_state(
|
|||
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
|
||||
)
|
||||
|
||||
_bump_membership_write_epoch(team_membership_reservation_cache_key(user_id=user_id, team_id=team_id))
|
||||
await evict_and_broadcast(
|
||||
cache_keys=(
|
||||
team_membership_auth_cache_key(team_id=team_id, user_id=user_id),
|
||||
|
|
|
|||
|
|
@ -6494,52 +6494,6 @@ async def test_get_team_membership_coalesces_parallel_db_fetches():
|
|||
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_membership_returns_before_redis_set_completes():
|
||||
"""Auth must not wait on DualCache Redis SET; L1 is enough for the next lookup."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_membership
|
||||
|
||||
membership_row = MagicMock()
|
||||
membership_row.dict = lambda: {"user_id": "u-redis", "team_id": "t-redis", "spend": 2.0}
|
||||
|
||||
hang_redis_set = asyncio.Event()
|
||||
|
||||
async def _hanging_redis_set(*args, **kwargs):
|
||||
await hang_redis_set.wait()
|
||||
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
redis_cache.async_set_cache = AsyncMock(side_effect=_hanging_redis_set)
|
||||
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
|
||||
|
||||
first = await asyncio.wait_for(
|
||||
get_team_membership(
|
||||
user_id="u-redis",
|
||||
team_id="t-redis",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
),
|
||||
timeout=0.5,
|
||||
)
|
||||
second = await get_team_membership(
|
||||
user_id="u-redis",
|
||||
team_id="t-redis",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
|
||||
hang_redis_set.set()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert first is not None and second is not None
|
||||
assert first.spend == 2.0
|
||||
assert second.spend == 2.0
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_calls_get_team_membership_once_per_request():
|
||||
"""Model-access, attribution, and member-budget must reuse one membership load."""
|
||||
|
|
@ -6588,33 +6542,85 @@ async def test_common_checks_calls_get_team_membership_once_per_request():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_membership_db_error_raises_503_not_none():
|
||||
"""A Prisma failure must fail closed as 503, not look like a missing membership row."""
|
||||
from fastapi import HTTPException
|
||||
async def test_common_checks_skips_membership_load_when_no_check_reads_it():
|
||||
"""A management route has no model and no budget gate, so the membership row is never loaded."""
|
||||
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():
|
||||
"""A Prisma failure reads as no membership, caches nothing, and the next call hits the DB again."""
|
||||
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"))
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(
|
||||
side_effect=[RuntimeError("db down"), membership_row]
|
||||
)
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await get_team_membership(
|
||||
user_id="u-fail",
|
||||
team_id="t-fail",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
cached = await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-fail", team_id="t-fail"))
|
||||
assert exc.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE
|
||||
assert cached is None
|
||||
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():
|
||||
"""Unit tests stub prisma_client as a string; that is not a lookup failure and must not 503."""
|
||||
"""Unit tests stub prisma_client as a string; the lookup fails and reads as no membership."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_membership
|
||||
|
||||
result = await get_team_membership(
|
||||
|
|
@ -6626,59 +6632,6 @@ async def test_get_team_membership_string_prisma_client_returns_none():
|
|||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_does_not_skip_member_limits_when_membership_lookup_fails():
|
||||
"""common_checks must not mark membership loaded-absent after a lookup error."""
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
|
||||
team = LiteLLM_TeamTable(team_id="t-fail-closed")
|
||||
token = UserAPIKeyAuth(
|
||||
token="k-fail-closed",
|
||||
user_id="u-fail-closed",
|
||||
team_id="t-fail-closed",
|
||||
models=["gpt-4o-mini"],
|
||||
)
|
||||
lookup_error = HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Failed to load team membership",
|
||||
)
|
||||
|
||||
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: injects membership lookup failure; common_checks has no seam
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=lookup_error,
|
||||
),
|
||||
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
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await common_checks(
|
||||
request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
|
||||
team_object=team,
|
||||
user_object=LiteLLM_UserTable(user_id="u-fail-closed"),
|
||||
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 exc.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load():
|
||||
"""Cancelling one coalesced waiter must not cancel the shared Prisma load."""
|
||||
|
|
@ -6721,51 +6674,6 @@ async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load():
|
|||
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_membership_redis_replicate_does_not_restore_after_invalidate():
|
||||
"""A delayed Redis SET must not rewrite L1 or leave Redis holding membership after invalidation."""
|
||||
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
|
||||
|
||||
membership_row = MagicMock()
|
||||
membership_row.dict = lambda: {"user_id": "u-stale", "team_id": "t-stale", "spend": 9.0}
|
||||
|
||||
hang_redis_set = asyncio.Event()
|
||||
|
||||
async def _hanging_redis_set(*args, **kwargs):
|
||||
await hang_redis_set.wait()
|
||||
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
redis_cache.async_set_cache = AsyncMock(side_effect=_hanging_redis_set)
|
||||
redis_cache.async_delete_cache = AsyncMock()
|
||||
redis_cache.delete_cache = MagicMock()
|
||||
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
|
||||
|
||||
loaded = await get_team_membership(
|
||||
user_id="u-stale",
|
||||
team_id="t-stale",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
cache_key = team_membership_reservation_cache_key(user_id="u-stale", team_id="t-stale")
|
||||
await invalidate_team_member_spend_state(user_id="u-stale", team_id="t-stale", user_api_key_cache=cache)
|
||||
after_invalidate = await cache.async_get_cache(key=cache_key, local_only=True)
|
||||
|
||||
hang_redis_set.set()
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
after_replicate = await cache.async_get_cache(key=cache_key, local_only=True)
|
||||
|
||||
assert loaded is not None
|
||||
assert after_invalidate is None
|
||||
assert after_replicate is None
|
||||
redis_cache.async_delete_cache.assert_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue