diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index d1523861492..47a8ed0b227 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -46,6 +46,7 @@ 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, @@ -3097,6 +3098,66 @@ 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 9ad78876043..b6fcfeb3aa2 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -205,8 +205,10 @@ 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 ( @@ -2010,9 +2012,6 @@ 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: @@ -2049,7 +2048,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) - return await _user_api_key_auth_builder( + user_api_key_dict: Final = await _user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header="", @@ -2058,6 +2057,15 @@ 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 b70e3a55326..ad99d5343f4 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 @@ -8194,6 +8194,76 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset 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 53645e62034..d6397375be0 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,7 +2552,9 @@ 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) + fake_proxy_server = types.SimpleNamespace( + master_key=master_key, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None + ) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), @@ -2629,7 +2631,9 @@ 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) + fake_proxy_server = types.SimpleNamespace( + master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None + ) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), @@ -2681,7 +2685,9 @@ 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) + fake_proxy_server = types.SimpleNamespace( + master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None + ) with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),