diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 405815e6162..bbc40433516 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -863,7 +863,9 @@ async def _selected_connections_refusal( global_mcp_server_manager, ) - for server in (global_mcp_server_manager.get_mcp_server_by_name(name) for name in dict.fromkeys(selected_servers)): + for server in ( + global_mcp_server_manager.get_mcp_server_by_id(server_id) for server_id in dict.fromkeys(selected_servers) + ): if server is None or not await lookup_server_reachability(flow.user_id, server.server_id): return _oauth_error(400, "invalid_request", "a selected MCP server is no longer available") if ( diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 0c037ebf893..eb3d8855951 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1258,7 +1258,7 @@ async def _collect_mcp_listing( return [], classify_list_exception(exc) results: Final = await asyncio.gather(*(fetch_one(server) for server in servers)) - failure: Final = listing_auth_error({server.name: result[1] for server, result in zip(servers, results)}) + failure: Final = listing_auth_error({server.server_id: result[1] for server, result in zip(servers, results)}) if failure is not None: raise failure return list(chain.from_iterable(items for items, _ in results)) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c6bd22d1729..bcd1f670298 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -127,8 +127,7 @@ class MCPAuthResponse: and message.get("more_body", False) and not any(line.startswith(b"data:") and line[5:].strip() for line in body.splitlines()) ): - if not body.startswith(b":"): - self._preamble = (*self._preamble, message) + self._preamble = (*self._preamble, message) return self._committed = True if self.challenge is not None: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 7093e8f58c2..9a752ab14a4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -77,7 +77,7 @@ def _salt_key(monkeypatch): from unittest.mock import patch with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_id", return_value=_scoped_mcp_server("public", auth_type="none"), ): yield @@ -346,7 +346,7 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u handle, cookies = _flow_cookie_from(authorize_response) denied = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -356,7 +356,7 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u assert denied.status_code == 403 anonymous = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -366,7 +366,7 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u assert anonymous.status_code == 401 completed = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -442,7 +442,7 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u @pytest.mark.asyncio async def test_complete_rejects_missing_tampered_and_expired_flows(): missing = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", method="POST"), flow_handle="nope", @@ -452,7 +452,7 @@ async def test_complete_rejects_missing_tampered_and_expired_flows(): assert missing.status_code == 400 tampered = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"), flow_handle="h1", @@ -535,7 +535,7 @@ async def test_token_gates_on_live_user_revalidation(failure, expected_status, e authorize_response = _authorize(client_id, session_user_id="deactivated-user") handle, cookies = _flow_cookie_from(authorize_response) completed = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -591,7 +591,7 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete(): handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1")) first = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -600,7 +600,7 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete(): ) assert first.status_code == 303 second = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -743,7 +743,7 @@ async def _complete(redirect_uri: str, delivery, cookies=None, handle=None, sess if cookies is None: handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=redirect_uri)) response = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -852,7 +852,7 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed(): handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=LOOPBACK_REDIRECT_URI)) rejected = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -864,7 +864,7 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed(): assert json.loads(rejected.body)["error"] == "invalid_request" retried = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -1026,8 +1026,9 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No handle, cookies = _flow_cookie_from(response) with patch(_MANAGER_PATCH) as manager: - manager.get_mcp_server_by_id.return_value = scoped_server - manager.get_mcp_server_by_name.return_value = scoped_server or _scoped_mcp_server("public", auth_type="none") + manager.get_mcp_server_by_id.side_effect = lambda server_id: ( + scoped_server or _scoped_mcp_server("public", auth_type="none") + ) if server_id == "public-id" else scoped_server return await complete_connect_flow( request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -1035,7 +1036,7 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No cache=cache or DualCache(), lookup_vendor_credential=vendor or _VendorCredential(), lookup_server_reachability=reachable or _ServerReachability(), - **{"selected_servers": ("public",), **overrides}, + **{"selected_servers": ("public-id",), **overrides}, ) @@ -1844,7 +1845,7 @@ async def test_mcp_wire_formats_carry_no_native_client_fields(): assert "audience" not in flow_wire assert "team_id" not in flow_wire completed = await complete_connect_flow( - selected_servers=("public",), + selected_servers=("public-id",), lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -2415,14 +2416,14 @@ async def test_unified_completion_checks_selected_upstream_and_permissions(crede vendor=vendor, reachable=_ServerReachability(reachable), cache=cache, - selected_servers=("github",), + selected_servers=("github-id",), ) assert completed.status_code == status if not reachable: assert vendor.calls == [] if status != 303: assert "location" not in completed.headers - retried = await _complete_page(response, scoped_server=server, cache=cache, selected_servers=("github",)) + retried = await _complete_page(response, scoped_server=server, cache=cache, selected_servers=("github-id",)) assert retried.status_code == 303 else: code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0] @@ -2475,12 +2476,34 @@ async def test_unified_completion_checks_every_selected_server_and_preserves_can ) assert missing.status_code == 400 unfinished = await complete_connect_flow( - **arguments, selected_servers=("github", "slack"), lookup_vendor_credential=credential + **arguments, selected_servers=("github-id", "slack-id"), lookup_vendor_credential=credential ) assert unfinished.status_code == 400 with pytest.raises(asyncio.CancelledError): - await complete_connect_flow(**arguments, selected_servers=("github",), lookup_vendor_credential=cancelled) + await complete_connect_flow(**arguments, selected_servers=("github-id",), lookup_vendor_credential=cancelled) completed = await complete_connect_flow( - **arguments, selected_servers=("github", "slack"), lookup_vendor_credential=_VendorCredential() + **arguments, selected_servers=("github-id", "slack-id"), lookup_vendor_credential=_VendorCredential() ) assert completed.status_code == 303 + + +@pytest.mark.asyncio +async def test_unified_completion_validates_selected_id_despite_alias_collision(): + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") + handle, cookies = _flow_cookie_from(response) + selected = _scoped_mcp_server("github", oauth2_flow="authorization_code") + other = _scoped_mcp_server("other", auth_type="none").model_copy(update={"alias": selected.server_id}) + cache = DualCache() + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = other + manager.get_mcp_server_by_id.side_effect = {selected.server_id: selected, other.server_id: other}.get + vendor = _VendorCredential("absent") + arguments = dict(request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", cache=cache, selected_servers=(selected.server_id,), lookup_server_reachability=_ServerReachability()) + refused = await complete_connect_flow(**arguments, lookup_vendor_credential=vendor) + assert refused.status_code == 400 + assert vendor.calls == [("u1", selected.server_id)] + completed = await complete_connect_flow(**arguments, lookup_vendor_credential=_VendorCredential()) + assert completed.status_code == 303 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 72af43d73f7..c625b80e1fb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -798,13 +798,14 @@ async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_fa @pytest.mark.asyncio -async def test_transport_preserves_sse_priming_event_on_success() -> None: +@pytest.mark.parametrize("preamble", (b"id: resume-token\r\ndata: \r\n\r\n", b": ping\r\n\r\n")) +async def test_transport_preserves_sse_priming_event_on_success(preamble: bytes) -> None: from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse send: Final = AsyncMock() response: Final = MCPAuthResponse(send) start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} - priming: Final = {"type": "http.response.body", "body": b"id: resume-token\r\ndata: \r\n\r\n", "more_body": True} + priming: Final = {"type": "http.response.body", "body": preamble, "more_body": True} tools: Final = { "type": "http.response.body", "body": b'data: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\n\n', @@ -947,3 +948,20 @@ async def test_optional_listing_challenges_auth_when_other_server_times_out( assert caught.value.www_authenticate == "Bearer" assert caught.value.server_name == "blocked" assert create.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates")) +@pytest.mark.parametrize("healthy_first", (False, True)) +async def test_optional_listing_preserves_healthy_duplicate_names(kind: str, healthy_first: bool, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import operations + healthy: Final = _http_server("healthy-id", "duplicate") + blocked: Final = _http_server("blocked-id", "duplicate") + manager: Final = MagicMock() + fetch: Final = AsyncMock(side_effect=[[], MCPUpstreamAuthError(401, "Bearer", "duplicate")] if healthy_first else [MCPUpstreamAuthError(401, "Bearer", "duplicate"), []]) + setattr(manager, f"get_{kind}_from_server", fetch) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))) + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[healthy, blocked] if healthy_first else [blocked, healthy])) + assert await getattr(operations, f"_list_mcp_{kind}")() == [] + assert fetch.await_count == 2 diff --git a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx index c9609405676..92abd1e420b 100644 --- a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.test.tsx @@ -125,7 +125,7 @@ describe("MCPAppsPanel connected-app reachability (LIT-4861)", () => { vi.mocked(fetchMCPServers).mockResolvedValue(connectServers); vi.mocked(listMCPTools).mockResolvedValue({ tools: [] }); - renderConnectPanel(true, ["reachable_srv", "unreachable_srv"]); + renderConnectPanel(true, ["s-reach", "s-unreach"]); expect(await screen.findByText("reachable_srv")).toBeInTheDocument(); expect(vi.mocked(fetchMCPServers)).toHaveBeenCalledWith("tok", undefined, true); @@ -300,3 +300,26 @@ describe("MCPAppsPanel connected-app reachability (LIT-4861)", () => { expect(screen.queryByText("revoked_srv")).not.toBeInTheDocument(); }); }); + +it.each([true, false])("selects an unambiguous upstream in connect mode=%s", async (connectMode) => { + const onChange = vi.fn(); + vi.mocked(fetchMCPServers).mockResolvedValue([ + { + server_id: "selected-id", + server_name: "github", + alias: "github-selected", + auth_type: "none", + connected_app_reachable: true, + }, + { server_id: "other-id", server_name: "other", alias: "github", auth_type: "none", connected_app_reachable: true }, + ] as MCPServer[]); + vi.mocked(listMCPTools).mockResolvedValue({ tools: [] }); + render( + + + , + ); + fireEvent.click(await screen.findByText("github")); + fireEvent.click(await screen.findByRole("button", { name: "Connect", exact: true })); + await waitFor(() => expect(onChange).toHaveBeenCalledWith([connectMode ? "selected-id" : "github"])); +}); diff --git a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx index 1fee8923e94..d249d441f88 100644 --- a/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx +++ b/ui/litellm-dashboard/src/components/chat/MCPAppsPanel.tsx @@ -140,6 +140,11 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, const nameOf = (s: MCPServer) => s.server_name ?? s.alias ?? s.server_id; + const selectionOf = useCallback( + (server: MCPServer) => (connectMode ? server.server_id : nameOf(server)), + [connectMode], + ); + const detailServer = servers.find((s) => s.server_id === detailServerId); const connectUnavailabilityLabel = useCallback( @@ -239,19 +244,20 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, .filter( (s) => oauthConnected.has(s.server_id) && - !selectedServersRef.current.includes(nameOf(s)) && + !selectedServersRef.current.includes(selectionOf(s)) && connectUnavailabilityLabel(s) === null, ) - .map(nameOf); + .map(selectionOf); if (namesToAdd.length > 0) { onChangeRef.current([...selectedServersRef.current, ...namesToAdd]); } - }, [oauthConnected, connectUnavailabilityLabel]); + }, [oauthConnected, connectUnavailabilityLabel, selectionOf]); const handleToggle = async (server: MCPServer, checked: boolean) => { const serverName = nameOf(server); + const selection = selectionOf(server); if (!checked) { - onChange(selectedServers.filter((s) => s !== serverName)); + onChange(selectedServers.filter((s) => s !== selection)); setOauthConnected((prev) => { const next = new Set(prev); next.delete(server.server_id); @@ -268,8 +274,8 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, return; } if (connectableNow(server.server_id) === undefined) return; - if (!selectedServersRef.current.includes(serverName)) { - onChange([...selectedServersRef.current, serverName]); + if (!selectedServersRef.current.includes(selection)) { + onChange([...selectedServersRef.current, selection]); } } catch { toast.warning(`Could not load tools for ${serverName}`); @@ -308,7 +314,7 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, /> ); } - if (selectedServers.includes(nameOf(server))) { + if (selectedServers.includes(selectionOf(server))) { return ; } return null; @@ -328,12 +334,12 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, name.toLowerCase().includes(query.toLowerCase()) || (s.description ?? "").toLowerCase().includes(query.toLowerCase()); const matchesTab = - activeTab === "all" || (selectedServers.includes(name) && connectUnavailabilityLabel(s) === null); + activeTab === "all" || (selectedServers.includes(selectionOf(s)) && connectUnavailabilityLabel(s) === null); return matchesQuery && matchesTab; }); const connectedCount = servers.filter( - (s) => selectedServers.includes(nameOf(s)) && connectUnavailabilityLabel(s) === null, + (s) => selectedServers.includes(selectionOf(s)) && connectUnavailabilityLabel(s) === null, ).length; const emptyStateText = () => { @@ -348,7 +354,7 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, if (detailServer) { const name = nameOf(detailServer); - const isConnected = selectedServers.includes(name); + const isConnected = selectedServers.includes(selectionOf(detailServer)); const isTogglingOn = togglingOn.has(name); const color = getAvatarColor(name); @@ -388,7 +394,7 @@ const MCPAppsPanel: React.FC = ({ accessToken, selectedServers, onChange, n.delete(detailServer.server_id); return n; }); - onChangeRef.current(selectedServersRef.current.filter((s) => s !== name)); + onChangeRef.current(selectedServersRef.current.filter((s) => s !== selectionOf(detailServer))); }} className="font-semibold h-[38px] min-w-[110px]" >