Skip to content

Commit d328751

Browse files
authored
fix: align protocol version negotiation (#855)
* fix: align protocol version negotiation * ci: relax semver-checks to allow minor changes
1 parent cc66e30 commit d328751

3 files changed

Lines changed: 108 additions & 12 deletions

File tree

.github/workflows/ci.yml

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,7 @@ jobs:
8585
cargo semver-checks \
8686
--package rmcp \
8787
--baseline-rev ${{ github.event.pull_request.base.sha }} \
88+
--release-type minor \
8889
--only-explicit-features \
8990
--features default
9091
@@ -97,6 +98,7 @@ jobs:
9798
cargo semver-checks \
9899
--package rmcp \
99100
--baseline-rev ${{ github.event.pull_request.base.sha }} \
101+
--release-type minor \
100102
--only-explicit-features \
101103
--features "$FEATURES"
102104

crates/rmcp/src/service/server.rs

Lines changed: 25 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -69,6 +69,10 @@ pub enum ServerInitializeError {
6969
#[error("initialize failed: {0}")]
7070
InitializeFailed(ErrorData),
7171

72+
#[deprecated(
73+
since = "1.8.0",
74+
note = "Negotiation now falls back to the server-configured version. This variant is never constructed and will be removed in a future major release."
75+
)]
7276
#[error("unsupported protocol version: {0}")]
7377
UnsupportedProtocolVersion(ProtocolVersion),
7478

@@ -155,6 +159,23 @@ where
155159
}
156160
}
157161

162+
/// Echoes the client-requested version if known; otherwise returns `server_fallback`.
163+
fn negotiate_protocol_version(
164+
client_requested: &ProtocolVersion,
165+
server_fallback: ProtocolVersion,
166+
) -> ProtocolVersion {
167+
if ProtocolVersion::KNOWN_VERSIONS.contains(client_requested) {
168+
client_requested.clone()
169+
} else {
170+
tracing::warn!(
171+
client_requested = %client_requested,
172+
server_fallback = %server_fallback,
173+
"client requested unsupported protocol version; falling back to server default"
174+
);
175+
server_fallback
176+
}
177+
}
178+
158179
async fn serve_server_with_ct_inner<S, T>(
159180
service: S,
160181
transport: T,
@@ -227,16 +248,10 @@ where
227248
return Err(ServerInitializeError::InitializeFailed(e));
228249
}
229250
};
230-
let peer_protocol_version = peer_info.params.protocol_version.clone();
231-
let protocol_version = match peer_protocol_version
232-
.partial_cmp(&init_response.protocol_version)
233-
.ok_or(ServerInitializeError::UnsupportedProtocolVersion(
234-
peer_protocol_version,
235-
))? {
236-
std::cmp::Ordering::Less => peer_info.params.protocol_version.clone(),
237-
_ => init_response.protocol_version,
238-
};
239-
init_response.protocol_version = protocol_version;
251+
init_response.protocol_version = negotiate_protocol_version(
252+
&peer_info.params.protocol_version,
253+
init_response.protocol_version,
254+
);
240255
transport
241256
.send(ServerJsonRpcMessage::response(
242257
ServerResult::InitializeResult(init_response),

crates/rmcp/tests/test_server_initialization.rs

Lines changed: 81 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,11 @@ mod common;
44

55
use common::handlers::TestServer;
66
use rmcp::{
7-
ServiceExt,
8-
model::{ClientJsonRpcMessage, ServerJsonRpcMessage, ServerResult},
7+
ServerHandler, ServiceExt,
8+
model::{
9+
ClientJsonRpcMessage, ProtocolVersion, ServerCapabilities, ServerInfo,
10+
ServerJsonRpcMessage, ServerResult,
11+
},
912
transport::{IntoTransport, Transport},
1013
};
1114

@@ -220,6 +223,82 @@ async fn server_init_buffers_request_before_initialized() {
220223
result.unwrap().cancel().await.unwrap();
221224
}
222225

226+
fn init_request_with_version(v: &str) -> ClientJsonRpcMessage {
227+
msg(&format!(
228+
r#"{{
229+
"jsonrpc": "2.0",
230+
"id": 1,
231+
"method": "initialize",
232+
"params": {{
233+
"protocolVersion": "{v}",
234+
"capabilities": {{}},
235+
"clientInfo": {{ "name": "test-client", "version": "0.0.1" }}
236+
}}
237+
}}"#
238+
))
239+
}
240+
241+
async fn negotiate_version<H>(handler: H, client_version: &str) -> ProtocolVersion
242+
where
243+
H: ServerHandler + 'static,
244+
{
245+
let (server_transport, client_transport) = tokio::io::duplex(4096);
246+
let _server = tokio::spawn(async move { handler.serve(server_transport).await });
247+
let mut client = IntoTransport::<rmcp::RoleClient, _, _>::into_transport(client_transport);
248+
249+
client
250+
.send(init_request_with_version(client_version))
251+
.await
252+
.unwrap();
253+
let response = client.receive().await.unwrap();
254+
let ServerJsonRpcMessage::Response(r) = response else {
255+
panic!("expected initialize response, got {response:?}");
256+
};
257+
let ServerResult::InitializeResult(init) = r.result else {
258+
panic!("expected InitializeResult");
259+
};
260+
init.protocol_version
261+
}
262+
263+
#[tokio::test]
264+
async fn server_echoes_client_protocol_version_when_known_old() {
265+
let negotiated = negotiate_version(TestServer::new(), "2024-11-05").await;
266+
assert_eq!(negotiated, ProtocolVersion::V_2024_11_05);
267+
}
268+
269+
#[tokio::test]
270+
async fn server_echoes_client_protocol_version_when_latest() {
271+
let negotiated = negotiate_version(TestServer::new(), "2025-11-25").await;
272+
assert_eq!(negotiated, ProtocolVersion::LATEST);
273+
}
274+
275+
#[tokio::test]
276+
async fn server_falls_back_when_client_protocol_version_unknown() {
277+
let negotiated = negotiate_version(TestServer::new(), "2099-99-99").await;
278+
assert_eq!(negotiated, ProtocolVersion::LATEST);
279+
}
280+
281+
struct PinnedServer;
282+
283+
impl ServerHandler for PinnedServer {
284+
fn get_info(&self) -> ServerInfo {
285+
ServerInfo::new(ServerCapabilities::builder().build())
286+
.with_protocol_version(ProtocolVersion::V_2025_06_18)
287+
}
288+
}
289+
290+
#[tokio::test]
291+
async fn server_pinned_version_does_not_override_known_client_request() {
292+
let negotiated = negotiate_version(PinnedServer, "2025-11-25").await;
293+
assert_eq!(negotiated, ProtocolVersion::LATEST);
294+
}
295+
296+
#[tokio::test]
297+
async fn server_pinned_version_used_as_fallback_for_unknown_client_request() {
298+
let negotiated = negotiate_version(PinnedServer, "2099-99-99").await;
299+
assert_eq!(negotiated, ProtocolVersion::V_2025_06_18);
300+
}
301+
223302
// Server buffers multiple requests before initialized and processes them in order.
224303
#[tokio::test]
225304
async fn server_init_buffers_multiple_requests_before_initialized() {

0 commit comments

Comments
 (0)