mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(mcp): route progress to the originating request stream
This commit is contained in:
parent
cc93368e2c
commit
5e36654b32
2 changed files with 13 additions and 12 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue