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.
This commit is contained in:
mateo-berri 2026-06-03 18:04:02 +00:00
parent d781b62112
commit 9764f54f4d
No known key found for this signature in database
2 changed files with 200 additions and 0 deletions

View file

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

View file

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