mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
fix(auth): load team membership once per request and skip prisma on an L1 hit
common_checks was querying get_team_membership twice, and DualCache awaited Redis SET on the auth path, so LRU eviction plus a hung Redis write showed up as two postgres spans
This commit is contained in:
parent
6e9e08475a
commit
5457f48290
2 changed files with 369 additions and 71 deletions
|
|
@ -327,6 +327,9 @@ _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_CACHE_MISS: Final = object()
|
||||
_team_membership_inflight: dict[str, asyncio.Task[LiteLLM_TeamMembership | None]] = {}
|
||||
|
||||
all_routes: Final = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
|
||||
|
||||
|
||||
|
|
@ -873,6 +876,21 @@ 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,
|
||||
|
|
@ -936,6 +954,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
|
||||
|
|
@ -987,6 +1007,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.
|
||||
|
|
@ -1096,6 +1118,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
|
||||
|
|
@ -2141,7 +2165,142 @@ async def get_tag_object(
|
|||
return tag_objects.get(tag_name)
|
||||
|
||||
|
||||
def _debug_team_membership_log(hypothesis_id: str, message: str, data: dict[str, object]) -> None:
|
||||
# #region agent log
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
with open("/Users/shijain/genai-apps/genai-proxy/.cursor/debug-86534f.log", "a", encoding="utf-8") as _f:
|
||||
_f.write(
|
||||
_json.dumps(
|
||||
{
|
||||
"sessionId": "86534f",
|
||||
"timestamp": int(time.time() * 1000),
|
||||
"location": "auth_checks.py:get_team_membership",
|
||||
"message": message,
|
||||
"hypothesisId": hypothesis_id,
|
||||
"data": data,
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
# #endregion
|
||||
|
||||
|
||||
def _membership_from_cached_payload(cached: object) -> LiteLLM_TeamMembership | None | object:
|
||||
"""Decode a DualCache payload. ``_TEAM_MEMBERSHIP_CACHE_MISS`` means try the next tier."""
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
"""Await in-memory write; replicate to Redis off the auth await path."""
|
||||
kwargs: dict[str, object] = {}
|
||||
if model_type is not None:
|
||||
kwargs["model_type"] = model_type
|
||||
if ttl is not None:
|
||||
kwargs["ttl"] = ttl
|
||||
await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, **kwargs)
|
||||
|
||||
async def _replicate_to_redis() -> None:
|
||||
try:
|
||||
await user_api_key_cache.async_set_cache(key=key, value=value, **kwargs)
|
||||
except Exception:
|
||||
return
|
||||
|
||||
asyncio.create_task(_replicate_to_redis())
|
||||
|
||||
|
||||
@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:
|
||||
"""Prisma read + L1 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}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
_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,
|
||||
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,
|
||||
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 redis_membership is not _TEAM_MEMBERSHIP_CACHE_MISS:
|
||||
# #region agent log
|
||||
_debug_team_membership_log(
|
||||
"H3",
|
||||
"membership redis hit after l1 miss",
|
||||
{"prisma": False, "source": "redis"},
|
||||
)
|
||||
# #endregion
|
||||
return cast(LiteLLM_TeamMembership | None, redis_membership)
|
||||
|
||||
# #region agent log
|
||||
_debug_team_membership_log("H1", "membership prisma fetch", {"prisma": True})
|
||||
# #endregion
|
||||
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 for user_id: %s, team_id: %s",
|
||||
user_id,
|
||||
team_id,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def get_team_membership(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2155,54 +2314,43 @@ 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 l1_membership is not _TEAM_MEMBERSHIP_CACHE_MISS:
|
||||
# #region agent log
|
||||
_debug_team_membership_log("H4", "membership l1 hit", {"prisma": False, "source": "l1"})
|
||||
# #endregion
|
||||
return cast(LiteLLM_TeamMembership | None, 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 = _team_membership_inflight.get(_key)
|
||||
if inflight is not None:
|
||||
# #region agent log
|
||||
_debug_team_membership_log("H5", "membership coalesced waiter", {"prisma": False, "coalesced": True})
|
||||
# #endregion
|
||||
return await 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,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
_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
|
||||
)
|
||||
_team_membership_inflight[_key] = task
|
||||
task.add_done_callback(lambda _t, k=_key: _team_membership_inflight.pop(k, None))
|
||||
return await task
|
||||
|
||||
|
||||
def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool:
|
||||
|
|
@ -4122,18 +4270,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)
|
||||
|
||||
|
||||
|
|
@ -4169,6 +4320,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 (
|
||||
|
|
@ -4180,6 +4333,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(
|
||||
|
|
@ -4274,6 +4429,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.
|
||||
|
|
@ -4319,6 +4476,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
|
||||
)
|
||||
|
|
@ -4334,6 +4493,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."""
|
||||
|
|
@ -4350,6 +4511,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)
|
||||
|
|
@ -5152,6 +5315,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 (
|
||||
|
|
@ -5160,23 +5325,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):
|
||||
|
|
@ -5195,7 +5362,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
|
||||
|
|
@ -5224,6 +5391,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.
|
||||
|
|
@ -5234,22 +5403,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,
|
||||
|
|
|
|||
|
|
@ -6454,6 +6454,133 @@ 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():
|
||||
"""Concurrent misses for the same member must share one Prisma round-trip."""
|
||||
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_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."""
|
||||
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("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new_callable=AsyncMock,
|
||||
return_value=membership,
|
||||
) as load_membership,
|
||||
patch("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_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue