Skip to content
Merged
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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ dependencies = [
"pywin32>=311; sys_platform == 'win32' and python_version >= '3.14'",
"pywin32>=310; sys_platform == 'win32' and python_version < '3.14'",
"pyjwt[crypto]>=2.10.1",
"cryptography>=49,<50; python_version < '3.11'",
"typing-extensions>=4.9.0",
"typing-inspection>=0.4.1",
]
Expand Down
12 changes: 10 additions & 2 deletions src/mcp/server/streamable_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
import re
from abc import ABC, abstractmethod
from collections.abc import AsyncGenerator, Awaitable, Callable
from contextlib import asynccontextmanager
from contextlib import asynccontextmanager, suppress
from dataclasses import dataclass
from functools import partial
from http import HTTPStatus
Expand All @@ -23,7 +23,7 @@
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
from pydantic import ValidationError
from sse_starlette import EventSourceResponse
from starlette.requests import Request
from starlette.requests import ClientDisconnect, Request
from starlette.responses import Response
from starlette.types import Receive, Scope, Send

Expand Down Expand Up @@ -688,6 +688,14 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re
finally:
await sse_stream_reader.aclose()

except ClientDisconnect:
# The client is already gone; send a response for ASGI middleware
# that expects one, while tolerating a closed socket.
logger.warning("Client disconnected during POST request")
response = self._create_json_response(None, 499) # type: ignore[arg-type]
with suppress(Exception):
await response(scope, receive, send)
return
except Exception as err:
logger.exception("Error handling POST request")
response = self._create_error_response(
Expand Down
26 changes: 26 additions & 0 deletions src/mcp/server/streamable_http_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import logging
import math
from collections.abc import AsyncIterator
from http import HTTPStatus
from typing import Any, Final
from uuid import uuid4

Expand Down Expand Up @@ -198,6 +199,31 @@ async def handle_request(
receive: ASGI receive function
send: ASGI send function
"""
# Stateless GET/DELETE cannot establish a stream or terminate a session.
# Reject them before the body limit middleware tries to read a body.
if (
self.stateless
and self._task_group is not None
and scope["type"] == "http"
and scope["method"] in ("GET", "DELETE")
):
method = scope["method"]
error = JSONRPCError(
jsonrpc="2.0",
id="",
error=ErrorData(
code=INVALID_REQUEST,
message=f"Method Not Allowed: {method} is not supported in stateless mode",
),
)
response = Response(
content=error.model_dump_json(by_alias=True, exclude_unset=True),
status_code=HTTPStatus.METHOD_NOT_ALLOWED,
headers={"Allow": "POST"},
media_type="application/json",
)
await response(scope, receive, send)
return
await self.asgi_app(scope, receive, send)

async def _handle_request(
Expand Down
92 changes: 86 additions & 6 deletions tests/server/test_streamable_http_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -351,6 +351,73 @@ async def mock_receive():
assert len(transport._request_streams) == 0, "Transport should have no active request streams"


@pytest.mark.anyio
@pytest.mark.parametrize("method", ["GET", "DELETE"])
async def test_stateless_http_rejects_session_methods(method: str) -> None:
"""Stateless HTTP has no standalone SSE stream or session to terminate."""
manager = StreamableHTTPSessionManager(app=Server("test-stateless-methods"), stateless=True)
sent: list[Message] = []

async def receive() -> Message:
raise AssertionError("GET and DELETE must not read a request body") # pragma: no cover

async def send(message: Message) -> None:
sent.append(message)

scope: Scope = {
"type": "http",
"method": method,
"path": "/mcp",
"headers": [
(b"accept", b"application/json, text/event-stream"),
(b"content-length", b"999999999"),
],
}
async with manager.run():
with anyio.fail_after(5):
await manager.handle_request(scope, receive, send)

response = next(message for message in sent if message["type"] == "http.response.start")
assert response["status"] == 405
assert (b"allow", b"POST") in response["headers"]
body = b"".join(message.get("body", b"") for message in sent if message["type"] == "http.response.body")
error = json.loads(body)
assert error["jsonrpc"] == "2.0"
assert error["id"] == ""
assert error["error"]["code"] == INVALID_REQUEST
assert error["error"]["message"] == f"Method Not Allowed: {method} is not supported in stateless mode"


@pytest.mark.anyio
async def test_disconnected_stateless_post_returns_client_closed_request(caplog: pytest.LogCaptureFixture) -> None:
"""A client leaving before its POST body arrives is not a server error."""
manager = StreamableHTTPSessionManager(app=Server("test-disconnected-post"), stateless=True)
sent: list[Message] = []

async def receive() -> Message:
return {"type": "http.disconnect"}

async def send(message: Message) -> None:
sent.append(message)

scope: Scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"content-type", b"application/json"),
(b"accept", b"application/json, text/event-stream"),
],
}
async with manager.run():
with anyio.fail_after(5):
await manager.handle_request(scope, receive, send)

response = next(message for message in sent if message["type"] == "http.response.start")
assert response["status"] == 499
assert not any(record.levelno >= logging.ERROR for record in caplog.records)


@pytest.mark.anyio
async def test_stateless_request_cancels_server_task_when_request_ends() -> None:
"""A completed stateless request must not leave its server task running."""
Expand Down Expand Up @@ -391,37 +458,50 @@ async def send(message: Message) -> None:

@pytest.mark.anyio
async def test_stateless_stream_closes_cleanly_on_shutdown() -> None:
"""Shutdown should send the final body chunk for an active SSE stream."""
manager = StreamableHTTPSessionManager(app=Server("test-stateless-shutdown"), stateless=True)
"""Shutdown should close an active stateless POST response stream."""
app = Server("test-stateless-shutdown")
tool_started = anyio.Event()

@app.call_tool()
async def handle_call_tool(name: str, arguments: dict[str, Any]) -> list[TextContent]:
tool_started.set()
await anyio.sleep_forever()
raise AssertionError("unreachable") # pragma: no cover

manager = StreamableHTTPSessionManager(app=app, stateless=True)
stream_started = anyio.Event()
stream_closed = anyio.Event()
request_received = False
body = json.dumps(
{"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "slow", "arguments": {}}}
).encode()

async def receive() -> Message:
nonlocal request_received
if not request_received:
request_received = True
return {"type": "http.request", "body": b"", "more_body": False}
return {"type": "http.request", "body": body, "more_body": False}
await anyio.sleep_forever()
raise AssertionError("unreachable") # pragma: no cover

async def send(message: Message) -> None:
if message["type"] == "http.response.start":
if message["type"] == "http.response.start" and message["status"] == 200:
stream_started.set()
if message["type"] == "http.response.body" and not message.get("more_body", False):
stream_closed.set()

scope: Scope = {
"type": "http",
"method": "GET",
"method": "POST",
"path": "/mcp",
"headers": [(b"accept", b"text/event-stream")],
"headers": [(b"content-type", b"application/json"), (b"accept", b"application/json, text/event-stream")],
}
with anyio.fail_after(5):
async with anyio.create_task_group() as requests:
async with manager.run():
requests.start_soon(manager.handle_request, scope, receive, send)
await stream_started.wait()
await tool_started.wait()
await stream_closed.wait()


Expand Down
Loading
Loading