fix(mcp): avoid relying on progress notification root shape
This commit is contained in:
@@ -47,12 +47,30 @@ _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":
|
||||
payload = _mcp_jsonrpc_payload(message)
|
||||
if _payload_value(payload, "method") != "notifications/progress":
|
||||
return False
|
||||
|
||||
params = getattr(root, "params", None)
|
||||
return not isinstance(params, Mapping) or "progressToken" not in params
|
||||
params = _payload_value(payload, "params")
|
||||
return not _progress_params_have_token(params)
|
||||
|
||||
|
||||
def _mcp_jsonrpc_payload(message: Any) -> Any:
|
||||
"""Return the JSON-RPC payload across current and future MCP SDK shapes."""
|
||||
envelope = getattr(message, "message", message)
|
||||
return getattr(envelope, "root", None) or envelope
|
||||
|
||||
|
||||
def _payload_value(payload: Any, key: str) -> Any:
|
||||
if isinstance(payload, Mapping):
|
||||
return payload.get(key)
|
||||
return getattr(payload, key, None)
|
||||
|
||||
|
||||
def _progress_params_have_token(params: Any) -> bool:
|
||||
if isinstance(params, Mapping):
|
||||
return "progressToken" in params
|
||||
return hasattr(params, "progressToken") or hasattr(params, "progress_token")
|
||||
|
||||
|
||||
class _MalformedProgressNotificationFilter:
|
||||
|
||||
@@ -36,6 +36,24 @@ def _mcp_notification(method: str, params: dict[str, Any] | None = None) -> Sess
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_progress_detection_accepts_flattened_sdk_message_shape():
|
||||
malformed = SimpleNamespace(
|
||||
message=SimpleNamespace(
|
||||
method="notifications/progress",
|
||||
params={"progress": 20, "total": 600},
|
||||
)
|
||||
)
|
||||
valid = SimpleNamespace(
|
||||
message=SimpleNamespace(
|
||||
method="notifications/progress",
|
||||
params={"progressToken": "req-1", "progress": 25},
|
||||
)
|
||||
)
|
||||
|
||||
assert mcp_runtime._is_malformed_mcp_progress_notification(malformed) is True
|
||||
assert mcp_runtime._is_malformed_mcp_progress_notification(valid) is False
|
||||
|
||||
|
||||
class _FakeMcpTool(Tool):
|
||||
def __init__(self, name: str) -> None:
|
||||
self._name = name
|
||||
|
||||
Reference in New Issue
Block a user