From 5e36654b32b38b18f75be0ea4ede7ec80c194342 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 12:22:42 -0700 Subject: [PATCH] fix(mcp): route progress to the originating request stream --- .../proxy/_experimental/mcp_server/server.py | 6 +----- .../mcp_server/test_mcp_tool_search.py | 19 ++++++++++++------- 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 555aebc7434..9b41472be6e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -870,11 +870,7 @@ if MCP_AVAILABLE: async def forward_progress(progress: float, total: float | None): """Forward progress notifications from external MCP to Host""" try: - await host_session.send_progress_notification( - progress_token=host_token, - progress=progress, - total=total, - ) + await host_session.report_progress(progress=progress, total=total) verbose_logger.debug("Forwarded progress %s/%s to Host", progress, total) except Exception as e: verbose_logger.error("Failed to forward progress to Host: %s", e) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 4575741aa8b..940cc817813 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -13,7 +13,7 @@ Covers: import json from collections.abc import Sequence from types import SimpleNamespace -from typing import Any +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -1159,7 +1159,11 @@ class TestCaptureHostProgressCallback: @pytest.mark.asyncio @pytest.mark.parametrize("token", ["tok12345", 12345, 0]) - async def test_forwards_wire_progress_token(self, _mcp_request_ctx, token) -> None: + @pytest.mark.parametrize("delivery_error", [None, RuntimeError("request stream closed")]) + async def test_forwards_progress_on_originating_request(self, _mcp_request_ctx, token, delivery_error) -> None: + from mcp.server.connection import Connection + from mcp.server.session import ServerSession + from mcp.shared.dispatcher import DispatchContext from mcp.types import CallToolRequestParams from litellm.proxy._experimental.mcp_server.server import _capture_host_progress_callback @@ -1167,13 +1171,14 @@ class TestCaptureHostProgressCallback: params = CallToolRequestParams.model_validate( {"name": "tool", "_meta": {"progressToken": token}}, by_name=False ) - session = AsyncMock() - callback = _capture_host_progress_callback(_mcp_request_ctx(meta=params.meta, session=session)) + request_channel: Final = AsyncMock(spec=DispatchContext, progress=AsyncMock(side_effect=delivery_error)) + connection: Final = MagicMock(spec=Connection, protocol_version="2025-11-25", outbound=AsyncMock()) + session: Final = ServerSession(request_channel, connection, request_meta=params.meta) + callback: Final = _capture_host_progress_callback(_mcp_request_ctx(meta=params.meta, session=session)) assert callback is not None await callback(0.5, 1.0) - session.send_progress_notification.assert_awaited_once_with( - progress_token=token, progress=0.5, total=1.0 - ) + request_channel.progress.assert_awaited_once_with(0.5, 1.0, None) + connection.outbound.notify.assert_not_awaited() class TestHandleListToolsVirtual: