11use std:: collections:: HashMap ;
2+ use std:: error:: Error as StdError ;
23use std:: ffi:: OsString ;
4+ use std:: fmt;
5+ use std:: future:: Future ;
36use std:: io;
47use std:: process:: Stdio ;
58use std:: sync:: Arc ;
@@ -24,6 +27,7 @@ use rmcp::service::{self};
2427use rmcp:: transport:: StreamableHttpClientTransport ;
2528use rmcp:: transport:: child_process:: TokioChildProcess ;
2629use rmcp:: transport:: streamable_http_client:: StreamableHttpClientTransportConfig ;
30+ use reqwest:: Error as ReqwestError ;
2731use tokio:: io:: AsyncBufReadExt ;
2832use tokio:: io:: BufReader ;
2933use tokio:: process:: Command ;
@@ -41,7 +45,11 @@ use crate::utils::run_with_timeout;
4145
4246enum 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
4755enum 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
5883pub 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+
232338fn 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
239345fn 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) ]
246365mod 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