Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 13 additions & 9 deletions src/mcp/client/streamable_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
26 changes: 26 additions & 0 deletions tests/client/test_streamable_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading