Skip to content

Commit f754b75

Browse files
fix: conservatively retry legacy server initialization (#1040)
1 parent e660b80 commit f754b75

5 files changed

Lines changed: 807 additions & 30 deletions

File tree

crates/rmcp/src/service/client.rs

Lines changed: 212 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -102,13 +102,18 @@ where
102102
.ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string()))
103103
}
104104

105+
enum StartupResponse {
106+
Response(Box<ServerResult>, RequestId),
107+
Error(ErrorData, Option<RequestId>),
108+
}
109+
105110
/// Helper function to expect a response from the stream
106111
async fn expect_response<T, S>(
107112
transport: &mut T,
108113
context: &str,
109114
service: &S,
110115
peer: Peer<RoleClient>,
111-
) -> Result<(ServerResult, RequestId), ClientInitializeError>
116+
) -> Result<StartupResponse, ClientInitializeError>
112117
where
113118
T: Transport<RoleClient>,
114119
S: Service<RoleClient>,
@@ -118,11 +123,11 @@ where
118123
match message {
119124
// Expected message to complete the initialization
120125
ServerJsonRpcMessage::Response(JsonRpcResponse { id, result, .. }) => {
121-
break Ok((result, id));
126+
break Ok(StartupResponse::Response(Box::new(result), id));
122127
}
123128
// Handle JSON-RPC error responses
124129
ServerJsonRpcMessage::Error(error) => {
125-
break Err(ClientInitializeError::JsonRpcError(error.error));
130+
break Ok(StartupResponse::Error(error.error, error.id));
126131
}
127132
// Server could send logging messages before handshake
128133
ServerJsonRpcMessage::Notification(mut notification) => {
@@ -550,6 +555,159 @@ pub enum ClientLifecycleMode {
550555
},
551556
}
552557

558+
#[derive(Debug)]
559+
struct DiscoverStartupError {
560+
error: ClientInitializeError,
561+
requested_version: Option<ProtocolVersion>,
562+
request_id: Option<RequestId>,
563+
response_id: Option<RequestId>,
564+
}
565+
566+
impl From<ClientInitializeError> for DiscoverStartupError {
567+
fn from(error: ClientInitializeError) -> Self {
568+
Self {
569+
error,
570+
requested_version: None,
571+
request_id: None,
572+
response_id: None,
573+
}
574+
}
575+
}
576+
577+
impl DiscoverStartupError {
578+
fn json_rpc(
579+
error: ErrorData,
580+
requested_version: ProtocolVersion,
581+
request_id: RequestId,
582+
response_id: Option<RequestId>,
583+
) -> Self {
584+
Self {
585+
error: ClientInitializeError::JsonRpcError(error),
586+
requested_version: Some(requested_version),
587+
request_id: Some(request_id),
588+
response_id,
589+
}
590+
}
591+
592+
fn is_legacy_server(&self) -> bool {
593+
match &self.error {
594+
ClientInitializeError::NoCompatibleProtocolVersion {
595+
server_supported, ..
596+
} => exclusively_historical_protocol_versions(server_supported),
597+
ClientInitializeError::JsonRpcError(error) => {
598+
let correlated = self
599+
.request_id
600+
.as_ref()
601+
.zip(self.response_id.as_ref())
602+
.is_some_and(|(request_id, response_id)| {
603+
request_id.matches_response_id(response_id)
604+
});
605+
606+
if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND {
607+
return correlated;
608+
}
609+
610+
if error.code != crate::model::ErrorCode(-32000)
611+
|| self.request_id.is_none()
612+
|| (!correlated && self.response_id.is_some())
613+
{
614+
return false;
615+
}
616+
617+
let message = error.message.trim().to_ascii_lowercase();
618+
let normalized_message = message.strip_prefix("bad request: ").unwrap_or(&message);
619+
if normalized_message == "no valid session id provided" {
620+
return true;
621+
}
622+
623+
let Some(requested_version) = &self.requested_version else {
624+
return false;
625+
};
626+
if !normalized_message.contains("unsupported protocol version")
627+
|| !normalized_message.contains(requested_version.as_str())
628+
{
629+
return false;
630+
}
631+
632+
let supported_data = error.data.as_ref().and_then(|data| data.get("supported"));
633+
let supported_from_data = match supported_data {
634+
Some(value) => {
635+
let Ok(versions) =
636+
serde_json::from_value::<Vec<ProtocolVersion>>(value.clone())
637+
else {
638+
return false;
639+
};
640+
Some(versions)
641+
}
642+
None => None,
643+
};
644+
let supported_from_message = historical_versions_from_message(normalized_message);
645+
if normalized_message.contains("supported versions:")
646+
&& supported_from_message.is_none()
647+
{
648+
return false;
649+
}
650+
651+
let supported = match (supported_from_data, supported_from_message) {
652+
(Some(data), Some(message))
653+
if data.len() == message.len()
654+
&& data.iter().all(|version| message.contains(version)) =>
655+
{
656+
Some(data)
657+
}
658+
(Some(_), Some(_)) => None,
659+
(Some(versions), None) | (None, Some(versions)) => Some(versions),
660+
(None, None) => None,
661+
};
662+
663+
supported
664+
.as_deref()
665+
.is_some_and(exclusively_historical_protocol_versions)
666+
}
667+
_ => false,
668+
}
669+
}
670+
}
671+
672+
fn exclusively_historical_protocol_versions(versions: &[ProtocolVersion]) -> bool {
673+
!versions.is_empty()
674+
&& versions.iter().all(|version| {
675+
(ProtocolVersion::KNOWN_VERSIONS.contains(version)
676+
&& version < &ProtocolVersion::V_2026_07_28)
677+
// Some deployed legacy servers also advertise this pre-release version.
678+
|| version.as_str() == "2024-10-07"
679+
})
680+
}
681+
682+
fn historical_versions_from_message(message: &str) -> Option<Vec<ProtocolVersion>> {
683+
let (_, supported_versions) = message.split_once("supported versions:")?;
684+
let supported_versions = supported_versions.split(')').next()?;
685+
let mut versions = Vec::new();
686+
for candidate in supported_versions.split(',') {
687+
let candidate = candidate
688+
.trim()
689+
.trim_matches(|character| matches!(character, '[' | ']' | '"' | '\''));
690+
let bytes = candidate.as_bytes();
691+
if bytes.len() != 10
692+
|| bytes.get(4) != Some(&b'-')
693+
|| bytes.get(7) != Some(&b'-')
694+
|| bytes
695+
.iter()
696+
.enumerate()
697+
.any(|(index, byte)| index != 4 && index != 7 && !byte.is_ascii_digit())
698+
{
699+
return None;
700+
}
701+
let version = serde_json::from_value::<ProtocolVersion>(serde_json::Value::String(
702+
candidate.to_owned(),
703+
))
704+
.ok()?;
705+
versions.push(version);
706+
}
707+
708+
(!versions.is_empty()).then_some(versions)
709+
}
710+
553711
/// Client-specific lifecycle entry points.
554712
pub trait ClientServiceExt: Service<RoleClient> + Sized {
555713
fn serve_with_lifecycle<T, E, A>(
@@ -676,7 +834,8 @@ where
676834
&client_info,
677835
preferred_versions,
678836
)
679-
.await?;
837+
.await
838+
.map_err(|error| error.error)?;
680839
}
681840
ClientLifecycleMode::Auto {
682841
preferred_versions,
@@ -693,17 +852,15 @@ where
693852
.await;
694853
match discover_result {
695854
Ok(()) => {}
696-
Err(ClientInitializeError::JsonRpcError(error))
697-
if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND =>
698-
{
855+
Err(error) if error.is_legacy_server() => {
699856
let mut legacy_info = client_info;
700857
if let Some(version) = legacy_version {
701858
legacy_info.protocol_version = version;
702859
}
703860
legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info)
704861
.await?;
705862
}
706-
Err(error) => return Err(error),
863+
Err(error) => return Err(error.error),
707864
}
708865
}
709866
}
@@ -739,7 +896,12 @@ where
739896
})?;
740897

741898
let (response, response_id) =
742-
expect_response(transport, "initialize response", service, peer.clone()).await?;
899+
match expect_response(transport, "initialize response", service, peer.clone()).await? {
900+
StartupResponse::Response(response, response_id) => (*response, response_id),
901+
StartupResponse::Error(error, _) => {
902+
return Err(ClientInitializeError::JsonRpcError(error));
903+
}
904+
};
743905

744906
if !id.matches_response_id(&response_id) {
745907
return Err(ClientInitializeError::ConflictInitResponseId(
@@ -773,13 +935,13 @@ async fn discover_startup<S, T>(
773935
peer: &Peer<RoleClient>,
774936
client_info: &ClientInfo,
775937
preferred_versions: Vec<ProtocolVersion>,
776-
) -> Result<(), ClientInitializeError>
938+
) -> Result<(), DiscoverStartupError>
777939
where
778940
S: Service<RoleClient>,
779941
T: Transport<RoleClient> + 'static,
780942
{
781943
if preferred_versions.is_empty() {
782-
return Err(ClientInitializeError::NoPreferredProtocolVersion);
944+
return Err(ClientInitializeError::NoPreferredProtocolVersion.into());
783945
}
784946

785947
let mut attempted = Vec::new();
@@ -805,21 +967,29 @@ where
805967
ClientInitializeError::transport::<T>(error, "send discover request")
806968
})?;
807969

808-
match expect_response(transport, "discover response", service, peer.clone()).await {
809-
Ok((ServerResult::DiscoverResult(result), response_id)) => {
970+
match expect_response(transport, "discover response", service, peer.clone()).await? {
971+
StartupResponse::Response(response, response_id) => {
972+
let result = match *response {
973+
ServerResult::DiscoverResult(result) => result,
974+
response => {
975+
return Err(
976+
ClientInitializeError::ExpectedInitResult(Some(response)).into()
977+
);
978+
}
979+
};
810980
if !id.matches_response_id(&response_id) {
811-
return Err(ClientInitializeError::ConflictInitResponseId(
812-
id,
813-
response_id,
814-
));
981+
return Err(
982+
ClientInitializeError::ConflictInitResponseId(id, response_id).into(),
983+
);
815984
}
816985
let Some(selected) =
817986
select_protocol_version(&preferred_versions, &result.supported_versions)
818987
else {
819988
return Err(ClientInitializeError::NoCompatibleProtocolVersion {
820989
client_supported: preferred_versions,
821990
server_supported: result.supported_versions,
822-
});
991+
}
992+
.into());
823993
};
824994
peer.set_peer_info(ServerInfo {
825995
protocol_version: selected.clone(),
@@ -835,12 +1005,28 @@ where
8351005
});
8361006
return Ok(());
8371007
}
838-
Ok((response, _)) => {
839-
return Err(ClientInitializeError::ExpectedInitResult(Some(response)));
840-
}
841-
Err(ClientInitializeError::JsonRpcError(error))
842-
if error.code == crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION =>
843-
{
1008+
StartupResponse::Error(error, response_id) => {
1009+
if let Some(response_id) = response_id.as_ref()
1010+
&& !id.matches_response_id(response_id)
1011+
{
1012+
return Err(ClientInitializeError::ConflictInitResponseId(
1013+
id,
1014+
response_id.clone(),
1015+
)
1016+
.into());
1017+
}
1018+
1019+
if error.code != crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION
1020+
|| response_id.is_none()
1021+
{
1022+
return Err(DiscoverStartupError::json_rpc(
1023+
error,
1024+
candidate,
1025+
id,
1026+
response_id,
1027+
));
1028+
}
1029+
8441030
let supported = error
8451031
.data
8461032
.as_ref()
@@ -865,11 +1051,11 @@ where
8651051
return Err(ClientInitializeError::NoCompatibleProtocolVersion {
8661052
client_supported: preferred_versions,
8671053
server_supported: supported,
868-
});
1054+
}
1055+
.into());
8691056
};
8701057
candidate = next;
8711058
}
872-
Err(error) => return Err(error),
8731059
}
8741060
}
8751061
}

crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -192,6 +192,13 @@ impl StreamableHttpClient for reqwest::Client {
192192
www_authenticate_header: header,
193193
}));
194194
}
195+
// Authentication failures must retain their HTTP meaning even
196+
// when the server includes a JSON-RPC error body. In particular,
197+
// Auto lifecycle negotiation must never interpret a 401 as a
198+
// legacy-protocol rejection and silently downgrade.
199+
return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned(
200+
format!("HTTP {}: authentication required", response.status()),
201+
)));
195202
}
196203
if response.status() == reqwest::StatusCode::FORBIDDEN {
197204
if let Some(header) = response.headers().get(WWW_AUTHENTICATE) {
@@ -208,6 +215,11 @@ impl StreamableHttpClient for reqwest::Client {
208215
},
209216
));
210217
}
218+
// A 403 is likewise an authorization decision, not evidence that
219+
// the endpoint only supports legacy initialization.
220+
return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned(
221+
format!("HTTP {}: access forbidden", response.status()),
222+
)));
211223
}
212224
let status = response.status();
213225
if matches!(

0 commit comments

Comments
 (0)