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:
|
def _is_malformed_mcp_progress_notification(message: Any) -> bool:
|
||||||
root = getattr(getattr(message, "message", None), "root", None)
|
payload = _mcp_jsonrpc_payload(message)
|
||||||
if getattr(root, "method", None) != "notifications/progress":
|
if _payload_value(payload, "method") != "notifications/progress":
|
||||||
return False
|
return False
|
||||||
|
|
||||||
params = getattr(root, "params", None)
|
params = _payload_value(payload, "params")
|
||||||
return not isinstance(params, Mapping) or "progressToken" not in 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:
|
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):
|
class _FakeMcpTool(Tool):
|
||||||
def __init__(self, name: str) -> None:
|
def __init__(self, name: str) -> None:
|
||||||
self._name = name
|
self._name = name
|
||||||
|
|||||||
Reference in New Issue
Block a user