fix(proxy): treat a missing team membership row as verified zero spend under fail-closed budgets

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-29 18:27:09 +00:00
parent 200224e5d6
commit 95585ba57a
3 changed files with 100 additions and 6 deletions

View file

@ -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):

View file

@ -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"

View file

@ -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