diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index a2ce3307061..1d62b325dec 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1228,6 +1228,23 @@ def _remaining_token_seconds(expires_at: str | None) -> int | None: return remaining if remaining > 0 else None +async def get_active_submitted_mcp_server_ids_for_user( + prisma_client: PrismaClient, + user_id: str, +) -> list[str]: + """Return active BYOM servers submitted by this user (creator visibility).""" + if not user_id: + return [] + + rows = await MCPServerRepository(prisma_client).table.find_many( + where={ + "submitted_by": user_id, + "approval_status": MCPApprovalStatus.active, + }, + ) + return [row.server_id for row in rows] + + async def approve_mcp_server( prisma_client: PrismaClient, server_id: str, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b6760e58852..95d00554034 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1246,6 +1246,67 @@ class MCPServerManager: """Return server IDs that bypass per-key restrictions.""" return [server.server_id for server in self.get_registry().values() if server.allow_all_keys is True] + @staticmethod + def get_byom_submitted_servers_cache_key(user_id: str) -> str: + return f"byom_submitted_servers:{user_id}" + + async def invalidate_byom_submitted_servers_cache(self, user_id: str | None) -> None: + if not user_id: + return + try: + from litellm.proxy.proxy_server import user_api_key_cache + + await user_api_key_cache.async_delete_cache(key=self.get_byom_submitted_servers_cache_key(user_id)) + except Exception as e: # noqa: BLE001 + verbose_logger.warning(f"Failed to invalidate BYOM submitted MCP server cache: {str(e)}") + + async def _get_active_submitted_mcp_server_ids_for_user( + self, user_api_key_auth: UserAPIKeyAuth | None + ) -> list[str]: + submitter_user_id = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + if not submitter_user_id: + return [] + + try: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + get_active_submitted_mcp_server_ids_for_user, + ) + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + except Exception as e: # noqa: BLE001 + verbose_logger.warning(f"Failed to load BYOM submitted MCP server cache dependencies: {str(e)}") + return [] + + byom_cache_key = self.get_byom_submitted_servers_cache_key(submitter_user_id) + submitted_server_ids: list[str] | None = None + try: + cached_submitted_server_ids = await user_api_key_cache.async_get_cache(key=byom_cache_key) + if cached_submitted_server_ids is not None: + submitted_server_ids = cast(list[str], cached_submitted_server_ids) + except Exception as e: # noqa: BLE001 + verbose_logger.warning(f"Failed to read BYOM submitted MCP server cache: {str(e)}") + + if submitted_server_ids is None: + if prisma_client is None: + submitted_server_ids = [] + else: + try: + submitted_server_ids = await get_active_submitted_mcp_server_ids_for_user( + prisma_client, submitter_user_id + ) + except Exception as e: # noqa: BLE001 + verbose_logger.warning(f"Failed to read BYOM submitted MCP servers from database: {str(e)}") + submitted_server_ids = [] + try: + await user_api_key_cache.async_set_cache( + key=byom_cache_key, + value=submitted_server_ids, + ttl=60, + ) + except Exception as e: # noqa: BLE001 + verbose_logger.warning(f"Failed to write BYOM submitted MCP server cache: {str(e)}") + + return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None] + async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> List[str]: """ Get the allowed MCP Servers for the user. @@ -1259,25 +1320,30 @@ class MCPServerManager: allow_all_server_ids = self.get_allow_all_keys_server_ids() + # The key explicitly opted out of every MCP server. Return zero before + # layering on allow_all_keys or submitted servers so the opt-out is absolute. + key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None + if key_object_permission is not None and ( + SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []) + ): + return [] + + # Check if object_permission.mcp_servers is explicitly set (not None, empty list is valid) + has_explicit_object_permission = key_object_permission is not None and ( + key_object_permission.mcp_servers is not None + ) + if has_explicit_object_permission: + verbose_logger.debug(f"Object permission mcp_servers explicitly set: {key_object_permission.mcp_servers}") + + # BYOM creator visibility never widens a key that was explicitly scoped: + # only keys without their own mcp_servers list get submitted servers unioned in. + submitted_server_ids = ( + [] + if has_explicit_object_permission + else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth) + ) + try: - # The key explicitly opted out of every MCP server. Return zero before - # layering on allow_all_keys servers so the opt-out is absolute. - key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None - if key_object_permission is not None and ( - SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []) - ): - return [] - - # Check if object_permission.mcp_servers is explicitly set - has_explicit_object_permission = False - if user_api_key_auth and user_api_key_auth.object_permission: - # Check if mcp_servers is explicitly set (not None, empty list is valid) - if user_api_key_auth.object_permission.mcp_servers is not None: - has_explicit_object_permission = True - verbose_logger.debug( - f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}" - ) - # If admin but NO explicit object permission, get all servers if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission: verbose_logger.debug("Admin user without explicit object_permission - returning all servers") @@ -1299,6 +1365,7 @@ class MCPServerManager: in_toolset_scope = _mcp_active_toolset_id.get() is not None if not in_toolset_scope: combined_servers.update(allow_all_server_ids) + combined_servers.update(submitted_server_ids) # For anonymous callers (no user_id, no role), also surface any # servers the operator has opted into upstream-delegated auth. @@ -1331,9 +1398,9 @@ class MCPServerManager: except Exception: # noqa: BLE001 verbose_logger.exception( "Failed to get allowed MCP servers; team-level object_permission " - "grants may be dropped. Falling back to global servers only." + "grants may be dropped. Falling back to global and submitted servers." ) - return allow_all_server_ids + return list(dict.fromkeys(allow_all_server_ids + submitted_server_ids)) async def resolve_toolset_tool_permissions( self, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index ab9d04a4eb4..09d9809e4fb 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1158,6 +1158,7 @@ if MCP_AVAILABLE: server_id, touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, ) + await global_mcp_server_manager.invalidate_byom_submitted_servers_cache(approved.submitted_by) await global_mcp_server_manager.reload_servers_from_database() return _redact_mcp_credentials(approved) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 0fc303737b4..3611a4a1401 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -1160,6 +1160,12 @@ async def update_mcp_semantic_filter_settings( Update MCP semantic filter settings in database. Settings will be picked up by all pods within approximately 10 seconds via background polling. """ + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Only proxy admins can update MCP semantic filter settings.", + ) + result = await _update_litellm_setting( settings=settings, settings_key="mcp_semantic_tool_filter", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index abefb2fd984..2b7c29ceff3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6498,3 +6498,86 @@ class TestMCPMetaTraceCarrier: assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None only_progress = RequestParams.Meta.model_validate({"progressToken": "p1"}) assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_includes_active_servers_submitted_by_user(): + """BYOM submitters can see approved servers they submitted without allow_all_keys.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + submitted_server = _make_mcp_server_for_scope_filter("submitted-1", "user_mcp") + submitter = UserAPIKeyAuth( + user_id="submitter-user", + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-submitter", + ) + other_user = UserAPIKeyAuth( + user_id="other-user", + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-other", + ) + + async def _submitted_ids(prisma_client, user_id): + return ["submitted-1"] if user_id == "submitter-user" else [] + + with ( + patch.object( + global_mcp_server_manager, + "get_registry", + return_value={"submitted-1": submitted_server}, + ), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(return_value=[]), + ), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.proxy._experimental.mcp_server.db.get_active_submitted_mcp_server_ids_for_user", + side_effect=_submitted_ids, + ), + ): + submitter_allowed = await global_mcp_server_manager.get_allowed_mcp_servers(submitter) + other_allowed = await global_mcp_server_manager.get_allowed_mcp_servers(other_user) + + assert "submitted-1" in submitter_allowed + assert "submitted-1" not in other_allowed + + +@pytest.mark.asyncio +async def test_get_active_submitted_mcp_server_ids_for_user_queries_active_rows(): + from litellm.proxy._experimental.mcp_server.db import ( + get_active_submitted_mcp_server_ids_for_user, + ) + from litellm.proxy._types import MCPApprovalStatus + + row = MagicMock() + row.server_id = "submitted-1" + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + + result = await get_active_submitted_mcp_server_ids_for_user(prisma_client, "submitter-user") + + assert result == ["submitted-1"] + prisma_client.db.litellm_mcpservertable.find_many.assert_awaited_once_with( + where={ + "submitted_by": "submitter-user", + "approval_status": MCPApprovalStatus.active, + }, + ) + + +@pytest.mark.asyncio +async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_db(): + from litellm.proxy._experimental.mcp_server.db import ( + get_active_submitted_mcp_server_ids_for_user, + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.find_many = AsyncMock() + + assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == [] + prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5dec580c771..c258ed1035f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -3198,6 +3198,250 @@ class TestMCPServerManager: assert result == [] mock_inner.assert_not_called() + @pytest.mark.asyncio + async def test_no_mcp_servers_sentinel_excludes_submitted_byom_servers(self): + from litellm.proxy import proxy_server as proxy_server_module + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth + + class _Cache: + async def async_get_cache(self, key: str): + return ["submitted-server"] + + manager = MCPServerManager() + manager.registry = { + "submitted-server": MCPServer( + server_id="submitted-server", + name="submitted", + transport=MCPTransport.http, + ) + } + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="perm_no_mcp", + mcp_servers=["no-mcp-servers"], + mcp_access_groups=[], + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="sk-test", + user_id="user-123", + object_permission=object_permission, + object_permission_id="perm_no_mcp", + ) + + with ( + patch.object(proxy_server_module, "user_api_key_cache", _Cache()), + patch.object(proxy_server_module, "prisma_client", None), + patch.object( + manager, "get_allow_all_keys_server_ids", return_value=["global-server"] + ), + patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=["leaked-server"], + ) as mock_inner, + ): + result = await manager.get_allowed_mcp_servers(user_api_key_auth) + + assert result == [] + mock_inner.assert_not_called() + + @pytest.mark.asyncio + async def test_explicitly_scoped_key_excludes_submitted_byom_servers(self): + from litellm.proxy import proxy_server as proxy_server_module + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth + + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=["submitted-server"]) + + manager = MCPServerManager() + manager.registry = { + "submitted-server": MCPServer( + server_id="submitted-server", + name="submitted", + transport=MCPTransport.http, + ), + "scoped-server": MCPServer( + server_id="scoped-server", + name="scoped", + transport=MCPTransport.http, + ), + } + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="perm_scoped", + mcp_servers=["scoped-server"], + mcp_access_groups=[], + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="sk-test", + user_id="user-123", + object_permission=object_permission, + object_permission_id="perm_scoped", + ) + + with ( + patch.object(proxy_server_module, "user_api_key_cache", cache), + patch.object(proxy_server_module, "prisma_client", None), + patch.object(manager, "get_allow_all_keys_server_ids", return_value=[]), + patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=["scoped-server"], + ), + ): + result = await manager.get_allowed_mcp_servers(user_api_key_auth) + + assert result == ["scoped-server"] + cache.async_get_cache.assert_not_awaited() + + @pytest.mark.asyncio + async def test_toolset_scope_excludes_submitted_byom_servers(self): + from litellm.proxy import proxy_server as proxy_server_module + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._experimental.mcp_server.mcp_context import ( + _mcp_active_toolset_id, + ) + from litellm.proxy._types import UserAPIKeyAuth + + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=["submitted-server"]) + + manager = MCPServerManager() + manager.registry = { + "submitted-server": MCPServer( + server_id="submitted-server", + name="submitted", + transport=MCPTransport.http, + ), + "toolset-server": MCPServer( + server_id="toolset-server", + name="toolset", + transport=MCPTransport.http, + ), + } + user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123") + + token = _mcp_active_toolset_id.set("toolset-abc") + try: + with ( + patch.object(proxy_server_module, "user_api_key_cache", cache), + patch.object(proxy_server_module, "prisma_client", None), + patch.object( + manager, "get_allow_all_keys_server_ids", return_value=["global-server"] + ), + patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=["toolset-server"], + ), + ): + result = await manager.get_allowed_mcp_servers(user_api_key_auth) + finally: + _mcp_active_toolset_id.reset(token) + + assert result == ["toolset-server"] + + @pytest.mark.asyncio + async def test_invalidate_byom_submitted_servers_cache_deletes_key(self): + from litellm.proxy import proxy_server as proxy_server_module + + cache = MagicMock() + cache.async_delete_cache = AsyncMock() + manager = MCPServerManager() + + with patch.object(proxy_server_module, "user_api_key_cache", cache): + await manager.invalidate_byom_submitted_servers_cache("user-123") + await manager.invalidate_byom_submitted_servers_cache(None) + + cache.async_delete_cache.assert_awaited_once_with(key="byom_submitted_servers:user-123") + + @pytest.mark.asyncio + async def test_get_active_submitted_ids_cache_miss_queries_db_and_caches(self): + from litellm.proxy import proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=None) + cache.async_set_cache = AsyncMock() + manager = MCPServerManager() + manager.registry = { + "submitted-server": MCPServer( + server_id="submitted-server", + name="submitted", + transport=MCPTransport.http, + ) + } + user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123") + + with ( + patch.object(proxy_server_module, "user_api_key_cache", cache), + patch.object(proxy_server_module, "prisma_client", MagicMock()), + patch( + "litellm.proxy._experimental.mcp_server.db.get_active_submitted_mcp_server_ids_for_user", + AsyncMock(return_value=["submitted-server", "unknown-server"]), + ), + ): + result = await manager._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth) + + assert result == ["submitted-server"] + cache.async_set_cache.assert_awaited_once_with( + key="byom_submitted_servers:user-123", + value=["submitted-server", "unknown-server"], + ttl=60, + ) + + @pytest.mark.asyncio + async def test_get_allowed_mcp_servers_fallback_keeps_submitted_byom_servers(self): + from litellm.proxy import proxy_server as proxy_server_module + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._types import UserAPIKeyAuth + + class _Cache: + async def async_get_cache(self, key: str): + assert key == "byom_submitted_servers:user-123" + return ["submitted-server"] + + manager = MCPServerManager() + manager.registry = { + "submitted-server": MCPServer( + server_id="submitted-server", + name="submitted", + transport=MCPTransport.http, + ) + } + user_api_key_auth = UserAPIKeyAuth( + api_key="sk-test", + user_id="user-123", + ) + + with ( + patch.object(proxy_server_module, "user_api_key_cache", _Cache()), + patch.object(proxy_server_module, "prisma_client", None), + patch.object( + manager, "get_allow_all_keys_server_ids", return_value=["global-server"] + ), + patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + side_effect=RuntimeError("permission resolver failed"), + ), + ): + result = await manager.get_allowed_mcp_servers(user_api_key_auth) + + assert set(result) == {"global-server", "submitted-server"} + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_anonymous_delegate_requires_oauth2(self): """Anonymous delegated auth listing should only include oauth2 servers.""" 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 f40904e234d..ef8433fdaf0 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 @@ -3223,8 +3223,10 @@ class TestMCPApprovalWorkflow: pending_server.approval_status = MCPApprovalStatus.pending_review approved_server = generate_mock_mcp_server_db_record() approved_server.approval_status = MCPApprovalStatus.active + approved_server.submitted_by = "submitter-user" mock_manager = MagicMock() + mock_manager.invalidate_byom_submitted_servers_cache = AsyncMock() mock_manager.reload_servers_from_database = AsyncMock() with ( @@ -3250,6 +3252,9 @@ class TestMCPApprovalWorkflow: ) mock_manager.reload_servers_from_database.assert_awaited_once() + mock_manager.invalidate_byom_submitted_servers_cache.assert_awaited_once_with( + "submitter-user" + ) assert result is not None @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index cb77c42fe9b..69845ec59c2 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -2367,3 +2367,29 @@ def test_update_ui_settings_writes_audit_log(monkeypatch): assert after["disable_custom_api_keys"] is True finally: app.dependency_overrides.pop(user_api_key_auth, None) + + +def test_update_mcp_semantic_filter_settings_requires_proxy_admin(monkeypatch): + """Non-admin callers must not mutate global MCP semantic filter settings.""" + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + + async def _internal_user_auth(): + return UserAPIKeyAuth( + user_id="internal-user-1", + api_key="hashed-internal-key", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + app.dependency_overrides[user_api_key_auth] = _internal_user_auth + try: + resp = client.patch( + "/update/mcp_semantic_filter_settings", + json={"enabled": True, "top_k": 99, "similarity_threshold": 0.01}, + ) + assert resp.status_code == 403 + assert "proxy admin" in resp.json()["detail"].lower() + finally: + app.dependency_overrides.pop(user_api_key_auth, None) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx index fc0a0779b24..100373a4571 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx @@ -97,7 +97,9 @@ function ConfirmDialog({ action, serverName, isCurrentlyActive, onConfirm, onCan

Are you sure you want to {action} "{serverName}"?{" "} - {isApprove ? "This will make it active and available for use." : rejectBody} + {isApprove + ? "This will activate the server. The submitting user will see it in their MCP Servers list once approved." + : rejectBody}

{!isApprove && (