diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8fbeaf18460..e8e41dc6688 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -5791,6 +5791,7 @@ async def _check_team_member_budget( counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}", fallback_spend=team_member_spend, max_budget=team_member_budget, + fallback_authoritative=loaded_membership is None, ) if not math.isfinite(team_member_budget): diff --git a/tests/e2e/quota_management/budgets/test_team_member_budget_e2e.py b/tests/e2e/quota_management/budgets/test_team_member_budget_e2e.py index aeba96a9509..f3add59b135 100644 --- a/tests/e2e/quota_management/budgets/test_team_member_budget_e2e.py +++ b/tests/e2e/quota_management/budgets/test_team_member_budget_e2e.py @@ -19,14 +19,14 @@ from dataclasses import dataclass import pytest from budget_client import BudgetClient, is_budget_block -from e2e_config import unique_marker +from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import Success, require_successful_call from lifecycle import ResourceManager from models import ChatBody, ChatMessage pytestmark = pytest.mark.e2e -MODEL = "claude-haiku-4-5" +MODEL = CHEAP_OPENAI_MODEL TEAM_BUDGET = 100.0 MEMBER_BUDGET = 3e-6 BURST = 6 @@ -102,9 +102,7 @@ class TestTeamMemberBudget: sent = frozenset(rid for rid in (_send(client, member.key) for _ in range(BURST)) if rid) assert sent, "no member call went through; cannot check attribution" - rows = client.proxy.poll_logs_for_key( - member.key, predicate=lambda rs: bool(sent & {r.request_id for r in rs}) - ) + rows = client.proxy.poll_logs_for_key(member.key, predicate=lambda rs: bool(sent & {r.request_id for r in rs})) logged = [row for row in rows if row.request_id in sent] assert logged, f"none of the member's {len(sent)} calls reached the spend logs" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index dabd97cff0b..be80a90439f 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -4,7 +4,7 @@ import sys import time from collections.abc import Iterator, Mapping from types import SimpleNamespace -from typing import TYPE_CHECKING, Final, Literal, Optional +from typing import TYPE_CHECKING, Final, Literal, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch if TYPE_CHECKING: @@ -65,6 +65,7 @@ from litellm.proxy.auth.auth_checks import ( route_skips_budget_checks, vector_store_access_check, ) +from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting @@ -78,6 +79,7 @@ from litellm.constants import ( from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.utils import PrismaClient from prisma.errors import DataError from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, @@ -8006,6 +8008,99 @@ async def test_check_team_member_budget_fails_closed_when_the_membership_read_hi await _check_team_member_budget(user_object=None, **_restricted_member_check_deps()) +def _unavailable_spend_counter_cache(monkeypatch: pytest.MonkeyPatch) -> DualCache: + redis_cache: Final = cast(RedisCache, MagicMock()) + monkeypatch.setattr(redis_cache, "async_get_cache", AsyncMock(side_effect=RuntimeError("redis unavailable"))) + return DualCache(redis_cache=redis_cache) + + +def _prisma_client_with_membership_lookup(find_unique: AsyncMock) -> PrismaClient: + return cast( + PrismaClient, + SimpleNamespace( + db=SimpleNamespace(litellm_teammembership=SimpleNamespace(find_unique=find_unique)), + ), + ) + + +@pytest.mark.asyncio +async def test_check_team_member_budget_missing_membership_is_verified_with_unavailable_counters( + monkeypatch: pytest.MonkeyPatch, +): + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_BudgetTable + from litellm.proxy.utils import ProxyLogging + + cache: Final = cast(UserApiKeyCache, MagicMock()) + monkeypatch.setattr( + cache, + "async_get_cache", + AsyncMock(return_value=LiteLLM_BudgetTable(budget_id="default-budget-100", max_budget=100.0)), + ) + membership_find_unique: Final = AsyncMock(return_value=None) + prisma_client: Final = _prisma_client_with_membership_lookup(membership_find_unique) + monkeypatch.setattr(proxy_server, "general_settings", {"fail_closed_budget_enforcement": True}) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "spend_counter_cache", _unavailable_spend_counter_cache(monkeypatch)) + + await _check_team_member_budget( + team_object=LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "default-budget-100"}, + ), + user_object=None, + valid_token=UserAPIKeyAuth(token="test-token", user_id="test-user", team_id="test-team"), + prisma_client=prisma_client, + user_api_key_cache=cache, + proxy_logging_obj=ProxyLogging(user_api_key_cache=cache), + team_membership=None, + team_membership_loaded=True, + ) + + membership_find_unique.assert_awaited() + + +@pytest.mark.asyncio +async def test_check_team_member_budget_existing_membership_still_fails_closed_when_counters_are_unavailable( + monkeypatch: pytest.MonkeyPatch, +): + from fastapi import HTTPException + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership + from litellm.proxy.utils import ProxyLogging + + cache: Final = cast(UserApiKeyCache, MagicMock()) + monkeypatch.setattr( + cache, + "async_get_cache", + AsyncMock(return_value=LiteLLM_BudgetTable(budget_id="default-budget-100", max_budget=100.0)), + ) + membership_find_unique: Final = AsyncMock(side_effect=RuntimeError("database unavailable")) + prisma_client: Final = _prisma_client_with_membership_lookup(membership_find_unique) + monkeypatch.setattr(proxy_server, "general_settings", {"fail_closed_budget_enforcement": True}) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "spend_counter_cache", _unavailable_spend_counter_cache(monkeypatch)) + + with pytest.raises(HTTPException) as exc_info: + await _check_team_member_budget( + team_object=LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "default-budget-100"}, + ), + user_object=None, + valid_token=UserAPIKeyAuth(token="test-token", user_id="test-user", team_id="test-team"), + prisma_client=prisma_client, + user_api_key_cache=cache, + proxy_logging_obj=ProxyLogging(user_api_key_cache=cache), + team_membership=LiteLLM_TeamMembership(user_id="test-user", team_id="test-team", spend=0.0), + team_membership_loaded=True, + ) + + assert exc_info.value.status_code == 503 + membership_find_unique.assert_awaited() + + @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