From 228162ff468774c034cd50a2f954a85681adbbdf Mon Sep 17 00:00:00 2001 From: daleselaji-dev Date: Thu, 17 Sep 2026 10:45:31 +0800 Subject: [PATCH] fix(client): specialize streamable HTTP request headers --- src/mcp/client/streamable_http.py | 22 +++++++++++++--------- tests/client/test_streamable_http.py | 26 ++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 9 deletions(-) diff --git a/src/mcp/client/streamable_http.py b/src/mcp/client/streamable_http.py index 82de50fd05..ac3e8874d2 100644 --- a/src/mcp/client/streamable_http.py +++ b/src/mcp/client/streamable_http.py @@ -131,7 +131,12 @@ def __init__(self, url: str) -> None: # `_consume_modern_cancellation`. Keys are verbatim-typed ("1" is not 1). self._in_flight_posts: dict[RequestId, _InFlightPost] = {} - def _prepare_headers(self) -> dict[str, str]: + def _prepare_headers( + self, + *, + accept: str = "application/json, text/event-stream", + content_type: str | None = "application/json", + ) -> dict[str, str]: """Build MCP-specific request headers for any outbound HTTP request. These are merged with the ``httpx2.AsyncClient`` defaults (these take @@ -141,10 +146,9 @@ def _prepare_headers(self) -> dict[str, str]: GET/DELETE — still carry the negotiated version. Per-message headers are layered on top by the caller. """ - headers: dict[str, str] = { - "accept": "application/json, text/event-stream", - "content-type": "application/json", - } + headers: dict[str, str] = {"accept": accept} + if content_type is not None: + headers["content-type"] = content_type if self.session_id: headers[MCP_SESSION_ID] = self.session_id if self._protocol_version_header: @@ -227,7 +231,7 @@ async def handle_get_stream(self, client: httpx2.AsyncClient, read_stream_writer if not self.session_id: return - headers = self._prepare_headers() + headers = self._prepare_headers(accept="text/event-stream", content_type=None) if last_event_id: headers[LAST_EVENT_ID] = last_event_id @@ -267,7 +271,7 @@ async def handle_get_stream(self, client: httpx2.AsyncClient, read_stream_writer async def _handle_resumption_request(self, ctx: RequestContext) -> None: """Handle a resumption request using GET with SSE.""" - headers = self._prepare_headers() + headers = self._prepare_headers(accept="text/event-stream", content_type=None) if ctx.metadata and ctx.metadata.resumption_token: headers[LAST_EVENT_ID] = ctx.metadata.resumption_token else: @@ -538,7 +542,7 @@ async def _handle_reconnection( delay_ms = retry_interval_ms if retry_interval_ms is not None else DEFAULT_RECONNECTION_DELAY_MS await anyio.sleep(delay_ms / 1000.0) - headers = self._prepare_headers() + headers = self._prepare_headers(accept="text/event-stream", content_type=None) headers[LAST_EVENT_ID] = last_event_id try: @@ -666,7 +670,7 @@ async def terminate_session(self, client: httpx2.AsyncClient) -> None: return # pragma: no cover try: - headers = self._prepare_headers() + headers = self._prepare_headers(content_type=None) response = await request_within_origin(client, "DELETE", self.url, headers=headers) if response.status_code == 405: diff --git a/tests/client/test_streamable_http.py b/tests/client/test_streamable_http.py index c6e62ad94a..80bf4c4977 100644 --- a/tests/client/test_streamable_http.py +++ b/tests/client/test_streamable_http.py @@ -87,6 +87,32 @@ def test_mcp_name_header_values_are_base64_wrapped_when_unsafe_for_an_http_field assert encoded == raw +def test_transport_headers_match_the_outbound_request_kind() -> None: + """POST, SSE GET, and bodyless DELETE advertise only what they send.""" + transport = StreamableHTTPTransport("http://test/mcp") + transport.session_id = "session-1" + transport._protocol_version_header = LATEST_MODERN_VERSION # pyright: ignore[reportPrivateUsage] + + post = transport._prepare_headers() # pyright: ignore[reportPrivateUsage] + sse_get = transport._prepare_headers( # pyright: ignore[reportPrivateUsage] + accept="text/event-stream", content_type=None + ) + delete = transport._prepare_headers(content_type=None) # pyright: ignore[reportPrivateUsage] + + assert post["accept"] == "application/json, text/event-stream" + assert post["content-type"] == "application/json" + assert sse_get == { + "accept": "text/event-stream", + "mcp-session-id": "session-1", + "mcp-protocol-version": LATEST_MODERN_VERSION, + } + assert delete == { + "accept": "application/json, text/event-stream", + "mcp-session-id": "session-1", + "mcp-protocol-version": LATEST_MODERN_VERSION, + } + + @pytest.mark.anyio async def test_post_request_merges_per_message_metadata_headers() -> None: """`ClientMessageMetadata.headers` on a `SessionMessage` are merged into the outgoing POST headers