Skip to content

Commit 8803d39

Browse files
authored
fix: accept stringified numeric response IDs (#1021)
1 parent 99ee024 commit 8803d39

6 files changed

Lines changed: 214 additions & 9 deletions

File tree

crates/rmcp/src/model.rs

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -240,6 +240,22 @@ impl NumberOrString {
240240
NumberOrString::String(s) => Value::String(s.to_string()),
241241
}
242242
}
243+
244+
pub(crate) fn numeric_string_value(&self) -> Option<i64> {
245+
match self {
246+
Self::String(id) => id.parse().ok(),
247+
Self::Number(_) => None,
248+
}
249+
}
250+
251+
pub(crate) fn matches_response_id(&self, response_id: &Self) -> bool {
252+
self == response_id
253+
|| matches!(
254+
self,
255+
Self::Number(request_id)
256+
if response_id.numeric_string_value() == Some(*request_id)
257+
)
258+
}
243259
}
244260

245261
impl std::fmt::Display for NumberOrString {

crates/rmcp/src/service.rs

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -305,6 +305,17 @@ pub trait ProgressTokenProvider: Send + Sync + 'static {
305305
pub type AtomicU32RequestIdProvider = AtomicU32Provider;
306306
pub type AtomicU32ProgressTokenProvider = AtomicU32Provider;
307307

308+
pub(crate) fn remove_pending_request<T>(
309+
pending_requests: &mut HashMap<RequestId, T>,
310+
response_id: &RequestId,
311+
) -> Option<T> {
312+
pending_requests.remove(response_id).or_else(|| {
313+
response_id
314+
.numeric_string_value()
315+
.and_then(|id| pending_requests.remove(&RequestId::Number(id)))
316+
})
317+
}
318+
308319
#[derive(Debug, Default)]
309320
pub struct AtomicU32Provider {
310321
id: AtomicU64,
@@ -1481,7 +1492,9 @@ where
14811492
id,
14821493
..
14831494
})) => {
1484-
if let Some(responder) = local_responder_pool.remove(&id) {
1495+
if let Some(responder) =
1496+
remove_pending_request(&mut local_responder_pool, &id)
1497+
{
14851498
let response_result = responder.send(Ok(result));
14861499
if let Err(_error) = response_result {
14871500
tracing::warn!(%id, "Error sending response");
@@ -1495,7 +1508,9 @@ where
14951508
tracing::debug!(?error, "received id-less peer error");
14961509
continue;
14971510
};
1498-
if let Some(responder) = local_responder_pool.remove(&id) {
1511+
if let Some(responder) =
1512+
remove_pending_request(&mut local_responder_pool, &id)
1513+
{
14991514
let service_error = if error.is_transport_closed() {
15001515
ServiceError::TransportClosed
15011516
} else {

crates/rmcp/src/service/client.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -687,7 +687,7 @@ where
687687
let (response, response_id) =
688688
expect_response(transport, "initialize response", service, peer.clone()).await?;
689689

690-
if id != response_id {
690+
if !id.matches_response_id(&response_id) {
691691
return Err(ClientInitializeError::ConflictInitResponseId(
692692
id,
693693
response_id,
@@ -753,7 +753,7 @@ where
753753

754754
match expect_response(transport, "discover response", service, peer.clone()).await {
755755
Ok((ServerResult::DiscoverResult(result), response_id)) => {
756-
if response_id != id {
756+
if !id.matches_response_id(&response_id) {
757757
return Err(ClientInitializeError::ConflictInitResponseId(
758758
id,
759759
response_id,

crates/rmcp/src/transport/streamable_http_client.rs

Lines changed: 47 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -499,8 +499,14 @@ impl<C: StreamableHttpClient> StreamableHttpClientWorker<C> {
499499
pending_stream_response_ids: &mut HashSet<RequestId>,
500500
message: &ServerJsonRpcMessage,
501501
) {
502-
if let Some(id) = Self::server_response_id(message) {
503-
pending_stream_response_ids.remove(id);
502+
let Some(response_id) = Self::server_response_id(message) else {
503+
return;
504+
};
505+
if pending_stream_response_ids.remove(response_id) {
506+
return;
507+
}
508+
if let Some(id) = response_id.numeric_string_value() {
509+
pending_stream_response_ids.remove(&RequestId::Number(id));
504510
}
505511
}
506512

@@ -1382,7 +1388,10 @@ impl<C: StreamableHttpClient> Worker for StreamableHttpClientWorker<C> {
13821388
}
13831389
Event::ServerMessage(mut json_rpc_message) => {
13841390
if let Some(response_id) = Self::server_response_id(&json_rpc_message)
1385-
&& let Some(stream_ct) = request_stream_cancellations.remove(response_id)
1391+
&& let Some(stream_ct) = crate::service::remove_pending_request(
1392+
&mut request_stream_cancellations,
1393+
response_id,
1394+
)
13861395
{
13871396
stream_ct.cancel();
13881397
}
@@ -1848,4 +1857,39 @@ mod tests {
18481857
vec!["legacy"]
18491858
);
18501859
}
1860+
1861+
#[cfg(feature = "transport-streamable-http-client-reqwest")]
1862+
#[test]
1863+
fn clear_stream_response_pending_accepts_stringified_numeric_id() {
1864+
let mut pending = HashSet::from([NumberOrString::Number(1)]);
1865+
let response = ServerJsonRpcMessage::response(
1866+
ServerResult::ListToolsResult(ListToolsResult::default()),
1867+
NumberOrString::String("1".into()),
1868+
);
1869+
1870+
StreamableHttpClientWorker::<reqwest::Client>::clear_stream_response_pending(
1871+
&mut pending,
1872+
&response,
1873+
);
1874+
1875+
assert!(pending.is_empty());
1876+
}
1877+
1878+
#[cfg(feature = "transport-streamable-http-client-reqwest")]
1879+
#[test]
1880+
fn clear_stream_response_pending_prefers_exact_string_id() {
1881+
let string_id = NumberOrString::String("1".into());
1882+
let mut pending = HashSet::from([NumberOrString::Number(1), string_id.clone()]);
1883+
let response = ServerJsonRpcMessage::response(
1884+
ServerResult::ListToolsResult(ListToolsResult::default()),
1885+
string_id,
1886+
);
1887+
1888+
StreamableHttpClientWorker::<reqwest::Client>::clear_stream_response_pending(
1889+
&mut pending,
1890+
&response,
1891+
);
1892+
1893+
assert_eq!(pending, HashSet::from([NumberOrString::Number(1)]));
1894+
}
18511895
}

crates/rmcp/tests/test_client_initialization.rs

Lines changed: 92 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,11 +9,102 @@ use common::handlers::TestClientHandler;
99
use rmcp::{
1010
ServiceExt,
1111
model::{
12-
ErrorCode, ErrorData, JsonRpcError, JsonRpcVersion2_0, RequestId, ServerJsonRpcMessage,
12+
ClientJsonRpcMessage, ErrorCode, ErrorData, InitializeResult, JsonRpcError,
13+
JsonRpcVersion2_0, RequestId, ServerCapabilities, ServerJsonRpcMessage, ServerResult,
1314
},
1415
transport::{IntoTransport, Transport},
1516
};
1617

18+
fn stringify_numeric_id(id: RequestId) -> RequestId {
19+
let RequestId::Number(id) = id else {
20+
panic!("expected a numeric request ID");
21+
};
22+
RequestId::String(id.to_string().into())
23+
}
24+
25+
#[tokio::test]
26+
async fn client_initialization_accepts_stringified_numeric_response_id() {
27+
let (server_transport, client_transport) = tokio::io::duplex(1024);
28+
let mut server = IntoTransport::<rmcp::RoleServer, _, _>::into_transport(server_transport);
29+
let server_task = tokio::spawn(async move {
30+
let ClientJsonRpcMessage::Request(request) =
31+
server.receive().await.expect("expected initialize request")
32+
else {
33+
panic!("expected initialize request");
34+
};
35+
server
36+
.send(ServerJsonRpcMessage::response(
37+
ServerResult::InitializeResult(
38+
InitializeResult::new(ServerCapabilities::default()),
39+
),
40+
stringify_numeric_id(request.id),
41+
))
42+
.await
43+
.expect("send initialize response");
44+
assert!(matches!(
45+
server.receive().await,
46+
Some(ClientJsonRpcMessage::Notification(_))
47+
));
48+
});
49+
50+
let client = TestClientHandler::new(true, true)
51+
.serve(client_transport)
52+
.await
53+
.expect("client should accept stringified initialize response ID");
54+
client.cancel().await.expect("cancel client");
55+
server_task.await.expect("server task");
56+
}
57+
58+
#[tokio::test]
59+
async fn client_correlates_stringified_numeric_response_id() {
60+
let (server_transport, client_transport) = tokio::io::duplex(1024);
61+
let mut server = IntoTransport::<rmcp::RoleServer, _, _>::into_transport(server_transport);
62+
let server_task = tokio::spawn(async move {
63+
let ClientJsonRpcMessage::Request(initialize) =
64+
server.receive().await.expect("expected initialize request")
65+
else {
66+
panic!("expected initialize request");
67+
};
68+
server
69+
.send(ServerJsonRpcMessage::response(
70+
ServerResult::InitializeResult(
71+
InitializeResult::new(ServerCapabilities::default()),
72+
),
73+
initialize.id,
74+
))
75+
.await
76+
.expect("send initialize response");
77+
assert!(matches!(
78+
server.receive().await,
79+
Some(ClientJsonRpcMessage::Notification(_))
80+
));
81+
82+
let ClientJsonRpcMessage::Request(request) =
83+
server.receive().await.expect("expected tools/list request")
84+
else {
85+
panic!("expected tools/list request");
86+
};
87+
server
88+
.send(ServerJsonRpcMessage::response(
89+
ServerResult::ListToolsResult(Default::default()),
90+
stringify_numeric_id(request.id),
91+
))
92+
.await
93+
.expect("send tools/list response");
94+
});
95+
96+
let client = TestClientHandler::new(true, true)
97+
.serve(client_transport)
98+
.await
99+
.expect("initialize client");
100+
client
101+
.list_tools(None)
102+
.await
103+
.expect("client should correlate stringified response ID");
104+
client.cancel().await.expect("cancel client");
105+
server_task.await.expect("server task");
106+
}
107+
17108
#[tokio::test]
18109
async fn test_client_init_handles_jsonrpc_error() {
19110
let (server_transport, client_transport) = tokio::io::duplex(1024);

crates/rmcp/tests/test_client_lifecycle_modes.rs

Lines changed: 40 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ use rmcp::{
44
ClientHandler, ClientLifecycleMode, ClientServiceExt, ServerHandler, ServiceExt,
55
model::{
66
ClientJsonRpcMessage, ClientRequest, DiscoverResult, ErrorCode, ErrorData, GetMeta,
7-
Implementation, InitializeResult, ProtocolVersion, ServerCapabilities,
7+
Implementation, InitializeResult, ProtocolVersion, RequestId, ServerCapabilities,
88
ServerJsonRpcMessage, ServerResult,
99
},
1010
service::PeerRequestOptions,
@@ -21,6 +21,45 @@ struct StatelessServer;
2121

2222
impl ServerHandler for StatelessServer {}
2323

24+
#[tokio::test]
25+
async fn discover_startup_accepts_stringified_numeric_response_id() {
26+
let (server_transport, client_transport) = tokio::io::duplex(4096);
27+
let mut server = IntoTransport::<rmcp::RoleServer, _, _>::into_transport(server_transport);
28+
let server_task = tokio::spawn(async move {
29+
let ClientJsonRpcMessage::Request(request) =
30+
server.receive().await.expect("expected discover request")
31+
else {
32+
panic!("expected discover request");
33+
};
34+
let RequestId::Number(response_id) = request.id else {
35+
panic!("expected a numeric request ID");
36+
};
37+
server
38+
.send(ServerJsonRpcMessage::response(
39+
ServerResult::DiscoverResult(DiscoverResult::new(
40+
vec![ProtocolVersion::V_2026_07_28],
41+
ServerCapabilities::default(),
42+
Implementation::new("discover-server", "1.0.0"),
43+
)),
44+
RequestId::String(response_id.to_string().into()),
45+
))
46+
.await
47+
.expect("send discover response");
48+
});
49+
50+
let client = DiscoverClient
51+
.serve_with_lifecycle(
52+
client_transport,
53+
ClientLifecycleMode::Discover {
54+
preferred_versions: vec![ProtocolVersion::V_2026_07_28],
55+
},
56+
)
57+
.await
58+
.expect("client should accept stringified discover response ID");
59+
client.cancel().await.expect("cancel client");
60+
server_task.await.expect("server task");
61+
}
62+
2463
#[tokio::test]
2564
async fn high_level_server_accepts_discover_startup_without_initialize() {
2665
let (server_transport, client_transport) = tokio::io::duplex(4096);

0 commit comments

Comments
 (0)