From 9764f54f4d8bf7b35ea1a4b2de5bb07d3cee7469 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 3 Jun 2026 18:04:02 +0000 Subject: [PATCH] fix(mcp): gate per-user env-var endpoints by server access The per-user env-var endpoints (GET/POST/DELETE /server/{id}/user-env-vars) fetched the server by id and returned its name, alias, and required-credential metadata, or persisted/cleared stored values, without checking that the caller can access that server. A non-admin could query or mutate env-var state for any server id. Apply the same access gate fetch_mcp_server uses so non-admins are limited to servers in their allowed set. --- .../mcp_management_endpoints.py | 30 ++++ .../test_mcp_management_endpoints.py | 170 ++++++++++++++++++ 2 files changed, 200 insertions(+) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f5b4b583599..b075ab19617 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2132,6 +2132,33 @@ if MCP_AVAILABLE: # ── Per-user MCP env var endpoints ──────────────────────────────────────── + async def _authorize_mcp_server_access( + prisma_client, + user_api_key_dict: UserAPIKeyAuth, + server_id: str, + ) -> None: + """Raise 403 if a non-admin caller cannot access this MCP server. + + Mirrors the access gate in ``fetch_mcp_server`` so the per-user env-var + endpoints can't be used to read or mutate state for servers outside the + caller's allowed set. + """ + if _user_has_admin_view(user_api_key_dict): + return + accessible = await get_all_mcp_servers_for_user( + prisma_client, user_api_key_dict + ) + if not does_mcp_server_exist(accessible, server_id): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + f"User does not have permission to access mcp server with id {server_id}. " + "You can only manage env vars for mcp servers that you have access to." + ) + }, + ) + def _compute_user_env_var_status( *, server: LiteLLM_MCPServerTable, @@ -2208,6 +2235,7 @@ if MCP_AVAILABLE: status_code=status.HTTP_404_NOT_FOUND, detail={"error": f"MCP Server {server_id} not found"}, ) + await _authorize_mcp_server_access(prisma_client, user_api_key_dict, server_id) stored = await get_user_env_vars(prisma_client, user_id, server_id) return _compute_user_env_var_status(server=server, stored_values=stored) @@ -2238,6 +2266,7 @@ if MCP_AVAILABLE: status_code=status.HTTP_404_NOT_FOUND, detail={"error": f"MCP Server {server_id} not found"}, ) + await _authorize_mcp_server_access(prisma_client, user_api_key_dict, server_id) # Filter to only known per-user var names declared by the admin — # never persist arbitrary keys the user invents. _, user_specs = parse_admin_env_vars(getattr(server, "env_vars", None)) @@ -2274,6 +2303,7 @@ if MCP_AVAILABLE: status_code=status.HTTP_404_NOT_FOUND, detail={"error": f"MCP Server {server_id} not found"}, ) + await _authorize_mcp_server_access(prisma_client, user_api_key_dict, server_id) try: await delete_user_env_vars(prisma_client, user_id, server_id) except Exception: 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 1162cf9c978..8c06b9a8883 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 @@ -3729,3 +3729,173 @@ class TestListMCPUserEnvVarStatus: ) assert [s.server_id for s in result] == ["srv-with"] assert result[0].missing_count == 1 + + +class TestMCPUserEnvVarsAccessControl: + """Per-server env-var endpoints must enforce the same access gate as + fetch_mcp_server: a non-admin caller can only touch servers in their + allowed set.""" + + @pytest.mark.asyncio + async def test_get_forbidden_for_non_admin_without_access(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED + ) + get_user_env_vars = AsyncMock(return_value={}) + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[_make_env_var_server(server_id="other")]), + ), + patch.object(mgmt_endpoints, "get_user_env_vars", get_user_env_vars), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.get_mcp_user_env_vars( + server_id="srv-1", + user_api_key_dict=generate_mock_user_api_key_auth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + assert exc.value.status_code == 403 + get_user_env_vars.assert_not_awaited() + + @pytest.mark.asyncio + async def test_store_forbidden_for_non_admin_without_access(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED + ) + store_mock = AsyncMock() + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[]), + ), + patch.object(mgmt_endpoints, "store_user_env_vars", store_mock), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.store_mcp_user_env_vars( + server_id="srv-1", + payload=mgmt_endpoints.MCPUserEnvVarsRequest( + values={"CORP_USERNAME": "alice"} + ), + user_api_key_dict=generate_mock_user_api_key_auth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + assert exc.value.status_code == 403 + store_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_clear_forbidden_for_non_admin_without_access(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED + ) + delete_mock = AsyncMock() + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[]), + ), + patch.object(mgmt_endpoints, "delete_user_env_vars", delete_mock), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.clear_mcp_user_env_vars( + server_id="srv-1", + user_api_key_dict=generate_mock_user_api_key_auth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + assert exc.value.status_code == 403 + delete_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_get_allowed_for_non_admin_with_access(self): + server = _make_env_var_server( + server_id="srv-1", + env_vars=_ENV_VARS_MIXED, + static_headers=_STATIC_HEADERS_MIXED, + ) + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[server]), + ), + patch.object( + mgmt_endpoints, + "get_user_env_vars", + AsyncMock(return_value={"CORP_USERNAME": "alice"}), + ), + ): + result = await mgmt_endpoints.get_mcp_user_env_vars( + server_id="srv-1", + user_api_key_dict=generate_mock_user_api_key_auth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + assert result.server_id == "srv-1" + assert result.missing_count == 1 + + @pytest.mark.asyncio + async def test_admin_bypasses_access_check(self): + """Proxy admins must not be filtered by get_all_mcp_servers_for_user.""" + server = _make_env_var_server( + server_id="srv-1", + env_vars=_ENV_VARS_MIXED, + static_headers=_STATIC_HEADERS_MIXED, + ) + access_list_mock = AsyncMock(return_value=[]) + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=server) + ), + patch.object( + mgmt_endpoints, "get_all_mcp_servers_for_user", access_list_mock + ), + patch.object( + mgmt_endpoints, "get_user_env_vars", AsyncMock(return_value={}) + ), + ): + result = await mgmt_endpoints.get_mcp_user_env_vars( + server_id="srv-1", + user_api_key_dict=generate_mock_user_api_key_auth( + user_id="admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ), + ) + assert result.server_id == "srv-1" + access_list_mock.assert_not_awaited()