Add SSE polling support (SEP-1699) (#1654)
This commit is contained in:
@@ -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"]
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user