diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index f877e9a10db..ae6fab83a7e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -850,6 +850,20 @@ BUDGET_ENFORCED_SIDE_EFFECT_ROUTES: Final = frozenset( } ) +MCP_DISCOVERY_ROUTES: Final = frozenset({"/mcp-rest/tools/list", "/v1/mcp/tools", "/mcp/tools", "/mcp/tools/list"}) + +MCP_ZERO_SPEND_JSONRPC_METHODS: Final = frozenset({"initialize", "notifications/initialized", "ping", "tools/list"}) + +MCP_TOOL_CALL_ROUTES: Final = frozenset({"/mcp/tools/call", "/mcp-rest/tools/call"}) + + +def is_mcp_discovery_request(route: str, request_body: dict) -> bool: + if route in MCP_DISCOVERY_ROUTES: + return True + if route in MCP_TOOL_CALL_ROUTES or not (route == "/mcp" or route.startswith("/mcp/")): + return False + return request_body.get("method") in MCP_ZERO_SPEND_JSONRPC_METHODS + async def common_checks( request_body: dict, @@ -900,7 +914,11 @@ async def common_checks( skip_all_budget_checks: Final = skip_budget_checks or ( route not in BUDGET_ENFORCED_SIDE_EFFECT_ROUTES - and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route)) + and ( + route in MODEL_DISCOVERY_ROUTES + or is_mcp_discovery_request(route=route, request_body=request_body) + or not RouteChecks.is_llm_api_route(route=route) + ) ) membership_user_id: Final = ( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9491f77ecfc..da09833038e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -59,6 +59,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_membership, get_team_object, get_user_object, + is_mcp_discovery_request, is_valid_fallback_model, jwt_key_mapping_cache_key, resolve_and_validate_end_user_id, @@ -1657,8 +1658,8 @@ async def _user_api_key_auth_builder( llm_router=llm_router, team_id=valid_token.team_id, ) - skip_budget_checks = False - if model is not None and llm_router is not None: + skip_budget_checks = is_mcp_discovery_request(route=route, request_body=request_data) + if not skip_budget_checks and model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) @@ -2098,8 +2099,8 @@ async def _user_api_key_auth_builder( llm_router=llm_router, team_id=valid_token.team_id, ) - skip_budget_checks = False - if model is not None and llm_router is not None: + skip_budget_checks = is_mcp_discovery_request(route=route, request_body=request_data) + if not skip_budget_checks and model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 78acc33165f..82889c14d90 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -5187,6 +5187,125 @@ async def test_inference_route_still_enforces_team_budget(): ) +MCP_DISCOVERY_REST_ROUTES = [ + "/mcp-rest/tools/list", + "/v1/mcp/tools", + "/mcp/tools", + "/mcp/tools/list", +] + + +@pytest.mark.parametrize("route", MCP_DISCOVERY_REST_ROUTES) +@pytest.mark.asyncio +async def test_mcp_tool_discovery_route_bypasses_team_budget(route): + """An exhausted budget must not block zero-spend MCP tool discovery.""" + from litellm.proxy.auth.auth_checks import common_checks + + team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0) + + result = await common_checks( + request_body={}, + team_object=team_object, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=AsyncMock(), + valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"), + request=MagicMock(), + ) + + assert result is True + + +@pytest.mark.parametrize("method", ["initialize", "notifications/initialized", "ping", "tools/list"]) +@pytest.mark.parametrize("route", ["/mcp", "/mcp/some-server"]) +@pytest.mark.asyncio +async def test_mcp_jsonrpc_zero_spend_method_bypasses_team_budget(route, method): + """An exhausted budget must not block zero-spend JSON-RPC methods on the /mcp transport.""" + from litellm.proxy.auth.auth_checks import common_checks + + team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0) + + result = await common_checks( + request_body={"jsonrpc": "2.0", "id": 1, "method": method, "params": {}}, + team_object=team_object, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=AsyncMock(), + valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"), + request=MagicMock(), + ) + + assert result is True + + +@pytest.mark.parametrize( + "route,request_body", + [ + ("/mcp", {"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "add"}}), + ("/mcp/some-server", {"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "add"}}), + ("/mcp-rest/tools/call", {"name": "add", "arguments": {}}), + ("/mcp/tools/call", {"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}}), + ], +) +@pytest.mark.asyncio +async def test_mcp_tool_call_route_still_enforces_team_budget(route, request_body): + """tools/call is spend-bearing and must stay budget-enforced, including when the + tools/call REST route carries a discovery-looking JSON-RPC body.""" + from litellm.proxy.auth.auth_checks import common_checks + + team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0) + + with pytest.raises(litellm.BudgetExceededError): + await common_checks( + request_body=request_body, + team_object=team_object, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=AsyncMock(), + valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"), + request=MagicMock(), + ) + + +@pytest.mark.parametrize( + "route,request_body,expected", + [ + ("/mcp-rest/tools/list", {}, True), + ("/v1/mcp/tools", {}, True), + ("/mcp/tools", {}, True), + ("/mcp/tools/list", {}, True), + ("/mcp", {"method": "initialize"}, True), + ("/mcp", {"method": "notifications/initialized"}, True), + ("/mcp", {"method": "ping"}, True), + ("/mcp", {"method": "tools/list"}, True), + ("/mcp/some-server", {"method": "tools/list"}, True), + ("/mcp", {"method": "tools/call"}, False), + ("/mcp/some-server", {"method": "tools/call"}, False), + ("/mcp", {}, False), + ("/mcp-rest/tools/call", {"method": "initialize"}, False), + ("/mcp/tools/call", {"method": "initialize"}, False), + ("/v1/chat/completions", {"method": "initialize"}, False), + ("/mcp-rest/other", {"method": "tools/list"}, False), + ], +) +def test_is_mcp_discovery_request(route, request_body, expected): + from litellm.proxy.auth.auth_checks import is_mcp_discovery_request + + assert is_mcp_discovery_request(route=route, request_body=request_body) is expected + + @pytest.mark.asyncio async def test_virtual_key_max_budget_error_names_the_key(): """BudgetExceededError for a virtual key must name the key (alias + masked key)