mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
d781b62112
commit
9764f54f4d
2 changed files with 200 additions and 0 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue