mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(mcp): preserve upstream identity through OAuth completion
This commit is contained in:
parent
812e58a146
commit
e9aaae31f2
7 changed files with 110 additions and 39 deletions
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false } } })}>
|
||||
<MCPAppsPanel accessToken="tok" selectedServers={[]} onChange={onChange} connectMode={connectMode} />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
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"]));
|
||||
});
|
||||
|
|
|
|||
|
|
@ -140,6 +140,11 @@ const MCPAppsPanel: React.FC<Props> = ({ 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<Props> = ({ 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<Props> = ({ 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<Props> = ({ accessToken, selectedServers, onChange,
|
|||
/>
|
||||
);
|
||||
}
|
||||
if (selectedServers.includes(nameOf(server))) {
|
||||
if (selectedServers.includes(selectionOf(server))) {
|
||||
return <span className="w-[7px] h-[7px] rounded-full bg-success shrink-0" />;
|
||||
}
|
||||
return null;
|
||||
|
|
@ -328,12 +334,12 @@ const MCPAppsPanel: React.FC<Props> = ({ 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<Props> = ({ 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<Props> = ({ 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]"
|
||||
>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue