fix(mcp): route progress to the originating request stream

This commit is contained in:
Joshua Valluru 2026-09-26 12:22:42 -07:00
parent cc93368e2c
commit 5e36654b32
2 changed files with 13 additions and 12 deletions

View file

@ -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)

View file

@ -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: