diff --git a/tests/tools/test_mcp_client_cert.py b/tests/tools/test_mcp_client_cert.py index 9a16f46c32..6b0ce2af89 100644 --- a/tests/tools/test_mcp_client_cert.py +++ b/tests/tools/test_mcp_client_cert.py @@ -6,10 +6,12 @@ Covers: errors, missing-file errors. 2. HTTP (new SDK ``streamable_http_client``) path forwards ``cert=`` into the - user-owned ``httpx.AsyncClient``. + inner ``AsyncHTTPTransport`` wrapped by the wire-body-cap transport that + the user-owned ``httpx.AsyncClient`` is built on. 3. SSE path forwards ``cert`` and ``ssl_verify`` via an ``httpx_client_factory`` - without breaking the OAuth/headers/timeout passthrough. + (always injected — it also installs the wire-body cap) without breaking the + OAuth/headers/timeout passthrough. """ from __future__ import annotations @@ -34,6 +36,33 @@ def _patch_sdk_async_client(dummy): return patch.object(sdk_httpx(), "AsyncClient", dummy) +def _patch_sdk_transport(dummy): + """Patch ``AsyncHTTPTransport`` on the SDK's httpx module. + + Since the wire-body cap (_make_mcp_body_cap_transport), verify/cert are + applied to an inner ``AsyncHTTPTransport`` rather than as AsyncClient + kwargs — a custom ``transport=`` makes client-level TLS kwargs inert. + """ + from tools.mcp_tool import sdk_httpx + + return patch.object(sdk_httpx(), "AsyncHTTPTransport", dummy) + + +class _DummyTransport: + """Capture-only stand-in for AsyncHTTPTransport.""" + + captured: dict = {} + + def __init__(self, **kwargs): + type(self).captured = dict(kwargs) + + async def handle_async_request(self, request): # pragma: no cover + raise AssertionError("not dispatched in these tests") + + async def aclose(self): # pragma: no cover + pass + + # --------------------------------------------------------------------------- # _resolve_client_cert helper # --------------------------------------------------------------------------- @@ -92,7 +121,7 @@ class TestResolveClientCert: class TestHTTPClientCert: def test_cert_forwarded_to_async_client(self, tmp_path): """When client_cert is set, the new-SDK HTTP path passes ``cert=`` - into ``httpx.AsyncClient``.""" + into the inner AsyncHTTPTransport under the body-cap wrapper.""" from tools.mcp_tool import MCPServerTask cert = tmp_path / "client.pem" @@ -138,6 +167,7 @@ class TestHTTPClientCert: with patch("tools.mcp_tool._MCP_HTTP_AVAILABLE", True), \ patch("tools.mcp_tool._MCP_NEW_HTTP", True), \ _patch_sdk_async_client(DummyAsyncClient), \ + _patch_sdk_transport(_DummyTransport), \ patch("tools.mcp_tool.streamable_http_client", return_value=DummyTransportCtx()), \ patch("tools.mcp_tool.ClientSession", DummySession), \ @@ -148,7 +178,13 @@ class TestHTTPClientCert: }) asyncio.run(_drive()) - assert captured.get("cert") == str(cert) + # cert/verify live on the inner transport (client-level TLS kwargs + # are inert once a custom transport is passed); the client itself + # receives the body-cap transport. + assert _DummyTransport.captured.get("cert") == str(cert) + assert _DummyTransport.captured.get("verify") is True + assert "cert" not in captured + assert "transport" in captured def test_missing_cert_file_surfaces_clear_error(self, tmp_path): @@ -217,9 +253,10 @@ def patch_sse_client(): class TestSSEClientCert: - def test_no_factory_when_defaults(self, patch_sse_client): - """With no cert and ssl_verify=True (default), the SDK's own factory is - used — we don't inject one.""" + def test_factory_always_injected_with_default_tls(self, patch_sse_client): + """The factory is always injected (it installs the wire-body cap); + with no cert and ssl_verify=True the inner transport keeps default + TLS settings.""" from tools.mcp_tool import MCPServerTask server = MCPServerTask("sse-test") @@ -242,7 +279,22 @@ class TestSSEClientCert: pass asyncio.run(drive()) - assert "httpx_client_factory" not in patch_sse_client + factory = patch_sse_client.get("httpx_client_factory") + assert factory is not None, "body-cap factory must always be injected" + + captured_client_kwargs: dict = {} + + class DummyAsyncClient: + def __init__(self, **kwargs): + captured_client_kwargs.update(kwargs) + + with _patch_sdk_async_client(DummyAsyncClient), \ + _patch_sdk_transport(_DummyTransport): + factory(headers=None, timeout=None, auth=None) + + assert _DummyTransport.captured.get("verify") is True + assert "cert" not in _DummyTransport.captured + assert "transport" in captured_client_kwargs def test_factory_injected_when_cert_set(self, patch_sse_client, tmp_path): """With client_cert set, an httpx_client_factory is injected that @@ -286,13 +338,15 @@ class TestSSEClientCert: captured_client_kwargs.update(kwargs) from tools.mcp_tool import sdk_httpx - with _patch_sdk_async_client(DummyAsyncClient): + with _patch_sdk_async_client(DummyAsyncClient), \ + _patch_sdk_transport(_DummyTransport): factory(headers={"x": "y"}, timeout=sdk_httpx().Timeout(30.0), auth=None) - assert captured_client_kwargs["cert"] == str(cert) - assert captured_client_kwargs["verify"] is True + assert _DummyTransport.captured["cert"] == str(cert) + assert _DummyTransport.captured["verify"] is True assert captured_client_kwargs["follow_redirects"] is True assert captured_client_kwargs["headers"] == {"x": "y"} + assert "transport" in captured_client_kwargs def test_factory_forwards_custom_ca_bundle(self, patch_sse_client, tmp_path): """ssl_verify as a path is forwarded to the factory's httpx client.""" @@ -332,8 +386,9 @@ class TestSSEClientCert: def __init__(self, **kwargs): captured_client_kwargs.update(kwargs) - with _patch_sdk_async_client(DummyAsyncClient): + with _patch_sdk_async_client(DummyAsyncClient), \ + _patch_sdk_transport(_DummyTransport): factory(headers=None, timeout=None, auth=None) - assert captured_client_kwargs["verify"] == str(ca_bundle) - assert "cert" not in captured_client_kwargs + assert _DummyTransport.captured["verify"] == str(ca_bundle) + assert "cert" not in _DummyTransport.captured diff --git a/tests/tools/test_mcp_http_body_cap.py b/tests/tools/test_mcp_http_body_cap.py new file mode 100644 index 0000000000..74aeed142e --- /dev/null +++ b/tests/tools/test_mcp_http_body_cap.py @@ -0,0 +1,132 @@ +"""Wire-body cap for MCP HTTP/SSE transports (port of openclaw/openclaw#123194). + +Exercises ``_make_mcp_body_cap_transport`` with a real httpx AsyncClient over +a MockTransport: oversized finite bodies and oversized SSE events fail with a +ReadError naming the byte cap; bodies/events under the cap pass; long-lived +SSE connections reset accounting at event boundaries so cumulative keepalive +traffic is unlimited. +""" + +import httpx +import pytest + +from tools.mcp_tool_errors import _MCP_HTTP_MAX_BODY_BYTES, _make_mcp_body_cap_transport + +LIMIT = 1024 # small cap for tests + + +def _client_for(handler, limit=LIMIT): + inner = httpx.MockTransport(handler) + capped = _make_mcp_body_cap_transport(httpx, inner, limit=limit) + return httpx.AsyncClient(transport=capped) + + +@pytest.mark.asyncio +async def test_small_json_body_passes(): + async def handler(request): + return httpx.Response(200, json={"ok": True}) + async with _client_for(handler) as client: + resp = await client.get("http://mcp.test/rpc") + assert resp.json() == {"ok": True} + + +@pytest.mark.asyncio +async def test_oversized_body_rejected_via_content_length(): + body = b"x" * (LIMIT + 1) + async def handler(request): + return httpx.Response(200, content=body) + async with _client_for(handler) as client: + with pytest.raises(httpx.ReadError, match=r"Content-Length"): + await client.get("http://mcp.test/rpc") + + +@pytest.mark.asyncio +async def test_oversized_streamed_body_rejected_without_content_length(): + # A streaming body with no Content-Length must still trip the cap. + async def gen(): + for _ in range(8): + yield b"y" * (LIMIT // 4) + + class _Stream(httpx.AsyncByteStream): + async def __aiter__(self): + async for c in gen(): + yield c + + async def handler(request): + return httpx.Response(200, stream=_Stream()) + async with _client_for(handler) as client: + with pytest.raises(httpx.ReadError, match=r"HTTP response exceeds"): + await client.get("http://mcp.test/rpc") + + +@pytest.mark.asyncio +async def test_sse_event_over_cap_rejected(): + async def gen(): + yield b"data: " + b"z" * (LIMIT + 64) # one giant unterminated event + + class _Stream(httpx.AsyncByteStream): + async def __aiter__(self): + async for c in gen(): + yield c + + async def handler(request): + return httpx.Response( + 200, stream=_Stream(), + headers={"content-type": "text/event-stream"}, + ) + async with _client_for(handler) as client: + with pytest.raises(httpx.ReadError, match=r"SSE event exceeds"): + async with client.stream("GET", "http://mcp.test/sse") as resp: + async for _ in resp.aiter_bytes(): + pass + + +@pytest.mark.asyncio +async def test_sse_cumulative_keepalives_unlimited(): + # Many small completed events whose TOTAL far exceeds the cap must all + # pass: accounting resets at every completed event boundary. + async def gen(): + for i in range(64): + yield b": keepalive %d\n\n" % i + b"data: {\"n\": %d}\n\n" % i + + class _Stream(httpx.AsyncByteStream): + async def __aiter__(self): + async for c in gen(): + yield c + + async def handler(request): + return httpx.Response( + 200, stream=_Stream(), + headers={"content-type": "text/event-stream"}, + ) + total = 0 + async with _client_for(handler, limit=64) as client: + async with client.stream("GET", "http://mcp.test/sse") as resp: + async for chunk in resp.aiter_bytes(): + total += len(chunk) + assert total > 64 # cumulative traffic exceeded the per-event cap + + +@pytest.mark.asyncio +async def test_sse_event_split_across_chunks_counts_prefix(): + # An event streamed in pieces (no boundary) accumulates until it + # crosses the cap. + async def gen(): + for _ in range(6): + yield b"data: " + b"q" * (LIMIT // 4) + + class _Stream(httpx.AsyncByteStream): + async def __aiter__(self): + async for c in gen(): + yield c + + async def handler(request): + return httpx.Response( + 200, stream=_Stream(), + headers={"content-type": "text/event-stream"}, + ) + async with _client_for(handler) as client: + with pytest.raises(httpx.ReadError, match=r"SSE event exceeds"): + async with client.stream("GET", "http://mcp.test/sse") as resp: + async for _ in resp.aiter_bytes(): + pass diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 7eecfaaf1d..13167317e5 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -3,6 +3,7 @@ headers, redirect header stripping, exception-group unwrapping, auth/session-exp method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py.""" import asyncio +import contextlib import errno import importlib import logging @@ -231,6 +232,73 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False, return _strip_on_cross_origin_redirect +# Wire-body cap, applied at the httpx transport before the SDK buffers/JSON-parses a response. A +# hostile or misbehaving remote MCP server can stream an unbounded catalog/tool-result body and none +# of the post-parse limits (resource cap, tool-result truncation) run before the parse blows up. +# Finite HTTP bodies are capped at this many bytes (a larger Content-Length is rejected up front); +# each SSE *event* is capped, with the counter reset at completed event boundaries so a long-lived +# stream and its keepalives have no cumulative limit. Violations raise the SDK httpx's ReadError and +# flow through the ordinary transport teardown/reconnect path (#66092). +_MCP_HTTP_MAX_BODY_BYTES = 10 * 1024 * 1024 +_SSE_EVENT_BOUNDARIES = (b"\n\n", b"\r\n\r\n") + + +def _make_mcp_body_cap_transport(httpx_mod, inner_transport, limit: int = _MCP_HTTP_MAX_BODY_BYTES): + """Wrap ``inner_transport`` so every response body is size-capped. ``httpx_mod`` must be the SDK's + own httpx module (``sdk_httpx()``): the transport is handed to that SDK's ``AsyncClient``.""" + + class _CappedStream(httpx_mod.AsyncByteStream): + def __init__(self, inner, is_sse: bool, url: str): + self._inner, self._is_sse, self._url = inner, is_sse, url + + def _reject(self, kind: str): + return httpx_mod.ReadError(f"MCP {kind} exceeds {limit} bytes (from {self._url})") + + async def __aiter__(self): + counted = 0 + async for chunk in self._inner: + if self._is_sse: + # Bytes up to the last completed event boundary belong to finished events (they must + # still fit the per-event cap together with the carried prefix); the remainder starts + # the next event's budget. + boundary_end = max(chunk.rfind(sep) + len(sep) if sep in chunk else -1 for sep in _SSE_EVENT_BOUNDARIES) + if boundary_end != -1: + if counted + boundary_end > limit: + raise self._reject("SSE event") + counted = len(chunk) - boundary_end + else: + counted += len(chunk) + else: + counted += len(chunk) + if counted > limit: + raise self._reject("SSE event" if self._is_sse else "HTTP response") + yield chunk + + async def aclose(self): + await self._inner.aclose() + + class _BodyCapTransport(httpx_mod.AsyncBaseTransport): + def __init__(self, inner): + self._inner = inner + + async def handle_async_request(self, request): + response = await self._inner.handle_async_request(request) + declared = response.headers.get("content-length") + with contextlib.suppress(ValueError): # malformed header: the streamed cap still applies + if declared is not None and int(declared) > limit: + await response.aclose() + raise httpx_mod.ReadError(f"MCP HTTP response declares Content-Length {declared} > {limit} " + f"bytes cap (from {request.url})") + is_sse = "text/event-stream" in response.headers.get("content-type", "").lower() + response.stream = _CappedStream(response.stream, is_sse, str(request.url)) + return response + + async def aclose(self): + await self._inner.aclose() + + return _BodyCapTransport(inner_transport) + + def _exc_children(exc: BaseException) -> List[BaseException]: """Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions.""" nested = getattr(exc, "exceptions", None) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 017e1fe2eb..d6d6751cb0 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -7,7 +7,7 @@ import asyncio import os from contextlib import asynccontextmanager from typing import Dict, Optional, Set -from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _is_streamable_http_rejection, _make_redirect_header_stripper, _resolve_client_cert, _unwrap_exception_group +from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _is_streamable_http_rejection, _make_mcp_body_cap_transport, _make_redirect_header_stripper, _resolve_client_cert, _unwrap_exception_group from tools.mcp_tool_lifecycle import _filter_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids from tools.mcp_tool_common import _core from tools import mcp_tool_config as _config @@ -350,14 +350,16 @@ class MCPServerTransportMixin: # Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded or OAuth SSE 401s silently. sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout), "sse_read_timeout": 300.0, **_present(auth=oauth_auth)} - if client_cert is not None or ssl_verify is not True: - # sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's (headers, - # auth, timeout) and layers TLS on top. Client MUST come from the SDK's httpx (httpx2 on mcp >= 2.0). - _httpx_mod = _core.sdk_httpx() - sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient( - follow_redirects=True, verify=ssl_verify, - timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0), - **_present(headers=headers, auth=auth, cert=client_cert)) + # Always own the client: the httpx_client_factory forwards the SDK's (headers, auth, timeout), + # installs the wire-body cap and layers TLS on the inner transport (client-level verify/cert are + # inert once a custom transport= is passed). Client MUST come from the SDK's httpx (httpx2 on mcp >= 2.0). + _httpx_mod = _core.sdk_httpx() + sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient( + follow_redirects=True, + timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0), + transport=_make_mcp_body_cap_transport( + _httpx_mod, _httpx_mod.AsyncHTTPTransport(verify=ssl_verify, **_present(cert=client_cert))), + **_present(headers=headers, auth=auth)) return _core.sse_client(**sse_kwargs) def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float, @@ -377,10 +379,13 @@ class MCPServerTransportMixin: httpx = _core.sdk_httpx() _strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper( httpx.URL(url), strict=strict_cfg_headers, configured_header_names=configured_header_names) + # verify/cert live on the inner transport: a custom transport= makes client-level TLS kwargs inert. client_kwargs: dict = {"follow_redirects": True, "timeout": httpx.Timeout(float(connect_timeout), read=300.0), - "verify": ssl_verify, **({"headers": headers} if headers else {}), + **({"headers": headers} if headers else {}), "event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]}, - **_present(auth=oauth_auth, cert=client_cert)} + "transport": _make_mcp_body_cap_transport( + httpx, httpx.AsyncHTTPTransport(verify=ssl_verify, **_present(cert=client_cert))), + **_present(auth=oauth_auth)} @asynccontextmanager async def _owned_client_streams(): # the SDK skips cleanup when http_client is provided