fix(mcp): ignore malformed progress notifications
This commit is contained in:
@@ -46,6 +46,58 @@ _RELOAD_LOCKS: WeakKeyDictionary[Any, asyncio.Lock] = WeakKeyDictionary()
|
|||||||
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
||||||
|
|
||||||
|
|
||||||
|
def _is_malformed_mcp_progress_notification(message: Any) -> bool:
|
||||||
|
root = getattr(getattr(message, "message", None), "root", None)
|
||||||
|
if getattr(root, "method", None) != "notifications/progress":
|
||||||
|
return False
|
||||||
|
|
||||||
|
params = getattr(root, "params", None)
|
||||||
|
return not isinstance(params, Mapping) or "progressToken" not in params
|
||||||
|
|
||||||
|
|
||||||
|
class _MalformedProgressNotificationFilter:
|
||||||
|
def __init__(self, read_stream: Any, server_name: str) -> None:
|
||||||
|
self._read_stream = read_stream
|
||||||
|
self._server_name = server_name
|
||||||
|
self._iterator: Any | None = None
|
||||||
|
|
||||||
|
async def __aenter__(self) -> "_MalformedProgressNotificationFilter":
|
||||||
|
await self._read_stream.__aenter__()
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> Any:
|
||||||
|
return await self._read_stream.__aexit__(exc_type, exc, tb)
|
||||||
|
|
||||||
|
def __aiter__(self) -> "_MalformedProgressNotificationFilter":
|
||||||
|
self._iterator = self._read_stream.__aiter__()
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __anext__(self) -> Any:
|
||||||
|
if self._iterator is None:
|
||||||
|
self._iterator = self._read_stream.__aiter__()
|
||||||
|
|
||||||
|
while True:
|
||||||
|
message = await self._iterator.__anext__()
|
||||||
|
if _is_malformed_mcp_progress_notification(message):
|
||||||
|
logger.debug(
|
||||||
|
"MCP server '{}': dropped progress notification without progressToken",
|
||||||
|
self._server_name,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
return message
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
close = getattr(self._read_stream, "aclose", None)
|
||||||
|
if close is not None:
|
||||||
|
await close()
|
||||||
|
|
||||||
|
|
||||||
|
def _filter_malformed_mcp_progress_notifications(read_stream: Any, server_name: str) -> Any:
|
||||||
|
if not all(hasattr(read_stream, name) for name in ("__aenter__", "__aexit__", "__aiter__")):
|
||||||
|
return read_stream
|
||||||
|
return _MalformedProgressNotificationFilter(read_stream, server_name)
|
||||||
|
|
||||||
|
|
||||||
def _sanitize_name(name: str) -> str:
|
def _sanitize_name(name: str) -> str:
|
||||||
"""Sanitize an MCP-derived name for model API compatibility."""
|
"""Sanitize an MCP-derived name for model API compatibility."""
|
||||||
return _SANITIZE_RE.sub("_", re.sub(r"[^a-zA-Z0-9_-]", "_", name))
|
return _SANITIZE_RE.sub("_", re.sub(r"[^a-zA-Z0-9_-]", "_", name))
|
||||||
@@ -681,6 +733,7 @@ async def connect_mcp_servers(
|
|||||||
await server_stack.aclose()
|
await server_stack.aclose()
|
||||||
return name, None
|
return name, None
|
||||||
|
|
||||||
|
read = _filter_malformed_mcp_progress_notifications(read, name)
|
||||||
session = await server_stack.enter_async_context(ClientSession(read, write))
|
session = await server_stack.enter_async_context(ClientSession(read, write))
|
||||||
await session.initialize()
|
await session.initialize()
|
||||||
|
|
||||||
|
|||||||
@@ -8,9 +8,11 @@ from types import SimpleNamespace
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import anyio
|
||||||
import pytest
|
import pytest
|
||||||
from mcp import types as mcp_types
|
from mcp import types as mcp_types
|
||||||
from mcp.shared.exceptions import McpError
|
from mcp.shared.exceptions import McpError
|
||||||
|
from mcp.shared.message import SessionMessage
|
||||||
from mcp.types import ErrorData
|
from mcp.types import ErrorData
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
@@ -22,6 +24,18 @@ from nanobot.config.loader import load_config, save_config
|
|||||||
from nanobot.config.schema import MCPServerConfig
|
from nanobot.config.schema import MCPServerConfig
|
||||||
|
|
||||||
|
|
||||||
|
def _mcp_notification(method: str, params: dict[str, Any] | None = None) -> SessionMessage:
|
||||||
|
return SessionMessage(
|
||||||
|
message=mcp_types.JSONRPCMessage(
|
||||||
|
mcp_types.JSONRPCNotification(
|
||||||
|
jsonrpc="2.0",
|
||||||
|
method=method,
|
||||||
|
params=params,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class _FakeMcpTool(Tool):
|
class _FakeMcpTool(Tool):
|
||||||
def __init__(self, name: str) -> None:
|
def __init__(self, name: str) -> None:
|
||||||
self._name = name
|
self._name = name
|
||||||
@@ -56,6 +70,33 @@ def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mcp_read_filter_drops_progress_notifications_without_progress_token():
|
||||||
|
send, receive = anyio.create_memory_object_stream(4)
|
||||||
|
malformed_progress = _mcp_notification(
|
||||||
|
"notifications/progress",
|
||||||
|
{"progress": 20, "total": 600, "message": "Polling"},
|
||||||
|
)
|
||||||
|
tool_change = _mcp_notification("notifications/tools/list_changed")
|
||||||
|
valid_progress = _mcp_notification(
|
||||||
|
"notifications/progress",
|
||||||
|
{"progressToken": "req-1", "progress": 25, "total": 600, "message": "Polling"},
|
||||||
|
)
|
||||||
|
|
||||||
|
await send.send(malformed_progress)
|
||||||
|
await send.send(tool_change)
|
||||||
|
await send.send(valid_progress)
|
||||||
|
await send.aclose()
|
||||||
|
|
||||||
|
wrapped = mcp_runtime._filter_malformed_mcp_progress_notifications(receive, "brightdata")
|
||||||
|
forwarded = []
|
||||||
|
async with wrapped:
|
||||||
|
async for message in wrapped:
|
||||||
|
forwarded.append(message)
|
||||||
|
|
||||||
|
assert forwarded == [tool_change, valid_progress]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
|
|||||||
Reference in New Issue
Block a user