Skip to content

Commit 80a7479

Browse files
authored
fix: negotiate protocol version in handler (#930)
* fix: negotiate protocol version in handler (fixes #916) * fix: use server pinned version as fallback in default initialize handler
1 parent 67a3085 commit 80a7479

6 files changed

Lines changed: 247 additions & 6 deletions

File tree

crates/rmcp/Cargo.toml

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -283,6 +283,16 @@ name = "test_streamable_http_protocol_version"
283283
required-features = ["server", "client", "transport-streamable-http-server", "reqwest"]
284284
path = "tests/test_streamable_http_protocol_version.rs"
285285

286+
[[test]]
287+
name = "test_stateless_protocol_version"
288+
required-features = ["server", "transport-streamable-http-server", "reqwest"]
289+
path = "tests/test_stateless_protocol_version.rs"
290+
291+
[[test]]
292+
name = "test_protocol_version_negotiation"
293+
required-features = ["server", "client"]
294+
path = "tests/test_protocol_version_negotiation.rs"
295+
286296
[[test]]
287297
name = "test_streamable_http_4xx_error_body"
288298
required-features = ["transport-streamable-http-client", "transport-streamable-http-client-reqwest"]

crates/rmcp/src/handler/server.rs

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ use crate::{
77
model::{TaskSupport, *},
88
service::{
99
MaybeSendFuture, NotificationContext, RequestContext, RoleServer, Service, ServiceRole,
10+
negotiate_protocol_version,
1011
},
1112
};
1213

@@ -202,8 +203,13 @@ macro_rules! server_handler_methods {
202203
request: InitializeRequestParams,
203204
context: RequestContext<RoleServer>,
204205
) -> impl Future<Output = Result<InitializeResult, McpError>> + MaybeSendFuture + '_ {
205-
context.peer.set_peer_info(request);
206-
std::future::ready(Ok(self.get_info()))
206+
context.peer.set_peer_info(request.clone());
207+
let mut info = self.get_info();
208+
info.protocol_version = negotiate_protocol_version(
209+
&request.protocol_version,
210+
info.protocol_version,
211+
);
212+
std::future::ready(Ok(info))
207213
}
208214
fn complete(
209215
&self,

crates/rmcp/src/service/server.rs

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,7 @@ where
162162
}
163163

164164
/// Echoes the client-requested version if known; otherwise returns `server_fallback`.
165-
fn negotiate_protocol_version(
165+
pub(crate) fn negotiate_protocol_version(
166166
client_requested: &ProtocolVersion,
167167
server_fallback: ProtocolVersion,
168168
) -> ProtocolVersion {
@@ -254,6 +254,11 @@ where
254254
&peer_info.params.protocol_version,
255255
init_response.protocol_version,
256256
);
257+
// Update peer_info so context.protocol_version() reflects the negotiated
258+
// version in all subsequent request handlers.
259+
let mut negotiated_peer_info = peer_info.params.clone();
260+
negotiated_peer_info.protocol_version = init_response.protocol_version.clone();
261+
peer.set_peer_info(negotiated_peer_info);
257262
transport
258263
.send(ServerJsonRpcMessage::response(
259264
ServerResult::InitializeResult(init_response),

crates/rmcp/src/transport/streamable_http_server/tower.rs

Lines changed: 41 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,9 @@ use super::session::{
1616
use crate::{
1717
RoleServer,
1818
model::{
19-
ClientJsonRpcMessage, ClientNotification, ClientRequest, ErrorData, GetExtensions,
20-
InitializeRequest, InitializedNotification, JsonRpcError, ProtocolVersion, RequestId,
19+
ClientCapabilities, ClientJsonRpcMessage, ClientNotification, ClientRequest, ErrorData,
20+
GetExtensions, Implementation, InitializeRequest, InitializeRequestParams,
21+
InitializedNotification, JsonRpcError, ProtocolVersion, RequestId,
2122
},
2223
serve_server,
2324
service::serve_directly,
@@ -1239,10 +1240,17 @@ where
12391240
.map_err(internal_error_response("get service"))?;
12401241
match message {
12411242
ClientJsonRpcMessage::Request(mut request) => {
1243+
// Build a peer_info so context.protocol_version() works inside handlers.
1244+
// serve_directly skips the handshake and receives None by default, making
1245+
// protocol_version() always return None in stateless mode. We reconstruct it:
1246+
// - initialize requests: version comes from the request body params
1247+
// - all other requests: version comes from the MCP-Protocol-Version header
1248+
// (already validated above; absent header defaults to 2025-03-26)
1249+
let peer_info = Self::peer_info_for_stateless_request(&request, &part.headers);
12421250
request.request.extensions_mut().insert(part);
12431251
let (transport, mut receiver) =
12441252
OneshotTransport::<RoleServer>::new(ClientJsonRpcMessage::Request(request));
1245-
let service = serve_directly(service, transport, None);
1253+
let service = serve_directly(service, transport, peer_info);
12461254
tokio::spawn(async move {
12471255
// on service created
12481256
let _ = service.waiting().await;
@@ -1331,4 +1339,34 @@ where
13311339
}
13321340
Ok(accepted_response())
13331341
}
1342+
1343+
/// Build a `ClientInfo` (peer_info) for a stateless request so that
1344+
/// `context.protocol_version()` returns the correct value inside handlers.
1345+
///
1346+
/// `serve_directly` skips the MCP handshake and accepts `peer_info = None`,
1347+
/// which means `context.protocol_version()` is always `None` in stateless mode.
1348+
/// We reconstruct the protocol version from the available signal per request type:
1349+
/// - initialize: version is in the request body params (authoritative)
1350+
/// - all other requests: version is in the MCP-Protocol-Version header
1351+
/// (validated before this point; absent header defaults to 2025-03-26)
1352+
fn peer_info_for_stateless_request(
1353+
request: &crate::model::JsonRpcRequest<ClientRequest>,
1354+
headers: &HeaderMap,
1355+
) -> Option<InitializeRequestParams> {
1356+
let version = if let ClientRequest::InitializeRequest(ref init) = request.request {
1357+
init.params.protocol_version.clone()
1358+
} else {
1359+
headers
1360+
.get(HEADER_MCP_PROTOCOL_VERSION)
1361+
.and_then(|v| v.to_str().ok())
1362+
.and_then(|s| serde_json::from_value(serde_json::Value::String(s.to_owned())).ok())
1363+
.unwrap_or(ProtocolVersion::V_2025_03_26)
1364+
};
1365+
Some(InitializeRequestParams {
1366+
meta: None,
1367+
protocol_version: version,
1368+
capabilities: ClientCapabilities::default(),
1369+
client_info: Implementation::default(),
1370+
})
1371+
}
13341372
}
Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,83 @@
1+
//! Tests for protocol version negotiation in the default ServerHandler::initialize impl.
2+
//!
3+
//! Known versions are echoed back; unknown versions fall back to LATEST.
4+
#![cfg(not(feature = "local"))]
5+
#![cfg(feature = "client")]
6+
7+
use rmcp::{
8+
ClientHandler, ServerHandler, ServiceExt,
9+
model::{ClientInfo, ProtocolVersion, ServerInfo},
10+
};
11+
12+
#[derive(Debug, Clone, Default)]
13+
struct EchoServer;
14+
15+
impl ServerHandler for EchoServer {
16+
fn get_info(&self) -> ServerInfo {
17+
ServerInfo::default()
18+
}
19+
}
20+
21+
#[derive(Debug, Clone)]
22+
struct VersionedClient {
23+
protocol_version: ProtocolVersion,
24+
}
25+
26+
impl ClientHandler for VersionedClient {
27+
fn get_info(&self) -> ClientInfo {
28+
let mut info = ClientInfo::default();
29+
info.protocol_version = self.protocol_version.clone();
30+
info
31+
}
32+
}
33+
34+
async fn negotiated_version(client_version: ProtocolVersion) -> ProtocolVersion {
35+
let (server_transport, client_transport) = tokio::io::duplex(4096);
36+
37+
tokio::spawn(async move {
38+
let _ = EchoServer
39+
.serve(server_transport)
40+
.await
41+
.expect("server should start")
42+
.waiting()
43+
.await;
44+
});
45+
46+
let client = VersionedClient {
47+
protocol_version: client_version,
48+
}
49+
.serve(client_transport)
50+
.await
51+
.expect("client should connect");
52+
53+
let version = client
54+
.peer_info()
55+
.expect("peer_info should be set")
56+
.protocol_version
57+
.clone();
58+
59+
client.cancel().await.expect("client should cancel");
60+
version
61+
}
62+
63+
#[tokio::test]
64+
async fn known_version_echoed_back() {
65+
for version in ProtocolVersion::KNOWN_VERSIONS {
66+
let negotiated = negotiated_version(version.clone()).await;
67+
assert_eq!(
68+
negotiated, *version,
69+
"known version {version} should be echoed back"
70+
);
71+
}
72+
}
73+
74+
#[tokio::test]
75+
async fn unknown_version_falls_back_to_latest() {
76+
let unknown: ProtocolVersion = serde_json::from_str(r#""1999-01-01""#).unwrap();
77+
let negotiated = negotiated_version(unknown).await;
78+
assert_eq!(
79+
negotiated,
80+
ProtocolVersion::LATEST,
81+
"unknown version should fall back to LATEST"
82+
);
83+
}
Lines changed: 99 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
1+
//! Tests for protocol version negotiation in stateless HTTP mode.
2+
//!
3+
//! Known versions are echoed back; unknown versions fall back to LATEST.
4+
#![cfg(not(feature = "local"))]
5+
6+
use rmcp::{
7+
model::ProtocolVersion,
8+
transport::streamable_http_server::{
9+
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
10+
},
11+
};
12+
use tokio_util::sync::CancellationToken;
13+
14+
mod common;
15+
use common::calculator::Calculator;
16+
17+
fn stateless_json_config() -> StreamableHttpServerConfig {
18+
StreamableHttpServerConfig::default()
19+
.with_stateful_mode(false)
20+
.with_json_response(true)
21+
.with_sse_keep_alive(None)
22+
.with_cancellation_token(CancellationToken::new())
23+
}
24+
25+
async fn spawn_server(
26+
config: StreamableHttpServerConfig,
27+
) -> (reqwest::Client, String, CancellationToken) {
28+
let ct = config.cancellation_token.clone();
29+
let service: StreamableHttpService<Calculator, LocalSessionManager> =
30+
StreamableHttpService::new(|| Ok(Calculator::new()), Default::default(), config);
31+
32+
let router = axum::Router::new().nest_service("/mcp", service);
33+
let tcp_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
34+
let addr = tcp_listener.local_addr().unwrap();
35+
36+
tokio::spawn({
37+
let ct = ct.clone();
38+
async move {
39+
let _ = axum::serve(tcp_listener, router)
40+
.with_graceful_shutdown(async move { ct.cancelled_owned().await })
41+
.await;
42+
}
43+
});
44+
45+
(reqwest::Client::new(), format!("http://{addr}/mcp"), ct)
46+
}
47+
48+
async fn post_init(client: &reqwest::Client, url: &str, body_version: &str) -> serde_json::Value {
49+
let body = serde_json::json!({
50+
"jsonrpc": "2.0",
51+
"id": 1,
52+
"method": "initialize",
53+
"params": {
54+
"protocolVersion": body_version,
55+
"capabilities": {},
56+
"clientInfo": {"name": "test", "version": "0.0.1"}
57+
}
58+
});
59+
let resp = client
60+
.post(url)
61+
.header("Content-Type", "application/json")
62+
.header("Accept", "application/json, text/event-stream")
63+
.body(body.to_string())
64+
.send()
65+
.await
66+
.expect("send request");
67+
assert!(resp.status().is_success(), "HTTP {}", resp.status());
68+
resp.json().await.expect("parse JSON")
69+
}
70+
71+
#[tokio::test]
72+
async fn stateless_init_echoes_known_version() {
73+
let (client, url, ct) = spawn_server(stateless_json_config()).await;
74+
75+
for version in ProtocolVersion::KNOWN_VERSIONS {
76+
let resp = post_init(&client, &url, version.as_str()).await;
77+
assert_eq!(
78+
resp["result"]["protocolVersion"],
79+
version.as_str(),
80+
"known version {version} should be echoed back"
81+
);
82+
}
83+
84+
ct.cancel();
85+
}
86+
87+
#[tokio::test]
88+
async fn stateless_init_unknown_version_falls_back_to_latest() {
89+
let (client, url, ct) = spawn_server(stateless_json_config()).await;
90+
91+
let resp = post_init(&client, &url, "1999-01-01").await;
92+
assert_eq!(
93+
resp["result"]["protocolVersion"],
94+
ProtocolVersion::LATEST.as_str(),
95+
"unknown version should fall back to LATEST"
96+
);
97+
98+
ct.cancel();
99+
}

0 commit comments

Comments
 (0)