Skip to content

Commit dc1de83

Browse files
committed
fix(rmcp-client): label phase errors, retry transient init (#378)
1 parent 427c24a commit dc1de83

3 files changed

Lines changed: 187 additions & 39 deletions

File tree

COMMIT_MESSAGE_ISSUE_378.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
fix(rmcp-client): label phase errors, retry transient init (#378)

PR_BODY_ISSUE_378.md

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,8 @@
1+
## Summary
2+
- add explicit phase context to initialize/list/call MCP failures so logs show where the request died
3+
- retry the Streamable HTTP initialize handshake once on obvious transient network errors before surfacing failure
4+
- cover the new helpers with unit tests for phase labeling and retry gating
5+
6+
## Testing
7+
- cargo test -p code-rmcp-client
8+
- ./build-fast.sh

code-rs/rmcp-client/src/rmcp_client.rs

Lines changed: 178 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,8 @@
11
use std::collections::HashMap;
2+
use std::error::Error as StdError;
23
use std::ffi::OsString;
4+
use std::fmt;
5+
use std::future::Future;
36
use std::io;
47
use std::process::Stdio;
58
use std::sync::Arc;
@@ -24,6 +27,7 @@ use rmcp::service::{self};
2427
use rmcp::transport::StreamableHttpClientTransport;
2528
use rmcp::transport::child_process::TokioChildProcess;
2629
use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig;
30+
use reqwest::Error as ReqwestError;
2731
use tokio::io::AsyncBufReadExt;
2832
use tokio::io::BufReader;
2933
use tokio::process::Command;
@@ -41,7 +45,11 @@ use crate::utils::run_with_timeout;
4145

4246
enum PendingTransport {
4347
ChildProcess(TokioChildProcess),
44-
StreamableHttp(StreamableHttpClientTransport<reqwest::Client>),
48+
StreamableHttp {
49+
transport: StreamableHttpClientTransport<reqwest::Client>,
50+
url: String,
51+
bearer_token: Option<String>,
52+
},
4553
}
4654

4755
enum ClientState {
@@ -53,6 +61,23 @@ enum ClientState {
5361
},
5462
}
5563

64+
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
65+
enum Phase {
66+
Initialize,
67+
ListTools,
68+
CallTool,
69+
}
70+
71+
impl Phase {
72+
fn as_str(self) -> &'static str {
73+
match self {
74+
Phase::Initialize => "initialize",
75+
Phase::ListTools => "list_tools",
76+
Phase::CallTool => "call_tool",
77+
}
78+
}
79+
}
80+
5681
/// MCP client implemented on top of the official `rmcp` SDK.
5782
/// https://github.com/modelcontextprotocol/rust-sdk
5883
pub struct RmcpClient {
@@ -105,16 +130,15 @@ impl RmcpClient {
105130
}
106131

107132
pub fn new_streamable_http_client(url: String, bearer_token: Option<String>) -> Result<Self> {
108-
let mut config = StreamableHttpClientTransportConfig::with_uri(url);
109-
if let Some(token) = bearer_token {
110-
config = config.auth_header(format!("Bearer {token}"));
111-
}
112-
113-
let transport = StreamableHttpClientTransport::from_config(config);
133+
let transport = build_streamable_http_transport(&url, bearer_token.as_deref());
114134

115135
Ok(Self {
116136
state: Mutex::new(ClientState::Connecting {
117-
transport: Some(PendingTransport::StreamableHttp(transport)),
137+
transport: Some(PendingTransport::StreamableHttp {
138+
transport,
139+
url,
140+
bearer_token,
141+
}),
118142
}),
119143
})
120144
}
@@ -126,52 +150,66 @@ impl RmcpClient {
126150
params: InitializeRequestParams,
127151
timeout: Option<Duration>,
128152
) -> Result<InitializeResult> {
129-
let transport = {
153+
let pending_transport = {
130154
let mut guard = self.state.lock().await;
131155
match &mut *guard {
132156
ClientState::Connecting { transport } => transport
133157
.take()
134158
.ok_or_else(|| anyhow!("client already initializing"))?,
135-
ClientState::Ready { .. } => {
136-
return Err(anyhow!("client already initialized"));
137-
}
159+
ClientState::Ready { .. } => return Err(anyhow!("client already initialized")),
138160
}
139161
};
140162

141-
let client_info = convert_to_rmcp::<_, InitializeRequestParam>(params.clone())?;
142-
let client_handler = LoggingClientHandler::new(client_info);
143-
let service_future = match transport {
163+
let service = match pending_transport {
144164
PendingTransport::ChildProcess(transport) => {
145-
service::serve_client(client_handler.clone(), transport).boxed()
165+
let client_info = convert_to_rmcp::<_, InitializeRequestParam>(params.clone())?;
166+
let client_handler = LoggingClientHandler::new(client_info);
167+
let service_future = service::serve_client(client_handler.clone(), transport).boxed();
168+
await_handshake(service_future, timeout)
169+
.await
170+
.map_err(|err| annotate_phase_error(Phase::Initialize, err))?
146171
}
147-
PendingTransport::StreamableHttp(transport) => {
148-
service::serve_client(client_handler, transport).boxed()
172+
PendingTransport::StreamableHttp {
173+
mut transport,
174+
url,
175+
bearer_token,
176+
} => {
177+
let mut attempt = 0;
178+
loop {
179+
let client_info = convert_to_rmcp::<_, InitializeRequestParam>(params.clone())?;
180+
let client_handler = LoggingClientHandler::new(client_info);
181+
let service_future = service::serve_client(client_handler.clone(), transport).boxed();
182+
match await_handshake(service_future, timeout).await {
183+
Ok(service) => break service,
184+
Err(err) => {
185+
let err = annotate_phase_error(Phase::Initialize, err);
186+
if should_retry_initialize(&err, attempt) {
187+
attempt += 1;
188+
time::sleep(Duration::from_millis(250)).await;
189+
transport = build_streamable_http_transport(&url, bearer_token.as_deref());
190+
continue;
191+
}
192+
return Err(err);
193+
}
194+
}
195+
}
149196
}
150197
};
151198

152-
let service = match timeout {
153-
Some(duration) => match time::timeout(duration, service_future).await {
154-
Ok(Ok(service)) => service,
155-
Ok(Err(err)) => return Err(handshake_failed_error(err)),
156-
Err(_) => return Err(handshake_timeout_error(duration)),
157-
},
158-
None => match service_future.await {
159-
Ok(service) => service,
160-
Err(err) => return Err(handshake_failed_error(err)),
161-
},
162-
};
163-
164199
let initialize_result_rmcp = service
165200
.peer()
166201
.peer_info()
167-
.ok_or_else(|| anyhow!("handshake succeeded but server info was missing"))?;
202+
.ok_or_else(|| annotate_phase_error(Phase::Initialize, anyhow!("handshake succeeded but server info was missing")))?;
168203
let initialize_result: InitializeResult = convert_to_mcp(initialize_result_rmcp)?;
169204

170205
if initialize_result.protocol_version != MCP_SCHEMA_VERSION {
171206
let reported_version = initialize_result.protocol_version.clone();
172-
return Err(anyhow!(
173-
"MCP server reported protocol version {reported_version}, but this client expects {}. Update either side so both speak the same schema.",
174-
MCP_SCHEMA_VERSION
207+
return Err(annotate_phase_error(
208+
Phase::Initialize,
209+
anyhow!(
210+
"MCP server reported protocol version {reported_version}, but this client expects {}. Update either side so both speak the same schema.",
211+
MCP_SCHEMA_VERSION
212+
),
175213
));
176214
}
177215

@@ -196,7 +234,9 @@ impl RmcpClient {
196234
.transpose()?;
197235

198236
let fut = service.list_tools(rmcp_params);
199-
let result = run_with_timeout(fut, timeout, "tools/list").await?;
237+
let result = run_with_timeout(fut, timeout, "tools/list")
238+
.await
239+
.map_err(|err| annotate_phase_error(Phase::ListTools, err))?;
200240
convert_to_mcp(result)
201241
}
202242

@@ -210,7 +250,9 @@ impl RmcpClient {
210250
let params = CallToolRequestParams { arguments, name };
211251
let rmcp_params: CallToolRequestParam = convert_to_rmcp(params)?;
212252
let fut = service.call_tool(rmcp_params);
213-
let rmcp_result = run_with_timeout(fut, timeout, "tools/call").await?;
253+
let rmcp_result = run_with_timeout(fut, timeout, "tools/call")
254+
.await
255+
.map_err(|err| annotate_phase_error(Phase::CallTool, err))?;
214256
convert_call_tool_result(rmcp_result)
215257
}
216258

@@ -229,6 +271,70 @@ impl RmcpClient {
229271
}
230272
}
231273

274+
async fn await_handshake<F, E>(
275+
future: F,
276+
timeout: Option<Duration>,
277+
) -> Result<RunningService<RoleClient, LoggingClientHandler>>
278+
where
279+
F: Future<
280+
Output = Result<
281+
RunningService<RoleClient, LoggingClientHandler>,
282+
E,
283+
>,
284+
>,
285+
E: Into<anyhow::Error>,
286+
{
287+
if let Some(duration) = timeout {
288+
match time::timeout(duration, future).await {
289+
Ok(Ok(service)) => Ok(service),
290+
Ok(Err(err)) => Err(handshake_failed_error(err)),
291+
Err(_) => Err(handshake_timeout_error(duration)),
292+
}
293+
} else {
294+
future.await.map_err(handshake_failed_error)
295+
}
296+
}
297+
298+
fn annotate_phase_error(phase: Phase, err: anyhow::Error) -> anyhow::Error {
299+
err.context(format!("phase={}", phase.as_str()))
300+
}
301+
302+
fn should_retry_initialize(err: &anyhow::Error, attempt: usize) -> bool {
303+
if attempt != 0 {
304+
return false;
305+
}
306+
307+
for source in err.chain() {
308+
if let Some(reqwest_err) = source.downcast_ref::<ReqwestError>() {
309+
if reqwest_err.is_timeout() || reqwest_err.is_connect() {
310+
return true;
311+
}
312+
}
313+
314+
if let Some(io_err) = source.downcast_ref::<io::Error>() {
315+
if matches!(
316+
io_err.kind(),
317+
io::ErrorKind::TimedOut | io::ErrorKind::ConnectionRefused
318+
) {
319+
return true;
320+
}
321+
}
322+
}
323+
324+
false
325+
}
326+
327+
fn build_streamable_http_transport(
328+
url: &str,
329+
bearer_token: Option<&str>,
330+
) -> StreamableHttpClientTransport<reqwest::Client> {
331+
let mut config = StreamableHttpClientTransportConfig::with_uri(url.to_string());
332+
if let Some(token) = bearer_token {
333+
config = config.auth_header(format!("Bearer {token}"));
334+
}
335+
StreamableHttpClientTransport::from_config(config)
336+
}
337+
232338
fn handshake_failed_error(err: impl Into<anyhow::Error>) -> anyhow::Error {
233339
let err = err.into();
234340
anyhow!(
@@ -237,14 +343,28 @@ fn handshake_failed_error(err: impl Into<anyhow::Error>) -> anyhow::Error {
237343
}
238344

239345
fn handshake_timeout_error(duration: Duration) -> anyhow::Error {
240-
anyhow!(
241-
"timed out handshaking with MCP server after {duration:?} (expected MCP schema version {MCP_SCHEMA_VERSION})"
242-
)
346+
anyhow!(HandshakeTimeoutError(duration))
243347
}
244348

349+
#[derive(Debug)]
350+
struct HandshakeTimeoutError(Duration);
351+
352+
impl fmt::Display for HandshakeTimeoutError {
353+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
354+
write!(
355+
f,
356+
"timed out awaiting MCP handshake after {:?}",
357+
self.0
358+
)
359+
}
360+
}
361+
362+
impl StdError for HandshakeTimeoutError {}
363+
245364
#[cfg(test)]
246365
mod tests {
247366
use super::*;
367+
use anyhow::anyhow;
248368

249369
#[test]
250370
fn mcp_schema_version_is_well_formed() {
@@ -257,4 +377,23 @@ mod tests {
257377
);
258378
assert!(parts.iter().all(|segment| !segment.trim().is_empty()));
259379
}
380+
381+
#[test]
382+
fn annotate_phase_error_adds_phase_label() {
383+
let err = annotate_phase_error(Phase::ListTools, anyhow!("boom"));
384+
let message = err.to_string();
385+
assert_eq!(message, "phase=list_tools");
386+
let sources: Vec<String> = err.chain().map(|source| source.to_string()).collect();
387+
assert!(sources.iter().any(|s| s.contains("boom")), "sources: {sources:?}");
388+
}
389+
390+
#[test]
391+
fn should_retry_initialize_detects_transient_errors() {
392+
let timeout_err = anyhow!(io::Error::new(io::ErrorKind::TimedOut, "timed out"));
393+
assert!(should_retry_initialize(&timeout_err, 0));
394+
assert!(!should_retry_initialize(&timeout_err, 1));
395+
396+
let mismatch_err = anyhow!("protocol mismatch");
397+
assert!(!should_retry_initialize(&mismatch_err, 0));
398+
}
260399
}

0 commit comments

Comments
 (0)