From 5a5dd410523a8bed386ce7d9300a86b314c4dccb Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 24 Sep 2026 14:35:10 -0700 Subject: [PATCH] fix(auth): keep main's team member budget enforcement, alert once per request Restore the builder's team member budget check and 422 message exactly as on main and drop the MCP-only gate. The builder sends the member alert only on the request it rejects; common_checks sends it for requests that get past the builder, so no request alerts twice. --- litellm/proxy/auth/auth_checks.py | 8 +- litellm/proxy/auth/user_api_key_auth.py | 134 +++---- .../mcp_management_endpoints.py | 16 +- .../proxy/auth/test_user_api_key_auth.py | 336 ++++++++---------- .../test_mcp_management_endpoints.py | 12 +- 5 files changed, 223 insertions(+), 283 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 0f15eb3da97..2c93aad820c 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -5648,16 +5648,12 @@ async def _check_team_member_budget( ) if team_member_spend >= team_member_budget: - entity_id: Final = f"{valid_token.user_id}:{team_object.team_id}" raise litellm.BudgetExceededError( current_cost=team_member_spend, max_budget=team_member_budget, - message=( - f"Budget has been exceeded! TeamMember={entity_id} " - f"Current cost: {team_member_spend}, Max budget: {team_member_budget}" - ), + message=f"Budget has been exceeded! User={valid_token.user_id} in Team={team_object.team_id} Current cost: {team_member_spend}, Max budget: {team_member_budget}", entity_type=Litellm_EntityType.TEAM_MEMBER.value, - entity_id=entity_id, + entity_id=f"{valid_token.user_id}:{team_object.team_id}", ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 47a8ed0b227..cad34ed5553 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -46,11 +46,11 @@ from litellm.proxy.auth.auth_checks import ( _cache_key_object, _can_object_call_model, _check_end_user_budget, - _check_team_member_budget, _delete_cache_key_object, _get_user_role, _is_model_cost_zero, _is_user_proxy_admin, + _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, _virtual_key_max_budget_check, _virtual_key_soft_budget_check, @@ -119,6 +119,7 @@ from litellm.proxy.common_utils.model_listing_utils import claude_code_requested from litellm.proxy.common_utils.realtime_utils import _realtime_request_body from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, + team_membership_auth_cache_key, ) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup @@ -133,6 +134,7 @@ from litellm.proxy.utils import ( ProxyLogging, normalize_route_for_root_path, ) +from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.secret_managers.main import get_secret_bool from litellm.types.services import ServiceTypes @@ -2237,6 +2239,76 @@ async def _user_api_key_auth_builder( if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) + # Check 3. Check if user is in their team budget + if not skip_budget_checks and valid_token.team_member_spend is not None: + _user_id: Final = valid_token.user_id + _team_id: Final = valid_token.team_id + if prisma_client is not None and _user_id is not None and _team_id is not None: + _cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id) + + team_member_info = await user_api_key_cache.async_get_cache( + key=_cache_key, + model_type=LiteLLM_TeamMembership, + ) + if team_member_info is None: + # read from DB + _db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first( + where={ + "user_id": _user_id, + "team_id": _team_id, + }, + include={"litellm_budget_table": True}, + ) + if _db_member is not None: + team_member_info = LiteLLM_TeamMembership(**_db_member.model_dump()) + await user_api_key_cache.async_set_cache( + key=_cache_key, + value=team_member_info, + model_type=LiteLLM_TeamMembership, + ttl=5, + ) + + if team_member_info is not None and team_member_info.litellm_budget_table is not None: + team_member_budget: Final = team_member_info.litellm_budget_table.effective_max_budget( + now=datetime.now(timezone.utc), + ) + if team_member_budget is not None and team_member_budget > 0: + # Read from cross-pod counter (Redis-first) if available + from litellm.proxy.proxy_server import get_current_spend + + team_member_spend = valid_token.team_member_spend + if valid_token.user_id is not None and valid_token.team_id is not None: + team_member_spend = await get_current_spend( + counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}", + fallback_spend=team_member_spend, + max_budget=team_member_budget, + ) + if team_member_spend >= team_member_budget: + # common_checks sends this alert on requests that get past here, so only the + # request rejected here sends it from the builder. + _team_member_max_budget_alert_check( + team_id=_team_id, + team_alias=valid_token.team_alias, + team_metadata=valid_token.team_metadata, + organization_id=valid_token.org_id, + user_id=_user_id, + user_email=user_obj.user_email if user_obj is not None else None, + proxy_logging_obj=proxy_logging_obj, + spend=team_member_spend, + max_budget=team_member_budget, + ) + _entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}" + raise litellm.BudgetExceededError( + current_cost=team_member_spend, + max_budget=team_member_budget, + message=( + f"Budget has been exceeded! TeamMember={_entity_id} " + f"Current cost: {team_member_spend}, Max budget: {team_member_budget}" + ), + entity_type=Litellm_EntityType.TEAM_MEMBER.value, + entity_id=_entity_id, + ) + # Check 3. If token is expired if valid_token.expires is not None: current_time = datetime.now(timezone.utc) @@ -3098,66 +3170,6 @@ def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) -> ) -async def enforce_team_member_budget_without_common_checks( - user_api_key_auth_obj: UserAPIKeyAuth, - request: Request, - request_data: dict, - route: str, - api_key: str, -) -> None: - """Team member budget gate for callers that stop at ``_user_api_key_auth_builder`` and never reach - ``common_checks`` (the MCP OAuth dependency), so an over-budget member stays blocked there as before. - Failures go through the builder's exception handler so the caller still gets the 422. - """ - from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache - - team_id: Final = user_api_key_auth_obj.team_id - user_id: Final = user_api_key_auth_obj.user_id - if prisma_client is None or team_id is None or team_id == UI_TEAM_ID or user_id is None: - return - parent_otel_span: Final = user_api_key_auth_obj.parent_otel_span - try: - team_object: Final = await get_team_object( - 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 HTTPException: # no team row means no member budget to enforce - return - try: - user_object = await get_user_object( - user_id=user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - user_id_upsert=False, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - except Exception: # noqa: BLE001 # the user row only supplies the alert email; enforcement does not need it - user_object = None - try: - await _check_team_member_budget( - team_object=team_object, - user_object=user_object, - valid_token=user_api_key_auth_obj, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - except litellm.BudgetExceededError as e: - await UserAPIKeyAuthExceptionHandler._handle_authentication_error( - e=e, - request=request, - request_data=request_data, - route=route, - parent_otel_span=parent_otel_span, - api_key=api_key, - resolved_identity=user_api_key_auth_obj, - ) - - async def _authorize_authenticated_request( user_api_key_auth_obj: UserAPIKeyAuth, request: Request, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index b6fcfeb3aa2..9ad78876043 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -205,10 +205,8 @@ if MCP_AVAILABLE: UserMCPManagementMode, is_per_server_oauth_discovery_eligible, ) - from litellm.proxy.auth.auth_utils import get_request_route from litellm.proxy.auth.user_api_key_auth import ( _user_api_key_auth_builder, - enforce_team_member_budget_without_common_checks, user_api_key_auth, ) from litellm.proxy.common_utils.http_parsing_utils import ( @@ -2012,6 +2010,9 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 global_mcp_server_manager, ) + from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415 + get_request_route, + ) server_id: Final[str] = request.path_params.get("server_id", "") if server_id: @@ -2048,7 +2049,7 @@ if MCP_AVAILABLE: request_data = await _read_request_body(request=request) request_data = populate_request_with_path_params(request_data=request_data, request=request) - user_api_key_dict: Final = await _user_api_key_auth_builder( + return await _user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header="", @@ -2057,15 +2058,6 @@ if MCP_AVAILABLE: azure_apim_header=None, request_data=request_data, ) - # This dependency never reaches common_checks, which is the only other place the member budget is enforced. - await enforce_team_member_budget_without_common_checks( - user_api_key_auth_obj=user_api_key_dict, - request=request, - request_data=request_data, - route=get_request_route(request), - api_key=api_key, - ) - return user_api_key_dict async def _get_cached_temporary_mcp_server_or_404( server_id: str, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ad99d5343f4..5238de9b6bb 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -7855,30 +7855,6 @@ async def test_temp_budget_increase_applied_for_cached_key(): assert cached_after.max_budget == 2.0 -async def _authenticate_and_authorize(mock_request, api_key): - """Builder then the single common_checks gate, the same sequence user_api_key_auth runs.""" - from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request - - request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} - auth_obj = await _user_api_key_auth_builder( - request=mock_request, - api_key=f"Bearer {api_key}", - azure_api_key_header="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - request_data=request_data, - ) - recovered = await _authorize_authenticated_request( - user_api_key_auth_obj=auth_obj, - request=mock_request, - request_data=request_data, - route="/v1/messages", - api_key=f"Bearer {api_key}", - ) - return recovered or auth_obj - - @pytest.mark.asyncio @pytest.mark.parametrize( "team_member_spend, expect_blocked", @@ -7892,7 +7868,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe """A team member counter sitting exactly at the cap (where a resized reservation lands it) must be rejected by the cached-key auth path like every other budget check.""" from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj - from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key + from litellm.proxy.common_utils.user_api_key_cache import team_membership_auth_cache_key from litellm.proxy.utils import hash_token api_key = "sk-team-member-exact-cap" @@ -7922,7 +7898,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe value=LiteLLM_UserTable(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER), ) await user_api_key_cache.async_set_cache( - key=team_membership_reservation_cache_key(team_id=team_id, user_id=user_id), + key=team_membership_auth_cache_key(team_id=team_id, user_id=user_id), value=LiteLLM_TeamMembership( user_id=user_id, team_id=team_id, @@ -7944,7 +7920,15 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) async def _auth(): - return await _authenticate_and_authorize(mock_request, api_key) + return await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]}, + ) with ( patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam @@ -7974,6 +7958,136 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe assert f"TeamMember={user_id}:{team_id}" in exc_info.value.message +@pytest.mark.asyncio +@pytest.mark.parametrize( + "expiry_offset, expect_blocked", + [ + (timedelta(days=1), False), + (timedelta(days=-1), True), + ], +) +async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset, expect_blocked): + """A member over their permanent cap is admitted while a temp_budget_increase is unexpired + and blocked again once it expires, on the cached-key auth path.""" + from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj + from litellm.proxy.common_utils.user_api_key_cache import team_membership_auth_cache_key + from litellm.proxy.utils import hash_token + + api_key = "sk-team-member-temp-budget" + hashed_token = hash_token(api_key) + team_id = "team-temp-budget" + user_id = "user-temp-budget" + team_member_spend = 2.5 + + user_api_key_cache = DualCache() + await _cache_key_object( + hashed_token=hashed_token, + user_api_key_obj=UserAPIKeyAuth( + token=hashed_token, + team_id=team_id, + user_id=user_id, + team_member_spend=team_member_spend, + ), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=None, + ) + await user_api_key_cache.async_set_cache( + key=f"team_id:{team_id}", + value=LiteLLM_TeamTableCachedObj(team_id=team_id), + ) + await user_api_key_cache.async_set_cache( + key=user_id, + value=LiteLLM_UserTable(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER), + ) + await user_api_key_cache.async_set_cache( + key=team_membership_auth_cache_key(team_id=team_id, user_id=user_id), + value=LiteLLM_TeamMembership( + user_id=user_id, + team_id=team_id, + spend=team_member_spend, + budget_id="budget-temp", + litellm_budget_table=LiteLLM_BudgetTable( + max_budget=2.0, + temp_budget_increase=1.0, + temp_budget_expiry=datetime.now(timezone.utc) + expiry_offset, + ), + ), + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/messages" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {api_key}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + async def _auth(): + return await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]}, + ) + + with ( + patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam + "litellm.proxy.proxy_server.general_settings", {"disable_budget_reservation": True} + ), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: module-global proxy state + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state + patch( # test-quality-ok: seed the cached key, team and membership without a DB + "litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache + ), + patch( # test-quality-ok: module-global proxy state + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ), + patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=team_member_spend), + ), + ): + if not expect_blocked: + result = await _auth() + assert result.team_member_spend == team_member_spend + return + with pytest.raises(ProxyException) as exc_info: + await _auth() + + assert exc_info.value.type == ProxyErrorTypes.budget_exceeded + assert "Max budget: 2.0" in exc_info.value.message + + +async def _authenticate_and_authorize(mock_request, api_key): + """Builder then the single common_checks gate, the same sequence user_api_key_auth runs.""" + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + + request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} + auth_obj = await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data=request_data, + ) + recovered = await _authorize_authenticated_request( + user_api_key_auth_obj=auth_obj, + request=mock_request, + request_data=request_data, + route="/v1/messages", + api_key=f"Bearer {api_key}", + ) + return recovered or auth_obj + + @pytest.mark.asyncio @pytest.mark.parametrize( "team_member_spend, expect_blocked, expected_alerts", @@ -8096,174 +8210,6 @@ async def test_cached_key_team_member_budget_emails_configured_thresholds( assert call_info.max_budget_alert_emails == alert_emails -@pytest.mark.asyncio -@pytest.mark.parametrize( - "expiry_offset, expect_blocked", - [ - (timedelta(days=1), False), - (timedelta(days=-1), True), - ], -) -async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset, expect_blocked): - """A member over their permanent cap is admitted while a temp_budget_increase is unexpired - and blocked again once it expires, on the cached-key auth path.""" - from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj - from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key - from litellm.proxy.utils import hash_token - - api_key = "sk-team-member-temp-budget" - hashed_token = hash_token(api_key) - team_id = "team-temp-budget" - user_id = "user-temp-budget" - team_member_spend = 2.5 - - user_api_key_cache = DualCache() - await _cache_key_object( - hashed_token=hashed_token, - user_api_key_obj=UserAPIKeyAuth( - token=hashed_token, - team_id=team_id, - user_id=user_id, - team_member_spend=team_member_spend, - ), - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=None, - ) - await user_api_key_cache.async_set_cache( - key=f"team_id:{team_id}", - value=LiteLLM_TeamTableCachedObj(team_id=team_id), - ) - await user_api_key_cache.async_set_cache( - key=user_id, - value=LiteLLM_UserTable(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER), - ) - await user_api_key_cache.async_set_cache( - key=team_membership_reservation_cache_key(team_id=team_id, user_id=user_id), - value=LiteLLM_TeamMembership( - user_id=user_id, - team_id=team_id, - spend=team_member_spend, - budget_id="budget-temp", - litellm_budget_table=LiteLLM_BudgetTable( - max_budget=2.0, - temp_budget_increase=1.0, - temp_budget_expiry=datetime.now(timezone.utc) + expiry_offset, - ), - ), - ) - - mock_request = MagicMock() - mock_request.url.path = "/v1/messages" - mock_request.method = "POST" - mock_request.headers = {"authorization": f"Bearer {api_key}"} - mock_request.query_params = {} - mock_request.state = SimpleNamespace() - - proxy_logging_obj = MagicMock() - proxy_logging_obj.budget_alerts = AsyncMock() - proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - - async def _auth(): - return await _authenticate_and_authorize(mock_request, api_key) - - with ( - patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam - "litellm.proxy.proxy_server.general_settings", {"disable_budget_reservation": True} - ), - patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: module-global proxy state - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state - patch( # test-quality-ok: seed the cached key, team and membership without a DB - "litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache - ), - patch( # test-quality-ok: module-global proxy state - "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj - ), - patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares - "litellm.proxy.proxy_server.get_current_spend", - new=AsyncMock(return_value=team_member_spend), - ), - ): - if not expect_blocked: - result = await _auth() - assert result.team_member_spend == team_member_spend - return - with pytest.raises(ProxyException) as exc_info: - await _auth() - - assert exc_info.value.type == ProxyErrorTypes.budget_exceeded - assert "Max budget: 2.0" in exc_info.value.message - - -@pytest.mark.asyncio -@pytest.mark.parametrize("team_member_spend, expect_blocked", [(2.4, True), (2.39, False)]) -async def test_team_member_budget_enforced_for_builder_only_callers(team_member_spend, expect_blocked): - """Callers that stop at the builder and skip common_checks (the MCP OAuth dependency) still reject a member - at their team member budget with the same 422 the full auth flow returns.""" - from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.user_api_key_auth import enforce_team_member_budget_without_common_checks - from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key - - team_id = "team-builder-only" - user_id = "user-builder-only" - user_api_key_cache = DualCache() - await user_api_key_cache.async_set_cache(key=f"team_id:{team_id}", value=LiteLLM_TeamTableCachedObj(team_id=team_id)) - await user_api_key_cache.async_set_cache( - key=user_id, value=LiteLLM_UserTable(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) - ) - await user_api_key_cache.async_set_cache( - key=team_membership_reservation_cache_key(team_id=team_id, user_id=user_id), - value=LiteLLM_TeamMembership( - user_id=user_id, - team_id=team_id, - spend=team_member_spend, - budget_id="budget-builder-only", - litellm_budget_table=LiteLLM_BudgetTable(max_budget=2.4), - ), - ) - - mock_request = MagicMock() - mock_request.url.path = "/v1/mcp/server/oauth/srv/authorize" - mock_request.method = "GET" - mock_request.headers = {} - mock_request.query_params = {} - mock_request.state = SimpleNamespace() - - proxy_logging_obj = MagicMock() - proxy_logging_obj.budget_alerts = AsyncMock() - proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - - async def _enforce(): - await enforce_team_member_budget_without_common_checks( - user_api_key_auth_obj=UserAPIKeyAuth(api_key="hashed", team_id=team_id, user_id=user_id), - request=mock_request, - request_data={}, - route="/v1/mcp/server/oauth/srv/authorize", - api_key="Bearer sk-builder-only", - ) - - with ( - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state - patch( # test-quality-ok: seed the team, user and membership without a DB - "litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache - ), - patch( # test-quality-ok: module-global proxy state - "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj - ), - patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares - "litellm.proxy.proxy_server.get_current_spend", - new=AsyncMock(return_value=team_member_spend), - ), - ): - if not expect_blocked: - await _enforce() - return - with pytest.raises(ProxyException) as exc_info: - await _enforce() - - assert exc_info.value.type == ProxyErrorTypes.budget_exceeded - assert f"TeamMember={user_id}:{team_id}" in exc_info.value.message - - async def _proxy_exception_for_key( api_key: str, general_settings: dict[str, bool], diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index d6397375be0..53645e62034 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -2552,9 +2552,7 @@ class TestTemporaryMCPSessionEndpoints: expected_auth = generate_mock_user_api_key_auth( user_role=LitellmUserRoles.PROXY_ADMIN, api_key=api_key_in_cookie ) - fake_proxy_server = types.SimpleNamespace( - master_key=master_key, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None - ) + fake_proxy_server = types.SimpleNamespace(master_key=master_key) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), @@ -2631,9 +2629,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = non_oauth_server mock_manager.get_mcp_server_by_name.return_value = None - fake_proxy_server = types.SimpleNamespace( - master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None - ) + fake_proxy_server = types.SimpleNamespace(master_key=None) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), @@ -2685,9 +2681,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = internal_server mock_manager.get_mcp_server_by_name.return_value = None - fake_proxy_server = types.SimpleNamespace( - master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None - ) + fake_proxy_server = types.SimpleNamespace(master_key=None) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),