From 855cf82581ec1d314c4540947eaeab489b293145 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 24 May 2026 03:42:54 +0000 Subject: [PATCH] test(mcp): cover per-user env-var management endpoints Adds unit tests for the four /v1/mcp per-user env-var handlers and the _compute_user_env_var_status helper, which were the dominant codecov/patch gap (72 uncovered lines in mcp_management_endpoints.py). Covers happy paths, missing-user-id (400), unknown-server (404), allowed/non-empty value filtering, the swallowed delete error, and the bulk status filter. https://claude.ai/code/session_01X5YQzqswkwcVLtsBbk7Qyh --- .../test_mcp_management_endpoints.py | 379 ++++++++++++++++++ 1 file changed, 379 insertions(+) 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 5d66c184495..49b7eb7b4bb 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 @@ -3286,3 +3286,382 @@ def test_sanitize_mcp_server_for_non_admin_clears_credential_fields(): # server without exposing secrets. assert sanitized.server_id == server.server_id assert sanitized.alias == server.alias + + +def _make_env_var_server( + *, + server_id: str = "srv-1", + server_name: str = "DB Server", + alias: str = "db_server", + env_vars=None, + static_headers=None, +): + """Lightweight server stand-in for the per-user env-var endpoints. + + The handlers only read ``server_id``/``server_name``/``alias``/``env_vars``/ + ``static_headers`` via ``getattr``, so a SimpleNamespace is enough and keeps + the test decoupled from the full Prisma model. + """ + return SimpleNamespace( + server_id=server_id, + server_name=server_name, + alias=alias, + env_vars=env_vars, + static_headers=static_headers, + ) + + +# env_vars with two referenced per-user fields, one unreferenced per-user field +# (must NOT be blocking), and a global value. +_ENV_VARS_MIXED = [ + {"name": "DB_PROTOCOL", "value": "postgres", "scope": "global"}, + { + "name": "CORP_USERNAME", + "value": "", + "scope": "user", + "description": "Your username", + }, + {"name": "CORP_PASSWORD", "value": "", "scope": "user"}, + {"name": "UNUSED_USER_VAR", "value": "", "scope": "user"}, +] +_STATIC_HEADERS_MIXED = { + "Authorization": "${DB_PROTOCOL}://${CORP_USERNAME}:${CORP_PASSWORD}@host/db", +} + + +class TestComputeUserEnvVarStatus: + """Unit tests for the _compute_user_env_var_status helper.""" + + def test_only_referenced_per_user_vars_are_required(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED + ) + status = mgmt_endpoints._compute_user_env_var_status( + server=server, stored_values={"CORP_USERNAME": "alice"} + ) + names = {spec.name for spec in status.required} + # UNUSED_USER_VAR is declared per-user but never referenced -> not blocking. + assert names == {"CORP_USERNAME", "CORP_PASSWORD"} + by_name = {spec.name: spec for spec in status.required} + assert by_name["CORP_USERNAME"].is_set is True + assert by_name["CORP_USERNAME"].value == "alice" + assert by_name["CORP_USERNAME"].description == "Your username" + assert by_name["CORP_PASSWORD"].is_set is False + assert status.missing_count == 1 + assert status.server_id == "srv-1" + assert status.server_name == "DB Server" + assert status.alias == "db_server" + # required is non-empty -> a setup URL is provided. + assert status.setup_url and "srv-1" in status.setup_url + + def test_all_filled_has_zero_missing(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers=_STATIC_HEADERS_MIXED + ) + status = mgmt_endpoints._compute_user_env_var_status( + server=server, + stored_values={"CORP_USERNAME": "alice", "CORP_PASSWORD": "s3cret"}, + ) + assert status.missing_count == 0 + assert all(spec.is_set for spec in status.required) + + def test_static_headers_as_json_string_is_parsed(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, + static_headers='{"Authorization": "${CORP_USERNAME}"}', + ) + status = mgmt_endpoints._compute_user_env_var_status( + server=server, stored_values={} + ) + # Only CORP_USERNAME is referenced via the JSON-string headers. + assert {spec.name for spec in status.required} == {"CORP_USERNAME"} + assert status.missing_count == 1 + + def test_static_headers_invalid_json_string_yields_no_required(self): + server = _make_env_var_server( + env_vars=_ENV_VARS_MIXED, static_headers="not-json{" + ) + status = mgmt_endpoints._compute_user_env_var_status( + server=server, stored_values={} + ) + assert status.required == [] + assert status.missing_count == 0 + # No required fields -> no setup URL. + assert status.setup_url is None + + def test_no_per_user_vars_referenced_yields_no_required(self): + server = _make_env_var_server( + env_vars=[{"name": "DB_PROTOCOL", "value": "postgres", "scope": "global"}], + static_headers={"Authorization": "${DB_PROTOCOL}://host"}, + ) + status = mgmt_endpoints._compute_user_env_var_status( + server=server, stored_values={} + ) + assert status.required == [] + assert status.setup_url is None + + +class TestGetMCPUserEnvVars: + @pytest.mark.asyncio + async def test_returns_status_for_server(self): + server = _make_env_var_server( + 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_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"), + ) + assert result.server_id == "srv-1" + assert result.missing_count == 1 + assert {s.name for s in result.required} == {"CORP_USERNAME", "CORP_PASSWORD"} + + @pytest.mark.asyncio + async def test_missing_user_id_raises_400(self): + with patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ): + 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=""), + ) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unknown_server_raises_404(self): + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None) + ), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.get_mcp_user_env_vars( + server_id="missing", + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), + ) + assert exc.value.status_code == 404 + + +class TestStoreMCPUserEnvVars: + @pytest.mark.asyncio + async def test_persists_only_allowed_non_empty_values(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, "store_user_env_vars", store_mock), + ): + result = await mgmt_endpoints.store_mcp_user_env_vars( + server_id="srv-1", + payload=mgmt_endpoints.MCPUserEnvVarsRequest( + values={ + "CORP_USERNAME": "alice", + "CORP_PASSWORD": "", # empty -> dropped + "NOT_A_DECLARED_VAR": "x", # unknown -> dropped + } + ), + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), + ) + # Only the declared, non-empty value is persisted. + store_mock.assert_awaited_once() + _, _, _, persisted = store_mock.await_args.args + assert persisted == {"CORP_USERNAME": "alice"} + # CORP_PASSWORD remains unset in the returned status. + assert result.missing_count == 1 + + @pytest.mark.asyncio + async def test_missing_user_id_raises_400(self): + with patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.store_mcp_user_env_vars( + server_id="srv-1", + payload=mgmt_endpoints.MCPUserEnvVarsRequest(values={}), + user_api_key_dict=generate_mock_user_api_key_auth(user_id=""), + ) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unknown_server_raises_404(self): + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None) + ), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.store_mcp_user_env_vars( + server_id="missing", + payload=mgmt_endpoints.MCPUserEnvVarsRequest(values={}), + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), + ) + assert exc.value.status_code == 404 + + +class TestClearMCPUserEnvVars: + @pytest.mark.asyncio + async def test_clears_and_returns_empty_status(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, "delete_user_env_vars", delete_mock), + ): + result = 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"), + ) + delete_mock.assert_awaited_once() + # Everything is now unset. + assert result.missing_count == 2 + assert all(not spec.is_set for spec in result.required) + + @pytest.mark.asyncio + async def test_delete_error_is_swallowed(self): + server = _make_env_var_server( + 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, + "delete_user_env_vars", + AsyncMock(side_effect=Exception("already gone")), + ), + ): + # Should not raise even though delete blows up. + result = 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"), + ) + assert result.missing_count == 2 + + @pytest.mark.asyncio + async def test_missing_user_id_raises_400(self): + with patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ): + 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=""), + ) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_unknown_server_raises_404(self): + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None) + ), + ): + with pytest.raises(HTTPException) as exc: + await mgmt_endpoints.clear_mcp_user_env_vars( + server_id="missing", + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice"), + ) + assert exc.value.status_code == 404 + + +class TestListMCPUserEnvVarStatus: + @pytest.mark.asyncio + async def test_no_user_id_returns_empty(self): + with patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ): + result = await mgmt_endpoints.list_mcp_user_env_var_status( + user_api_key_dict=generate_mock_user_api_key_auth(user_id="") + ) + assert result == [] + + @pytest.mark.asyncio + async def test_no_accessible_servers_returns_empty(self): + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[]), + ), + ): + result = await mgmt_endpoints.list_mcp_user_env_var_status( + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice") + ) + assert result == [] + + @pytest.mark.asyncio + async def test_only_servers_with_required_fields_are_returned(self): + server_with = _make_env_var_server( + server_id="srv-with", + env_vars=_ENV_VARS_MIXED, + static_headers=_STATIC_HEADERS_MIXED, + ) + # No per-user var is referenced -> contributes no status entry. + server_without = _make_env_var_server( + server_id="srv-without", + env_vars=[{"name": "DB_PROTOCOL", "value": "postgres", "scope": "global"}], + static_headers={"Authorization": "${DB_PROTOCOL}://host"}, + ) + with ( + patch.object( + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( + mgmt_endpoints, + "get_all_mcp_servers_for_user", + AsyncMock(return_value=[server_with, server_without]), + ), + patch.object( + mgmt_endpoints, + "get_user_env_vars_bulk", + AsyncMock(return_value={"srv-with": {"CORP_USERNAME": "alice"}}), + ), + ): + result = await mgmt_endpoints.list_mcp_user_env_var_status( + user_api_key_dict=generate_mock_user_api_key_auth(user_id="alice") + ) + assert [s.server_id for s in result] == ["srv-with"] + assert result[0].missing_count == 1