Add SSE polling support (SEP-1699) (#1654)

This commit is contained in:
Felix Weinberger
2025-12-02 11:44:49 +00:00
committed by GitHub
parent 2cd178a962
commit 281fd4765e
21 changed files with 1463 additions and 47 deletions
@@ -0,0 +1,30 @@
# MCP SSE Polling Demo Client
Demonstrates client-side auto-reconnect for the SSE polling pattern (SEP-1699).
## Features
- Connects to SSE polling demo server
- Automatically reconnects when server closes SSE stream
- Resumes from Last-Event-ID to avoid missing messages
- Respects server-provided retry interval
## Usage
```bash
# First start the server:
uv run mcp-sse-polling-demo --port 3000
# Then run this client:
uv run mcp-sse-polling-client --url http://localhost:3000/mcp
# Custom options:
uv run mcp-sse-polling-client --url http://localhost:3000/mcp --items 20 --checkpoint-every 5
```
## Options
- `--url`: Server URL (default: <http://localhost:3000/mcp>)
- `--items`: Number of items to process (default: 10)
- `--checkpoint-every`: Checkpoint interval (default: 3)
- `--log-level`: Logging level (default: DEBUG)
@@ -0,0 +1 @@
"""SSE Polling Demo Client - demonstrates auto-reconnect for long-running tasks."""
@@ -0,0 +1,105 @@
"""
SSE Polling Demo Client
Demonstrates the client-side auto-reconnect for SSE polling pattern.
This client connects to the SSE Polling Demo server and calls process_batch,
which triggers periodic server-side stream closes. The client automatically
reconnects using Last-Event-ID and resumes receiving messages.
Run with:
# First start the server:
uv run mcp-sse-polling-demo --port 3000
# Then run this client:
uv run mcp-sse-polling-client --url http://localhost:3000/mcp
"""
import asyncio
import logging
import click
from mcp import ClientSession
from mcp.client.streamable_http import streamablehttp_client
logger = logging.getLogger(__name__)
async def run_demo(url: str, items: int, checkpoint_every: int) -> None:
"""Run the SSE polling demo."""
print(f"\n{'=' * 60}")
print("SSE Polling Demo Client")
print(f"{'=' * 60}")
print(f"Server URL: {url}")
print(f"Processing {items} items with checkpoints every {checkpoint_every}")
print(f"{'=' * 60}\n")
async with streamablehttp_client(url) as (read_stream, write_stream, _):
async with ClientSession(read_stream, write_stream) as session:
# Initialize the connection
print("Initializing connection...")
await session.initialize()
print("Connected!\n")
# List available tools
tools = await session.list_tools()
print(f"Available tools: {[t.name for t in tools.tools]}\n")
# Call the process_batch tool
print(f"Calling process_batch(items={items}, checkpoint_every={checkpoint_every})...\n")
print("-" * 40)
result = await session.call_tool(
"process_batch",
{
"items": items,
"checkpoint_every": checkpoint_every,
},
)
print("-" * 40)
if result.content:
content = result.content[0]
text = getattr(content, "text", str(content))
print(f"\nResult: {text}")
else:
print("\nResult: No content")
print(f"{'=' * 60}\n")
@click.command()
@click.option(
"--url",
default="http://localhost:3000/mcp",
help="Server URL",
)
@click.option(
"--items",
default=10,
help="Number of items to process",
)
@click.option(
"--checkpoint-every",
default=3,
help="Checkpoint interval",
)
@click.option(
"--log-level",
default="INFO",
help="Logging level",
)
def main(url: str, items: int, checkpoint_every: int, log_level: str) -> None:
"""Run the SSE Polling Demo client."""
logging.basicConfig(
level=getattr(logging, log_level.upper()),
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
# Suppress noisy HTTP client logging
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
asyncio.run(run_demo(url, items, checkpoint_every))
if __name__ == "__main__":
main()
@@ -0,0 +1,36 @@
[project]
name = "mcp-sse-polling-client"
version = "0.1.0"
description = "Demo client for SSE polling with auto-reconnect"
readme = "README.md"
requires-python = ">=3.10"
authors = [{ name = "Anthropic, PBC." }]
keywords = ["mcp", "sse", "polling", "client"]
license = { text = "MIT" }
dependencies = ["click>=8.2.0", "mcp"]
[project.scripts]
mcp-sse-polling-client = "mcp_sse_polling_client.main:main"
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["mcp_sse_polling_client"]
[tool.pyright]
include = ["mcp_sse_polling_client"]
venvPath = "."
venv = ".venv"
[tool.ruff.lint]
select = ["E", "F", "I"]
ignore = []
[tool.ruff]
line-length = 120
target-version = "py310"
[dependency-groups]
dev = ["pyright>=1.1.378", "pytest>=8.3.3", "ruff>=0.6.9"]
@@ -14,6 +14,7 @@ import click
from mcp.server.fastmcp import Context, FastMCP
from mcp.server.fastmcp.prompts.base import UserMessage
from mcp.server.session import ServerSession
from mcp.server.streamable_http import EventCallback, EventMessage, EventStore
from mcp.types import (
AudioContent,
Completion,
@@ -21,6 +22,7 @@ from mcp.types import (
CompletionContext,
EmbeddedResource,
ImageContent,
JSONRPCMessage,
PromptReference,
ResourceTemplateReference,
SamplingMessage,
@@ -31,6 +33,43 @@ from pydantic import AnyUrl, BaseModel, Field
logger = logging.getLogger(__name__)
# Type aliases for event store
StreamId = str
EventId = str
class InMemoryEventStore(EventStore):
"""Simple in-memory event store for SSE resumability testing."""
def __init__(self) -> None:
self._events: list[tuple[StreamId, EventId, JSONRPCMessage | None]] = []
self._event_id_counter = 0
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId:
"""Store an event and return its ID."""
self._event_id_counter += 1
event_id = str(self._event_id_counter)
self._events.append((stream_id, event_id, message))
return event_id
async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None:
"""Replay events after the specified ID."""
target_stream_id = None
for stream_id, event_id, _ in self._events:
if event_id == last_event_id:
target_stream_id = stream_id
break
if target_stream_id is None:
return None
last_event_id_int = int(last_event_id)
for stream_id, event_id, message in self._events:
if stream_id == target_stream_id and int(event_id) > last_event_id_int:
# Skip priming events (None message)
if message is not None:
await send_callback(EventMessage(message, event_id))
return target_stream_id
# Test data
TEST_IMAGE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg=="
TEST_AUDIO_BASE64 = "UklGRiYAAABXQVZFZm10IBAAAAABAAEAQB8AAAB9AAACABAAZGF0YQIAAAA="
@@ -39,8 +78,13 @@ TEST_AUDIO_BASE64 = "UklGRiYAAABXQVZFZm10IBAAAAABAAEAQB8AAAB9AAACABAAZGF0YQIAAAA
resource_subscriptions: set[str] = set()
watched_resource_content = "Watched resource content"
# Create event store for SSE resumability (SEP-1699)
event_store = InMemoryEventStore()
mcp = FastMCP(
name="mcp-conformance-test-server",
event_store=event_store,
retry_interval=100, # 100ms retry interval for SSE polling
)
@@ -263,6 +307,19 @@ def test_error_handling() -> str:
raise RuntimeError("This tool intentionally returns an error for testing")
@mcp.tool()
async def test_reconnection(ctx: Context[ServerSession, None]) -> str:
"""Tests SSE polling by closing stream mid-call (SEP-1699)"""
await ctx.info("Before disconnect")
await ctx.close_sse_stream()
await asyncio.sleep(0.2) # Wait for client to reconnect
await ctx.info("After reconnect")
return "Reconnection test completed"
# Resources
@mcp.resource("test://static-text")
def static_text_resource() -> str:
@@ -24,7 +24,7 @@ class EventEntry:
event_id: EventId
stream_id: StreamId
message: JSONRPCMessage
message: JSONRPCMessage | None
class InMemoryEventStore(EventStore):
@@ -48,7 +48,7 @@ class InMemoryEventStore(EventStore):
# event_id -> EventEntry for quick lookup
self.event_index: dict[EventId, EventEntry] = {}
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage) -> EventId:
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId:
"""Stores an event with a generated event ID."""
event_id = str(uuid4())
event_entry = EventEntry(event_id=event_id, stream_id=stream_id, message=message)
@@ -88,7 +88,9 @@ class InMemoryEventStore(EventStore):
found_last = False
for event in stream_events:
if found_last:
await send_callback(EventMessage(event.message, event.event_id))
# Skip priming events (None message)
if event.message is not None:
await send_callback(EventMessage(event.message, event.event_id))
elif event.event_id == last_event_id:
found_last = True
@@ -0,0 +1,36 @@
# MCP SSE Polling Demo Server
Demonstrates the SSE polling pattern with server-initiated stream close for long-running tasks (SEP-1699).
## Features
- Priming events (automatic with EventStore)
- Server-initiated stream close via `close_sse_stream()` callback
- Client auto-reconnect with Last-Event-ID
- Progress notifications during long-running tasks
- Configurable retry interval
## Usage
```bash
# Start server on default port
uv run mcp-sse-polling-demo --port 3000
# Custom retry interval (milliseconds)
uv run mcp-sse-polling-demo --port 3000 --retry-interval 100
```
## Tool: process_batch
Processes items with periodic checkpoints that trigger SSE stream closes:
- `items`: Number of items to process (1-100, default: 10)
- `checkpoint_every`: Close stream after this many items (1-20, default: 3)
## Client
Use the companion `mcp-sse-polling-client` to test:
```bash
uv run mcp-sse-polling-client --url http://localhost:3000/mcp
```
@@ -0,0 +1 @@
"""SSE Polling Demo Server - demonstrates close_sse_stream for long-running tasks."""
@@ -0,0 +1,6 @@
"""Entry point for the SSE Polling Demo server."""
from .server import main
if __name__ == "__main__":
main()
@@ -0,0 +1,100 @@
"""
In-memory event store for demonstrating resumability functionality.
This is a simple implementation intended for examples and testing,
not for production use where a persistent storage solution would be more appropriate.
"""
import logging
from collections import deque
from dataclasses import dataclass
from uuid import uuid4
from mcp.server.streamable_http import EventCallback, EventId, EventMessage, EventStore, StreamId
from mcp.types import JSONRPCMessage
logger = logging.getLogger(__name__)
@dataclass
class EventEntry:
"""Represents an event entry in the event store."""
event_id: EventId
stream_id: StreamId
message: JSONRPCMessage | None # None for priming events
class InMemoryEventStore(EventStore):
"""
Simple in-memory implementation of the EventStore interface for resumability.
This is primarily intended for examples and testing, not for production use
where a persistent storage solution would be more appropriate.
This implementation keeps only the last N events per stream for memory efficiency.
"""
def __init__(self, max_events_per_stream: int = 100):
"""Initialize the event store.
Args:
max_events_per_stream: Maximum number of events to keep per stream
"""
self.max_events_per_stream = max_events_per_stream
# for maintaining last N events per stream
self.streams: dict[StreamId, deque[EventEntry]] = {}
# event_id -> EventEntry for quick lookup
self.event_index: dict[EventId, EventEntry] = {}
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId:
"""Stores an event with a generated event ID.
Args:
stream_id: ID of the stream the event belongs to
message: The message to store, or None for priming events
"""
event_id = str(uuid4())
event_entry = EventEntry(event_id=event_id, stream_id=stream_id, message=message)
# Get or create deque for this stream
if stream_id not in self.streams:
self.streams[stream_id] = deque(maxlen=self.max_events_per_stream)
# If deque is full, the oldest event will be automatically removed
# We need to remove it from the event_index as well
if len(self.streams[stream_id]) == self.max_events_per_stream:
oldest_event = self.streams[stream_id][0]
self.event_index.pop(oldest_event.event_id, None)
# Add new event
self.streams[stream_id].append(event_entry)
self.event_index[event_id] = event_entry
return event_id
async def replay_events_after(
self,
last_event_id: EventId,
send_callback: EventCallback,
) -> StreamId | None:
"""Replays events that occurred after the specified event ID."""
if last_event_id not in self.event_index:
logger.warning(f"Event ID {last_event_id} not found in store")
return None
# Get the stream and find events after the last one
last_event = self.event_index[last_event_id]
stream_id = last_event.stream_id
stream_events = self.streams.get(last_event.stream_id, deque())
# Events in deque are already in chronological order
found_last = False
for event in stream_events:
if found_last:
# Skip priming events (None messages) during replay
if event.message is not None:
await send_callback(EventMessage(event.message, event.event_id))
elif event.event_id == last_event_id:
found_last = True
return stream_id
@@ -0,0 +1,177 @@
"""
SSE Polling Demo Server
Demonstrates the SSE polling pattern with close_sse_stream() for long-running tasks.
Features demonstrated:
- Priming events (automatic with EventStore)
- Server-initiated stream close via close_sse_stream callback
- Client auto-reconnect with Last-Event-ID
- Progress notifications during long-running tasks
Run with:
uv run mcp-sse-polling-demo --port 3000
"""
import contextlib
import logging
from collections.abc import AsyncIterator
from typing import Any
import anyio
import click
import mcp.types as types
from mcp.server.lowlevel import Server
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
from starlette.applications import Starlette
from starlette.routing import Mount
from starlette.types import Receive, Scope, Send
from .event_store import InMemoryEventStore
logger = logging.getLogger(__name__)
@click.command()
@click.option("--port", default=3000, help="Port to listen on")
@click.option(
"--log-level",
default="INFO",
help="Logging level (DEBUG, INFO, WARNING, ERROR)",
)
@click.option(
"--retry-interval",
default=100,
help="SSE retry interval in milliseconds (sent to client)",
)
def main(port: int, log_level: str, retry_interval: int) -> int:
"""Run the SSE Polling Demo server."""
logging.basicConfig(
level=getattr(logging, log_level.upper()),
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
# Create the lowlevel server
app = Server("sse-polling-demo")
@app.call_tool()
async def call_tool(name: str, arguments: dict[str, Any]) -> list[types.ContentBlock]:
"""Handle tool calls."""
ctx = app.request_context
if name == "process_batch":
items = arguments.get("items", 10)
checkpoint_every = arguments.get("checkpoint_every", 3)
if items < 1 or items > 100:
return [types.TextContent(type="text", text="Error: items must be between 1 and 100")]
if checkpoint_every < 1 or checkpoint_every > 20:
return [types.TextContent(type="text", text="Error: checkpoint_every must be between 1 and 20")]
await ctx.session.send_log_message(
level="info",
data=f"Starting batch processing of {items} items...",
logger="process_batch",
related_request_id=ctx.request_id,
)
for i in range(1, items + 1):
# Simulate work
await anyio.sleep(0.5)
# Report progress
await ctx.session.send_log_message(
level="info",
data=f"[{i}/{items}] Processing item {i}",
logger="process_batch",
related_request_id=ctx.request_id,
)
# Checkpoint: close stream to trigger client reconnect
if i % checkpoint_every == 0 and i < items:
await ctx.session.send_log_message(
level="info",
data=f"Checkpoint at item {i} - closing SSE stream for polling",
logger="process_batch",
related_request_id=ctx.request_id,
)
if ctx.close_sse_stream:
logger.info(f"Closing SSE stream at checkpoint {i}")
await ctx.close_sse_stream()
# Wait for client to reconnect (must be > retry_interval of 100ms)
await anyio.sleep(0.2)
return [
types.TextContent(
type="text",
text=f"Successfully processed {items} items with checkpoints every {checkpoint_every} items",
)
]
return [types.TextContent(type="text", text=f"Unknown tool: {name}")]
@app.list_tools()
async def list_tools() -> list[types.Tool]:
"""List available tools."""
return [
types.Tool(
name="process_batch",
description=(
"Process a batch of items with periodic checkpoints. "
"Demonstrates SSE polling where server closes stream periodically."
),
inputSchema={
"type": "object",
"properties": {
"items": {
"type": "integer",
"description": "Number of items to process (1-100)",
"default": 10,
},
"checkpoint_every": {
"type": "integer",
"description": "Close stream after this many items (1-20)",
"default": 3,
},
},
},
)
]
# Create event store for resumability
event_store = InMemoryEventStore()
# Create session manager with event store and retry interval
session_manager = StreamableHTTPSessionManager(
app=app,
event_store=event_store,
retry_interval=retry_interval,
)
async def handle_streamable_http(scope: Scope, receive: Receive, send: Send) -> None:
await session_manager.handle_request(scope, receive, send)
@contextlib.asynccontextmanager
async def lifespan(starlette_app: Starlette) -> AsyncIterator[None]:
async with session_manager.run():
logger.info(f"SSE Polling Demo server started on port {port}")
logger.info("Try: POST /mcp with tools/call for 'process_batch'")
yield
logger.info("Server shutting down...")
starlette_app = Starlette(
debug=True,
routes=[
Mount("/mcp", app=handle_streamable_http),
],
lifespan=lifespan,
)
import uvicorn
uvicorn.run(starlette_app, host="127.0.0.1", port=port)
return 0
if __name__ == "__main__":
main()
@@ -0,0 +1,36 @@
[project]
name = "mcp-sse-polling-demo"
version = "0.1.0"
description = "Demo server showing SSE polling with close_sse_stream for long-running tasks"
readme = "README.md"
requires-python = ">=3.10"
authors = [{ name = "Anthropic, PBC." }]
keywords = ["mcp", "sse", "polling", "streamable", "http"]
license = { text = "MIT" }
dependencies = ["anyio>=4.5", "click>=8.2.0", "httpx>=0.27", "mcp", "starlette", "uvicorn"]
[project.scripts]
mcp-sse-polling-demo = "mcp_sse_polling_demo.server:main"
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["mcp_sse_polling_demo"]
[tool.pyright]
include = ["mcp_sse_polling_demo"]
venvPath = "."
venv = ".venv"
[tool.ruff.lint]
select = ["E", "F", "I"]
ignore = []
[tool.ruff]
line-length = 120
target-version = "py310"
[dependency-groups]
dev = ["pyright>=1.1.378", "pytest>=8.3.3", "ruff>=0.6.9"]
+138 -25
View File
@@ -42,6 +42,10 @@ GetSessionIdCallback = Callable[[], str | None]
MCP_SESSION_ID = "mcp-session-id"
MCP_PROTOCOL_VERSION = "mcp-protocol-version"
LAST_EVENT_ID = "last-event-id"
# Reconnection defaults
DEFAULT_RECONNECTION_DELAY_MS = 1000 # 1 second fallback when server doesn't provide retry
MAX_RECONNECTION_ATTEMPTS = 2 # Max retry attempts before giving up
CONTENT_TYPE = "content-type"
ACCEPT = "accept"
@@ -160,8 +164,11 @@ class StreamableHTTPTransport:
) -> bool:
"""Handle an SSE event, returning True if the response is complete."""
if sse.event == "message":
# Skip empty data (keep-alive pings)
# Handle priming events (empty data with ID) for resumability
if not sse.data:
# Call resumption callback for priming events that have an ID
if sse.id and resumption_callback:
await resumption_callback(sse.id)
return False
try:
message = JSONRPCMessage.model_validate_json(sse.data)
@@ -199,28 +206,55 @@ class StreamableHTTPTransport:
client: httpx.AsyncClient,
read_stream_writer: StreamWriter,
) -> None:
"""Handle GET stream for server-initiated messages."""
try:
if not self.session_id:
"""Handle GET stream for server-initiated messages with auto-reconnect."""
last_event_id: str | None = None
retry_interval_ms: int | None = None
attempt: int = 0
while attempt < MAX_RECONNECTION_ATTEMPTS: # pragma: no branch
try:
if not self.session_id:
return
headers = self._prepare_request_headers(self.request_headers)
if last_event_id:
headers[LAST_EVENT_ID] = last_event_id # pragma: no cover
async with aconnect_sse(
client,
"GET",
self.url,
headers=headers,
timeout=httpx.Timeout(self.timeout, read=self.sse_read_timeout),
) as event_source:
event_source.response.raise_for_status()
logger.debug("GET SSE connection established")
async for sse in event_source.aiter_sse():
# Track last event ID for reconnection
if sse.id:
last_event_id = sse.id # pragma: no cover
# Track retry interval from server
if sse.retry is not None:
retry_interval_ms = sse.retry # pragma: no cover
await self._handle_sse_event(sse, read_stream_writer)
# Stream ended normally (server closed) - reset attempt counter
attempt = 0
except Exception as exc: # pragma: no cover
logger.debug(f"GET stream error: {exc}")
attempt += 1
if attempt >= MAX_RECONNECTION_ATTEMPTS: # pragma: no cover
logger.debug(f"GET stream max reconnection attempts ({MAX_RECONNECTION_ATTEMPTS}) exceeded")
return
headers = self._prepare_request_headers(self.request_headers)
async with aconnect_sse(
client,
"GET",
self.url,
headers=headers,
timeout=httpx.Timeout(self.timeout, read=self.sse_read_timeout),
) as event_source:
event_source.response.raise_for_status()
logger.debug("GET SSE connection established")
async for sse in event_source.aiter_sse():
await self._handle_sse_event(sse, read_stream_writer)
except Exception as exc:
logger.debug(f"GET stream error (non-fatal): {exc}") # pragma: no cover
# Wait before reconnecting
delay_ms = retry_interval_ms if retry_interval_ms is not None else DEFAULT_RECONNECTION_DELAY_MS
logger.info(f"GET stream disconnected, reconnecting in {delay_ms}ms...")
await anyio.sleep(delay_ms / 1000.0)
async def _handle_resumption_request(self, ctx: RequestContext) -> None:
"""Handle a resumption request using GET with SSE."""
@@ -326,9 +360,20 @@ class StreamableHTTPTransport:
is_initialization: bool = False,
) -> None:
"""Handle SSE response from the server."""
last_event_id: str | None = None
retry_interval_ms: int | None = None
try:
event_source = EventSource(response)
async for sse in event_source.aiter_sse(): # pragma: no branch
# Track last event ID for potential reconnection
if sse.id:
last_event_id = sse.id
# Track retry interval from server
if sse.retry is not None:
retry_interval_ms = sse.retry
is_complete = await self._handle_sse_event(
sse,
ctx.read_stream_writer,
@@ -339,10 +384,78 @@ class StreamableHTTPTransport:
# break the loop
if is_complete:
await response.aclose()
break
except Exception as e:
logger.exception("Error reading SSE stream:") # pragma: no cover
await ctx.read_stream_writer.send(e) # pragma: no cover
return # Normal completion, no reconnect needed
except Exception as e: # pragma: no cover
logger.debug(f"SSE stream ended: {e}")
# Stream ended without response - reconnect if we received an event with ID
if last_event_id is not None: # pragma: no branch
logger.info("SSE stream disconnected, reconnecting...")
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms)
async def _handle_reconnection(
self,
ctx: RequestContext,
last_event_id: str,
retry_interval_ms: int | None = None,
attempt: int = 0,
) -> None:
"""Reconnect with Last-Event-ID to resume stream after server disconnect."""
# Bail if max retries exceeded
if attempt >= MAX_RECONNECTION_ATTEMPTS: # pragma: no cover
logger.debug(f"Max reconnection attempts ({MAX_RECONNECTION_ATTEMPTS}) exceeded")
return
# Always wait - use server value or default
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_request_headers(ctx.headers)
headers[LAST_EVENT_ID] = last_event_id
# Extract original request ID to map responses
original_request_id = None
if isinstance(ctx.session_message.message.root, JSONRPCRequest): # pragma: no branch
original_request_id = ctx.session_message.message.root.id
try:
async with aconnect_sse(
ctx.client,
"GET",
self.url,
headers=headers,
timeout=httpx.Timeout(self.timeout, read=self.sse_read_timeout),
) as event_source:
event_source.response.raise_for_status()
logger.info("Reconnected to SSE stream")
# Track for potential further reconnection
reconnect_last_event_id: str = last_event_id
reconnect_retry_ms = retry_interval_ms
async for sse in event_source.aiter_sse():
if sse.id: # pragma: no branch
reconnect_last_event_id = sse.id
if sse.retry is not None:
reconnect_retry_ms = sse.retry
is_complete = await self._handle_sse_event(
sse,
ctx.read_stream_writer,
original_request_id,
ctx.metadata.on_resumption_token_update if ctx.metadata else None,
)
if is_complete:
await event_source.response.aclose()
return
# Stream ended again without response - reconnect again (reset attempt counter)
logger.info("SSE stream disconnected, reconnecting...")
await self._handle_reconnection(ctx, reconnect_last_event_id, reconnect_retry_ms, 0)
except Exception as e: # pragma: no cover
logger.debug(f"Reconnection failed: {e}")
# Try to reconnect again if we still have an event ID
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms, attempt + 1)
async def _handle_unexpected_content_type(
self,
+35
View File
@@ -153,6 +153,7 @@ class FastMCP(Generic[LifespanResultT]):
auth_server_provider: (OAuthAuthorizationServerProvider[Any, Any, Any] | None) = None,
token_verifier: TokenVerifier | None = None,
event_store: EventStore | None = None,
retry_interval: int | None = None,
*,
tools: list[Tool] | None = None,
debug: bool = False,
@@ -221,6 +222,7 @@ class FastMCP(Generic[LifespanResultT]):
if auth_server_provider and not token_verifier: # pragma: no cover
self._token_verifier = ProviderTokenVerifier(auth_server_provider)
self._event_store = event_store
self._retry_interval = retry_interval
self._custom_starlette_routes: list[Route] = []
self.dependencies = self.settings.dependencies
self._session_manager: StreamableHTTPSessionManager | None = None
@@ -940,6 +942,7 @@ class FastMCP(Generic[LifespanResultT]):
self._session_manager = StreamableHTTPSessionManager(
app=self._mcp_server,
event_store=self._event_store,
retry_interval=self._retry_interval,
json_response=self.settings.json_response,
stateless=self.settings.stateless_http, # Use the stateless setting
security_settings=self.settings.transport_security,
@@ -1282,6 +1285,38 @@ class Context(BaseModel, Generic[ServerSessionT, LifespanContextT, RequestT]):
"""Access to the underlying session for advanced usage."""
return self.request_context.session
async def close_sse_stream(self) -> None:
"""Close the SSE stream to trigger client reconnection.
This method closes the HTTP connection for the current request, triggering
client reconnection. Events continue to be stored in the event store and will
be replayed when the client reconnects with Last-Event-ID.
Use this to implement polling behavior during long-running operations -
client will reconnect after the retry interval specified in the priming event.
Note:
This is a no-op if not using StreamableHTTP transport with event_store.
The callback is only available when event_store is configured.
"""
if self._request_context and self._request_context.close_sse_stream: # pragma: no cover
await self._request_context.close_sse_stream()
async def close_standalone_sse_stream(self) -> None:
"""Close the standalone GET SSE stream to trigger client reconnection.
This method closes the HTTP connection for the standalone GET stream used
for unsolicited server-to-client notifications. The client SHOULD reconnect
with Last-Event-ID to resume receiving notifications.
Note:
This is a no-op if not using StreamableHTTP transport with event_store.
Currently, client reconnection for standalone GET streams is NOT
implemented - this is a known gap.
"""
if self._request_context and self._request_context.close_standalone_sse_stream: # pragma: no cover
await self._request_context.close_standalone_sse_stream()
# Convenience methods for common log levels
async def debug(self, message: str, **extra: Any) -> None:
"""Send a debug log message."""
+7 -1
View File
@@ -713,12 +713,16 @@ class Server(Generic[LifespanResultT, RequestT]):
token = None
try:
# Extract request context from message metadata
# Extract request context and close_sse_stream from message metadata
request_data = None
close_sse_stream_cb = None
close_standalone_sse_stream_cb = None
if message.message_metadata is not None and isinstance(
message.message_metadata, ServerMessageMetadata
): # pragma: no cover
request_data = message.message_metadata.request_context
close_sse_stream_cb = message.message_metadata.close_sse_stream
close_standalone_sse_stream_cb = message.message_metadata.close_standalone_sse_stream
# Set our global state that can be retrieved via
# app.get_request_context()
@@ -741,6 +745,8 @@ class Server(Generic[LifespanResultT, RequestT]):
_task_support=task_support,
),
request=request_data,
close_sse_stream=close_sse_stream_cb,
close_standalone_sse_stream=close_standalone_sse_stream_cb,
)
)
response = await handler(req)
+116 -4
View File
@@ -15,6 +15,7 @@ from collections.abc import AsyncGenerator, Awaitable, Callable
from contextlib import asynccontextmanager
from dataclasses import dataclass
from http import HTTPStatus
from typing import Any
import anyio
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
@@ -87,13 +88,13 @@ class EventStore(ABC):
"""
@abstractmethod
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage) -> EventId:
async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId:
"""
Stores an event for later retrieval.
Args:
stream_id: ID of the stream the event belongs to
message: The JSON-RPC message to store
message: The JSON-RPC message to store, or None for priming events
Returns:
The generated event ID for the stored event
@@ -140,6 +141,7 @@ class StreamableHTTPServerTransport:
is_json_response_enabled: bool = False,
event_store: EventStore | None = None,
security_settings: TransportSecuritySettings | None = None,
retry_interval: int | None = None,
) -> None:
"""
Initialize a new StreamableHTTP server transport.
@@ -153,6 +155,10 @@ class StreamableHTTPServerTransport:
resumability will be enabled, allowing clients to
reconnect and resume messages.
security_settings: Optional security settings for DNS rebinding protection.
retry_interval: Retry interval in milliseconds to suggest to clients in SSE
retry field. When set, the server will send a retry field in
SSE priming events to control client reconnection timing for
polling behavior. Only used when event_store is provided.
Raises:
ValueError: If the session ID contains invalid characters.
@@ -164,6 +170,7 @@ class StreamableHTTPServerTransport:
self.is_json_response_enabled = is_json_response_enabled
self._event_store = event_store
self._security = TransportSecurityMiddleware(security_settings)
self._retry_interval = retry_interval
self._request_streams: dict[
RequestId,
tuple[
@@ -171,6 +178,7 @@ class StreamableHTTPServerTransport:
MemoryObjectReceiveStream[EventMessage],
],
] = {}
self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[dict[str, str]]] = {}
self._terminated = False
@property
@@ -178,6 +186,91 @@ class StreamableHTTPServerTransport:
"""Check if this transport has been explicitly terminated."""
return self._terminated
def close_sse_stream(self, request_id: RequestId) -> None: # pragma: no cover
"""Close SSE connection for a specific request without terminating the stream.
This method closes the HTTP connection for the specified request, triggering
client reconnection. Events continue to be stored in the event store and will
be replayed when the client reconnects with Last-Event-ID.
Use this to implement polling behavior during long-running operations -
client will reconnect after the retry interval specified in the priming event.
Args:
request_id: The request ID whose SSE stream should be closed.
Note:
This is a no-op if there is no active stream for the request ID.
Requires event_store to be configured for events to be stored during
the disconnect.
"""
writer = self._sse_stream_writers.pop(request_id, None)
if writer:
writer.close()
# Also close and remove request streams
if request_id in self._request_streams:
send_stream, receive_stream = self._request_streams.pop(request_id)
send_stream.close()
receive_stream.close()
def close_standalone_sse_stream(self) -> None: # pragma: no cover
"""Close the standalone GET SSE stream, triggering client reconnection.
This method closes the HTTP connection for the standalone GET stream used
for unsolicited server-to-client notifications. The client SHOULD reconnect
with Last-Event-ID to resume receiving notifications.
Use this to implement polling behavior for the notification stream -
client will reconnect after the retry interval specified in the priming event.
Note:
This is a no-op if there is no active standalone SSE stream.
Requires event_store to be configured for events to be stored during
the disconnect.
Currently, client reconnection for standalone GET streams is NOT
implemented - this is a known gap (see test_standalone_get_stream_reconnection).
"""
self.close_sse_stream(GET_STREAM_KEY)
def _create_session_message( # pragma: no cover
self,
message: JSONRPCMessage,
request: Request,
request_id: RequestId,
) -> SessionMessage:
"""Create a session message with metadata including close_sse_stream callback."""
async def close_stream_callback() -> None:
self.close_sse_stream(request_id)
async def close_standalone_stream_callback() -> None:
self.close_standalone_sse_stream()
metadata = ServerMessageMetadata(
request_context=request,
close_sse_stream=close_stream_callback,
close_standalone_sse_stream=close_standalone_stream_callback,
)
return SessionMessage(message, metadata=metadata)
async def _send_priming_event( # pragma: no cover
self,
request_id: RequestId,
sse_stream_writer: MemoryObjectSendStream[dict[str, Any]],
) -> None:
"""Send priming event for SSE resumability if event_store is configured."""
if not self._event_store:
return
priming_event_id = await self._event_store.store_event(
str(request_id), # Convert RequestId to StreamId (str)
None, # Priming event has no payload
)
priming_event: dict[str, str | int] = {"id": priming_event_id, "data": ""}
if self._retry_interval is not None:
priming_event["retry"] = self._retry_interval
await sse_stream_writer.send(priming_event)
def _create_error_response(
self,
error_message: str,
@@ -459,10 +552,16 @@ class StreamableHTTPServerTransport:
# Create SSE stream
sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, str]](0)
# Store writer reference so close_sse_stream() can close it
self._sse_stream_writers[request_id] = sse_stream_writer
async def sse_writer():
# Get the request ID from the incoming request message
try:
async with sse_stream_writer, request_stream_reader:
# Send priming event for SSE resumability
await self._send_priming_event(request_id, sse_stream_writer)
# Process messages from the request-specific stream
async for event_message in request_stream_reader:
# Build the event data
@@ -475,10 +574,14 @@ class StreamableHTTPServerTransport:
JSONRPCResponse | JSONRPCError,
):
break
except anyio.ClosedResourceError:
# Expected when close_sse_stream() is called
logger.debug("SSE stream closed by close_sse_stream()")
except Exception:
logger.exception("Error in SSE writer")
finally:
logger.debug("Closing SSE writer")
self._sse_stream_writers.pop(request_id, None)
await self._clean_up_memory_streams(request_id)
# Create and start EventSourceResponse
@@ -502,8 +605,7 @@ class StreamableHTTPServerTransport:
async with anyio.create_task_group() as tg:
tg.start_soon(response, scope, receive, send)
# Then send the message to be processed by the server
metadata = ServerMessageMetadata(request_context=request)
session_message = SessionMessage(message, metadata=metadata)
session_message = self._create_session_message(message, request, request_id)
await writer.send(session_message)
except Exception:
logger.exception("SSE response error")
@@ -778,6 +880,13 @@ class StreamableHTTPServerTransport:
# If stream ID not in mapping, create it
if stream_id and stream_id not in self._request_streams:
# Register SSE writer so close_sse_stream() can close it
self._sse_stream_writers[stream_id] = sse_stream_writer
# Send priming event for this new connection
await self._send_priming_event(stream_id, sse_stream_writer)
# Create new request streams for this connection
self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage](0)
msg_reader = self._request_streams[stream_id][1]
@@ -787,6 +896,9 @@ class StreamableHTTPServerTransport:
event_data = self._create_event_data(event_message)
await sse_stream_writer.send(event_data)
except anyio.ClosedResourceError:
# Expected when close_sse_stream() is called
logger.debug("Replay SSE stream closed by close_sse_stream()")
except Exception:
logger.exception("Error in replay sender")
@@ -51,6 +51,9 @@ class StreamableHTTPSessionManager:
json_response: Whether to use JSON responses instead of SSE streams
stateless: If True, creates a completely fresh transport for each request
with no session tracking or state persistence between requests.
security_settings: Optional transport security settings.
retry_interval: Retry interval in milliseconds to suggest to clients in SSE
retry field. Used for SSE polling behavior.
"""
def __init__(
@@ -60,12 +63,14 @@ class StreamableHTTPSessionManager:
json_response: bool = False,
stateless: bool = False,
security_settings: TransportSecuritySettings | None = None,
retry_interval: int | None = None,
):
self.app = app
self.event_store = event_store
self.json_response = json_response
self.stateless = stateless
self.security_settings = security_settings
self.retry_interval = retry_interval
# Session tracking (only used if not stateless)
self._session_creation_lock = anyio.Lock()
@@ -226,6 +231,7 @@ class StreamableHTTPSessionManager:
is_json_response_enabled=self.json_response,
event_store=self.event_store, # May be None (no resumability)
security_settings=self.security_settings,
retry_interval=self.retry_interval,
)
assert http_transport.mcp_session_id is not None
+3
View File
@@ -7,6 +7,7 @@ from typing import Any, Generic
from typing_extensions import TypeVar
from mcp.shared.message import CloseSSEStreamCallback
from mcp.shared.session import BaseSession
from mcp.types import RequestId, RequestParams
@@ -27,3 +28,5 @@ class RequestContext(Generic[SessionT, LifespanContextT, RequestT]):
# The Server sets this to an Experimental instance at runtime.
experimental: Any = field(default=None)
request: RequestT | None = None
close_sse_stream: CloseSSEStreamCallback | None = None
close_standalone_sse_stream: CloseSSEStreamCallback | None = None
+7
View File
@@ -14,6 +14,9 @@ ResumptionToken = str
ResumptionTokenUpdateCallback = Callable[[ResumptionToken], Awaitable[None]]
# Callback type for closing SSE streams without terminating
CloseSSEStreamCallback = Callable[[], Awaitable[None]]
@dataclass
class ClientMessageMetadata:
@@ -30,6 +33,10 @@ class ServerMessageMetadata:
related_request_id: RequestId | None = None
# Request-specific context (e.g., headers, auth info)
request_context: object | None = None
# Callback to close SSE stream for the current request without terminating
close_sse_stream: CloseSSEStreamCallback | None = None
# Callback to close the standalone GET SSE stream (for unsolicited notifications)
close_standalone_sse_stream: CloseSSEStreamCallback | None = None
MessageMetadata = ClientMessageMetadata | ServerMessageMetadata | None
+493 -14
View File
@@ -7,6 +7,7 @@ Contains tests for both server and client sides of the StreamableHTTP transport.
import json
import multiprocessing
import socket
import time
from collections.abc import Generator
from typing import Any
@@ -76,10 +77,12 @@ class SimpleEventStore(EventStore):
"""Simple in-memory event store for testing."""
def __init__(self):
self._events: list[tuple[StreamId, EventId, types.JSONRPCMessage]] = []
self._events: list[tuple[StreamId, EventId, types.JSONRPCMessage | None]] = []
self._event_id_counter = 0
async def store_event(self, stream_id: StreamId, message: types.JSONRPCMessage) -> EventId: # pragma: no cover
async def store_event( # pragma: no cover
self, stream_id: StreamId, message: types.JSONRPCMessage | None
) -> EventId:
"""Store an event and return its ID."""
self._event_id_counter += 1
event_id = str(self._event_id_counter)
@@ -109,7 +112,9 @@ class SimpleEventStore(EventStore):
# Replay only events from the same stream with ID > last_event_id
for stream_id, event_id, message in self._events:
if stream_id == target_stream_id and int(event_id) > last_event_id_int:
await send_callback(EventMessage(message, event_id))
# Skip priming events (None message)
if message is not None:
await send_callback(EventMessage(message, event_id))
return target_stream_id
@@ -164,6 +169,32 @@ class ServerTest(Server): # pragma: no cover
description="A tool that releases the lock",
inputSchema={"type": "object", "properties": {}},
),
Tool(
name="tool_with_stream_close",
description="A tool that closes SSE stream mid-operation",
inputSchema={"type": "object", "properties": {}},
),
Tool(
name="tool_with_multiple_notifications_and_close",
description="Tool that sends notification1, closes stream, sends notification2, notification3",
inputSchema={"type": "object", "properties": {}},
),
Tool(
name="tool_with_multiple_stream_closes",
description="Tool that closes SSE stream multiple times during execution",
inputSchema={
"type": "object",
"properties": {
"checkpoints": {"type": "integer", "default": 3},
"sleep_time": {"type": "number", "default": 0.2},
},
},
),
Tool(
name="tool_with_standalone_stream_close",
description="Tool that closes standalone GET stream mid-operation",
inputSchema={"type": "object", "properties": {}},
),
]
@self.call_tool()
@@ -255,17 +286,107 @@ class ServerTest(Server): # pragma: no cover
self._lock.set()
return [TextContent(type="text", text="Lock released")]
elif name == "tool_with_stream_close":
# Send notification before closing
await ctx.session.send_log_message(
level="info",
data="Before close",
logger="stream_close_tool",
related_request_id=ctx.request_id,
)
# Close SSE stream (triggers client reconnect)
assert ctx.close_sse_stream is not None
await ctx.close_sse_stream()
# Continue processing (events stored in event_store)
await anyio.sleep(0.1)
await ctx.session.send_log_message(
level="info",
data="After close",
logger="stream_close_tool",
related_request_id=ctx.request_id,
)
return [TextContent(type="text", text="Done")]
elif name == "tool_with_multiple_notifications_and_close":
# Send notification1
await ctx.session.send_log_message(
level="info",
data="notification1",
logger="multi_notif_tool",
related_request_id=ctx.request_id,
)
# Close SSE stream
assert ctx.close_sse_stream is not None
await ctx.close_sse_stream()
# Send notification2, notification3 (stored in event_store)
await anyio.sleep(0.1)
await ctx.session.send_log_message(
level="info",
data="notification2",
logger="multi_notif_tool",
related_request_id=ctx.request_id,
)
await ctx.session.send_log_message(
level="info",
data="notification3",
logger="multi_notif_tool",
related_request_id=ctx.request_id,
)
return [TextContent(type="text", text="All notifications sent")]
elif name == "tool_with_multiple_stream_closes":
num_checkpoints = args.get("checkpoints", 3)
sleep_time = args.get("sleep_time", 0.2)
for i in range(num_checkpoints):
await ctx.session.send_log_message(
level="info",
data=f"checkpoint_{i}",
logger="multi_close_tool",
related_request_id=ctx.request_id,
)
if ctx.close_sse_stream:
await ctx.close_sse_stream()
await anyio.sleep(sleep_time)
return [TextContent(type="text", text=f"Completed {num_checkpoints} checkpoints")]
elif name == "tool_with_standalone_stream_close":
# Test for GET stream reconnection
# 1. Send unsolicited notification via GET stream (no related_request_id)
await ctx.session.send_resource_updated(uri=AnyUrl("http://notification_1"))
# Small delay to ensure notification is flushed before closing
await anyio.sleep(0.1)
# 2. Close the standalone GET stream
if ctx.close_standalone_sse_stream:
await ctx.close_standalone_sse_stream()
# 3. Wait for client to reconnect (uses retry_interval from server, default 1000ms)
await anyio.sleep(1.5)
# 4. Send another notification on the new GET stream connection
await ctx.session.send_resource_updated(uri=AnyUrl("http://notification_2"))
return [TextContent(type="text", text="Standalone stream close test done")]
return [TextContent(type="text", text=f"Called {name}")]
def create_app(
is_json_response_enabled: bool = False, event_store: EventStore | None = None
is_json_response_enabled: bool = False,
event_store: EventStore | None = None,
retry_interval: int | None = None,
) -> Starlette: # pragma: no cover
"""Create a Starlette application for testing using the session manager.
Args:
is_json_response_enabled: If True, use JSON responses instead of SSE streams.
event_store: Optional event store for testing resumability.
retry_interval: Retry interval in milliseconds for SSE polling.
"""
# Create server instance
server = ServerTest()
@@ -279,6 +400,7 @@ def create_app(
event_store=event_store,
json_response=is_json_response_enabled,
security_settings=security_settings,
retry_interval=retry_interval,
)
# Create an ASGI application that uses the session manager
@@ -294,7 +416,10 @@ def create_app(
def run_server(
port: int, is_json_response_enabled: bool = False, event_store: EventStore | None = None
port: int,
is_json_response_enabled: bool = False,
event_store: EventStore | None = None,
retry_interval: int | None = None,
) -> None: # pragma: no cover
"""Run the test server.
@@ -302,9 +427,10 @@ def run_server(
port: Port to listen on.
is_json_response_enabled: If True, use JSON responses instead of SSE streams.
event_store: Optional event store for testing resumability.
retry_interval: Retry interval in milliseconds for SSE polling.
"""
app = create_app(is_json_response_enabled, event_store)
app = create_app(is_json_response_enabled, event_store, retry_interval)
# Configure server
config = uvicorn.Config(
app=app,
@@ -379,10 +505,10 @@ def event_server_port() -> int:
def event_server(
event_server_port: int, event_store: SimpleEventStore
) -> Generator[tuple[SimpleEventStore, str], None, None]:
"""Start a server with event store enabled."""
"""Start a server with event store and retry_interval enabled."""
proc = multiprocessing.Process(
target=run_server,
kwargs={"port": event_server_port, "event_store": event_store},
kwargs={"port": event_server_port, "event_store": event_store, "retry_interval": 500},
daemon=True,
)
proc.start()
@@ -883,7 +1009,7 @@ async def test_streamablehttp_client_tool_invocation(initialized_client_session:
"""Test client tool invocation."""
# First list tools
tools = await initialized_client_session.list_tools()
assert len(tools.tools) == 6
assert len(tools.tools) == 10
assert tools.tools[0].name == "test_tool"
# Call the tool
@@ -920,7 +1046,7 @@ async def test_streamablehttp_client_session_persistence(basic_server: None, bas
# Make multiple requests to verify session persistence
tools = await session.list_tools()
assert len(tools.tools) == 6
assert len(tools.tools) == 10
# Read a resource
resource = await session.read_resource(uri=AnyUrl("foobar://test-persist"))
@@ -949,7 +1075,7 @@ async def test_streamablehttp_client_json_response(json_response_server: None, j
# Check tool listing
tools = await session.list_tools()
assert len(tools.tools) == 6
assert len(tools.tools) == 10
# Call a tool and verify JSON response handling
result = await session.call_tool("test_tool", {})
@@ -962,7 +1088,6 @@ async def test_streamablehttp_client_json_response(json_response_server: None, j
async def test_streamablehttp_client_get_stream(basic_server: None, basic_server_url: str):
"""Test GET stream functionality for server-initiated messages."""
import mcp.types as types
from mcp.shared.session import RequestResponder
notifications_received: list[types.ServerNotification] = []
@@ -1020,7 +1145,7 @@ async def test_streamablehttp_client_session_termination(basic_server: None, bas
# Make a request to confirm session is working
tools = await session.list_tools()
assert len(tools.tools) == 6
assert len(tools.tools) == 10
headers: dict[str, str] = {} # pragma: no cover
if captured_session_id: # pragma: no cover
@@ -1086,7 +1211,7 @@ async def test_streamablehttp_client_session_termination_204(
# Make a request to confirm session is working
tools = await session.list_tools()
assert len(tools.tools) == 6
assert len(tools.tools) == 10
headers: dict[str, str] = {} # pragma: no cover
if captured_session_id: # pragma: no cover
@@ -1633,3 +1758,357 @@ async def test_handle_sse_event_skips_empty_data():
finally:
await write_stream.aclose()
await read_stream.aclose()
@pytest.mark.anyio
async def test_streamablehttp_client_receives_priming_event(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""Client should receive priming event (resumption token update) on POST SSE stream."""
_, server_url = event_server
captured_resumption_tokens: list[str] = []
async def on_resumption_token_update(token: str) -> None:
captured_resumption_tokens.append(token)
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
# Call tool with resumption token callback via send_request
metadata = ClientMessageMetadata(
on_resumption_token_update=on_resumption_token_update,
)
result = await session.send_request(
types.ClientRequest(
types.CallToolRequest(
params=types.CallToolRequestParams(name="test_tool", arguments={}),
)
),
types.CallToolResult,
metadata=metadata,
)
assert result is not None
# Should have received priming event token BEFORE response data
# Priming event = 1 token (empty data, id only)
# Response = 1 token (actual JSON-RPC response)
# Total = 2 tokens minimum
assert len(captured_resumption_tokens) >= 2, (
f"Server must send priming event before response. "
f"Expected >= 2 tokens (priming + response), got {len(captured_resumption_tokens)}"
)
assert captured_resumption_tokens[0] is not None
@pytest.mark.anyio
async def test_server_close_sse_stream_via_context(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""Server tool can call ctx.close_sse_stream() to close connection."""
_, server_url = event_server
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
# Call tool that closes stream mid-operation
# This should NOT raise NotImplementedError when fully implemented
result = await session.call_tool("tool_with_stream_close", {})
# Client should still receive complete response (via auto-reconnect)
assert result is not None
assert len(result.content) > 0
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "Done"
@pytest.mark.anyio
async def test_streamablehttp_client_auto_reconnects(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""Client should auto-reconnect with Last-Event-ID when server closes after priming event."""
_, server_url = event_server
captured_notifications: list[str] = []
async def message_handler(
message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception,
) -> None:
if isinstance(message, Exception): # pragma: no branch
return # pragma: no cover
if isinstance(message, types.ServerNotification): # pragma: no branch
if isinstance(message.root, types.LoggingMessageNotification): # pragma: no branch
captured_notifications.append(str(message.root.params.data))
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(
read_stream,
write_stream,
message_handler=message_handler,
) as session:
await session.initialize()
# Call tool that:
# 1. Sends notification
# 2. Closes SSE stream
# 3. Sends more notifications (stored in event_store)
# 4. Returns response
result = await session.call_tool("tool_with_stream_close", {})
# Client should have auto-reconnected and received ALL notifications
assert len(captured_notifications) >= 2, (
"Client should auto-reconnect and receive notifications sent both before and after stream close"
)
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "Done"
@pytest.mark.anyio
async def test_streamablehttp_client_respects_retry_interval(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""Client MUST respect retry field, waiting specified ms before reconnecting."""
_, server_url = event_server
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
start_time = time.monotonic()
result = await session.call_tool("tool_with_stream_close", {})
elapsed = time.monotonic() - start_time
# Verify result was received
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "Done"
# The elapsed time should include at least the retry interval
# if reconnection occurred. This test may be flaky depending on
# implementation details, but demonstrates the expected behavior.
# Note: This assertion may need adjustment based on actual implementation
assert elapsed >= 0.4, f"Client should wait ~500ms before reconnecting, but elapsed time was {elapsed:.3f}s"
@pytest.mark.anyio
async def test_streamablehttp_sse_polling_full_cycle(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""End-to-end test: server closes stream, client reconnects, receives all events."""
_, server_url = event_server
all_notifications: list[str] = []
async def message_handler(
message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception,
) -> None:
if isinstance(message, Exception): # pragma: no branch
return # pragma: no cover
if isinstance(message, types.ServerNotification): # pragma: no branch
if isinstance(message.root, types.LoggingMessageNotification): # pragma: no branch
all_notifications.append(str(message.root.params.data))
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(
read_stream,
write_stream,
message_handler=message_handler,
) as session:
await session.initialize()
# Call tool that simulates polling pattern:
# 1. Server sends priming event
# 2. Server sends "Before close" notification
# 3. Server closes stream (calls close_sse_stream)
# 4. (client reconnects automatically)
# 5. Server sends "After close" notification
# 6. Server sends final response
result = await session.call_tool("tool_with_stream_close", {})
# Verify all notifications received in order
assert "Before close" in all_notifications, "Should receive notification sent before stream close"
assert "After close" in all_notifications, (
"Should receive notification sent after stream close (via auto-reconnect)"
)
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "Done"
@pytest.mark.anyio
async def test_streamablehttp_events_replayed_after_disconnect(
event_server: tuple[SimpleEventStore, str],
) -> None:
"""Events sent while client is disconnected should be replayed on reconnect."""
_, server_url = event_server
notification_data: list[str] = []
async def message_handler(
message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception,
) -> None:
if isinstance(message, Exception): # pragma: no branch
return # pragma: no cover
if isinstance(message, types.ServerNotification): # pragma: no branch
if isinstance(message.root, types.LoggingMessageNotification): # pragma: no branch
notification_data.append(str(message.root.params.data))
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(
read_stream,
write_stream,
message_handler=message_handler,
) as session:
await session.initialize()
# Tool sends: notification1, close_stream, notification2, notification3, response
# Client should receive all notifications even though 2&3 were sent during disconnect
result = await session.call_tool("tool_with_multiple_notifications_and_close", {})
assert "notification1" in notification_data, "Should receive notification1 (sent before close)"
assert "notification2" in notification_data, "Should receive notification2 (sent after close, replayed)"
assert "notification3" in notification_data, "Should receive notification3 (sent after close, replayed)"
# Verify order: notification1 should come before notification2 and notification3
idx1 = notification_data.index("notification1")
idx2 = notification_data.index("notification2")
idx3 = notification_data.index("notification3")
assert idx1 < idx2 < idx3, "Notifications should be received in order"
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "All notifications sent"
@pytest.mark.anyio
async def test_streamablehttp_multiple_reconnections(
event_server: tuple[SimpleEventStore, str],
):
"""Verify multiple close_sse_stream() calls each trigger a client reconnect.
Server uses retry_interval=500ms, tool sleeps 600ms after each close to ensure
client has time to reconnect before the next checkpoint.
With 3 checkpoints, we expect 8 resumption tokens:
- 1 priming (initial POST connection)
- 3 notifications (checkpoint_0, checkpoint_1, checkpoint_2)
- 3 priming (one per reconnect after each close)
- 1 response
"""
_, server_url = event_server
resumption_tokens: list[str] = []
async def on_resumption_token(token: str) -> None:
resumption_tokens.append(token)
async with streamablehttp_client(f"{server_url}/mcp") as (read_stream, write_stream, _):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
# Use send_request with metadata to track resumption tokens
metadata = ClientMessageMetadata(on_resumption_token_update=on_resumption_token)
result = await session.send_request(
types.ClientRequest(
types.CallToolRequest(
method="tools/call",
params=types.CallToolRequestParams(
name="tool_with_multiple_stream_closes",
# retry_interval=500ms, so sleep 600ms to ensure reconnect completes
arguments={"checkpoints": 3, "sleep_time": 0.6},
),
)
),
types.CallToolResult,
metadata=metadata,
)
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert "Completed 3 checkpoints" in result.content[0].text
# 4 priming + 3 notifications + 1 response = 8 tokens
assert len(resumption_tokens) == 8, ( # pragma: no cover
f"Expected 8 resumption tokens (4 priming + 3 notifs + 1 response), "
f"got {len(resumption_tokens)}: {resumption_tokens}"
)
@pytest.mark.anyio
async def test_standalone_get_stream_reconnection(basic_server: None, basic_server_url: str) -> None:
"""
Test that standalone GET stream automatically reconnects after server closes it.
Verifies:
1. Client receives notification 1 via GET stream
2. Server closes GET stream
3. Client reconnects with Last-Event-ID
4. Client receives notification 2 on new connection
"""
server_url = basic_server_url
received_notifications: list[str] = []
async def message_handler(
message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception,
) -> None:
if isinstance(message, Exception):
return # pragma: no cover
if isinstance(message, types.ServerNotification): # pragma: no branch
if isinstance(message.root, types.ResourceUpdatedNotification): # pragma: no branch
received_notifications.append(str(message.root.params.uri))
async with streamablehttp_client(f"{server_url}/mcp") as (
read_stream,
write_stream,
_,
):
async with ClientSession(
read_stream,
write_stream,
message_handler=message_handler,
) as session:
await session.initialize()
# Call tool that:
# 1. Sends notification_1 via GET stream
# 2. Closes standalone GET stream
# 3. Sends notification_2 (stored in event_store)
# 4. Returns response
result = await session.call_tool("tool_with_standalone_stream_close", {})
# Verify the tool completed
assert result.content[0].type == "text"
assert isinstance(result.content[0], TextContent)
assert result.content[0].text == "Standalone stream close test done"
# Verify both notifications were received
assert "http://notification_1/" in received_notifications, (
f"Should receive notification 1 (sent before GET stream close), got: {received_notifications}"
)
assert "http://notification_2/" in received_notifications, (
f"Should receive notification 2 after reconnect, got: {received_notifications}"
)
Generated
+68
View File
@@ -21,6 +21,8 @@ members = [
"mcp-simple-task-interactive-client",
"mcp-simple-tool",
"mcp-snippets",
"mcp-sse-polling-client",
"mcp-sse-polling-demo",
"mcp-structured-output-lowlevel",
]
@@ -1364,6 +1366,72 @@ dependencies = [
[package.metadata]
requires-dist = [{ name = "mcp", editable = "." }]
[[package]]
name = "mcp-sse-polling-client"
version = "0.1.0"
source = { editable = "examples/clients/sse-polling-client" }
dependencies = [
{ name = "click" },
{ name = "mcp" },
]
[package.dev-dependencies]
dev = [
{ name = "pyright" },
{ name = "pytest" },
{ name = "ruff" },
]
[package.metadata]
requires-dist = [
{ name = "click", specifier = ">=8.2.0" },
{ name = "mcp", editable = "." },
]
[package.metadata.requires-dev]
dev = [
{ name = "pyright", specifier = ">=1.1.378" },
{ name = "pytest", specifier = ">=8.3.3" },
{ name = "ruff", specifier = ">=0.6.9" },
]
[[package]]
name = "mcp-sse-polling-demo"
version = "0.1.0"
source = { editable = "examples/servers/sse-polling-demo" }
dependencies = [
{ name = "anyio" },
{ name = "click" },
{ name = "httpx" },
{ name = "mcp" },
{ name = "starlette" },
{ name = "uvicorn" },
]
[package.dev-dependencies]
dev = [
{ name = "pyright" },
{ name = "pytest" },
{ name = "ruff" },
]
[package.metadata]
requires-dist = [
{ name = "anyio", specifier = ">=4.5" },
{ name = "click", specifier = ">=8.2.0" },
{ name = "httpx", specifier = ">=0.27" },
{ name = "mcp", editable = "." },
{ name = "starlette" },
{ name = "uvicorn" },
]
[package.metadata.requires-dev]
dev = [
{ name = "pyright", specifier = ">=1.1.378" },
{ name = "pytest", specifier = ">=8.3.3" },
{ name = "ruff", specifier = ">=0.6.9" },
]
[[package]]
name = "mcp-structured-output-lowlevel"
version = "0.1.0"