From f2e5f39df5323e621bf4db38bc4915878b6f2eaa Mon Sep 17 00:00:00 2001 From: Theodore Ni <3806110+tjni@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:13:13 +0000 Subject: [PATCH] fix(auth): preserve OAuth discovery transport errors --- crates/rmcp/src/error.rs | 14 ++ crates/rmcp/src/transport/auth.rs | 205 +++++++++++++++++++++++------- 2 files changed, 174 insertions(+), 45 deletions(-) diff --git a/crates/rmcp/src/error.rs b/crates/rmcp/src/error.rs index 74f7d4383..1e8635a3a 100644 --- a/crates/rmcp/src/error.rs +++ b/crates/rmcp/src/error.rs @@ -17,6 +17,20 @@ impl Display for ErrorData { impl std::error::Error for ErrorData {} +#[cfg(all(feature = "auth", any(feature = "client", feature = "server")))] +pub(crate) struct ErrorChain<'a>(pub(crate) &'a (dyn std::error::Error + 'static)); + +#[cfg(all(feature = "auth", any(feature = "client", feature = "server")))] +impl Display for ErrorChain<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0)?; + for source in std::iter::successors(self.0.source(), |source| source.source()) { + write!(f, "\n Caused by: {source}")?; + } + Ok(()) + } +} + /// This is an unified error type for the errors could be returned by the service. #[derive(Debug, thiserror::Error)] #[allow(clippy::large_enum_variant)] diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index c86dfb8bb..12344f31a 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -78,16 +78,32 @@ impl OAuthHttpRequest { /// Error returned by a custom OAuth HTTP client. #[derive(Debug, Error)] -#[error("{message}")] +#[error(transparent)] pub struct OAuthHttpClientError { - message: String, + inner: OAuthHttpClientErrorKind, +} + +#[derive(Debug, Error)] +enum OAuthHttpClientErrorKind { + #[error("{0}")] + Message(String), + + #[error("{0}")] + Source(#[source] Box), } impl OAuthHttpClientError { - /// Create an error from a transport-provided message. + /// Create an error from a message. pub fn new(message: impl Into) -> Self { Self { - message: message.into(), + inner: OAuthHttpClientErrorKind::Message(message.into()), + } + } + + /// Create an error from its underlying cause. + pub fn from_error(source: impl Into>) -> Self { + Self { + inner: OAuthHttpClientErrorKind::Source(source.into()), } } } @@ -137,12 +153,12 @@ impl OAuthHttpClient for ReqwestOAuthHttpClient { OAuthHttpRedirectPolicy::Follow => &self.follow_redirects, OAuthHttpRedirectPolicy::Stop => &self.stop_redirects, }; - let request = reqwest::Request::try_from(request) - .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + let request = + reqwest::Request::try_from(request).map_err(OAuthHttpClientError::from_error)?; let response = client .execute(request) .await - .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + .map_err(OAuthHttpClientError::from_error)?; let mut builder = oauth2::http::Response::builder() .status(response.status()) @@ -153,7 +169,7 @@ impl OAuthHttpClient for ReqwestOAuthHttpClient { let mut body = Vec::new(); let mut body_stream = response.bytes_stream(); while let Some(chunk) = body_stream.next().await { - let chunk = chunk.map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + let chunk = chunk.map_err(OAuthHttpClientError::from_error)?; if chunk.len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES - body.len() { return Err(OAuthHttpClientError::new(format!( "OAuth HTTP response body exceeds {MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES} bytes" @@ -1984,13 +2000,10 @@ impl AuthorizationManager { discovery_url: &Url, ) -> Result, AuthError> { debug!("discovery url: {:?}", discovery_url); - let response = match self.discovery_get(discovery_url).await { - Ok(r) => r, - Err(e) => { - debug!("discovery request failed: {}", e); - return Ok(None); - } - }; + let response = self + .discovery_get(discovery_url) + .await + .map_err(|error| Self::discovery_failed(discovery_url, error))?; if response.status() != StatusCode::OK { debug!("discovery returned non-200: {}", response.status()); @@ -2200,8 +2213,9 @@ impl AuthorizationManager { } async fn discover_resource_metadata_url(&self) -> Result, AuthError> { - if let Ok(Some(resource_metadata_url)) = - self.fetch_resource_metadata_url(&self.base_url, true).await + if let Some(resource_metadata_url) = self + .fetch_resource_metadata_url(&self.base_url, true) + .await? { return Ok(Some(resource_metadata_url)); } @@ -2215,9 +2229,9 @@ impl AuthorizationManager { discovery_url.set_query(None); discovery_url.set_fragment(None); discovery_url.set_path(&candidate_path); - if let Ok(Some(resource_metadata_url)) = self + if let Some(resource_metadata_url) = self .fetch_resource_metadata_url(&discovery_url, false) - .await + .await? { return Ok(Some(resource_metadata_url)); } @@ -2233,13 +2247,10 @@ impl AuthorizationManager { url: &Url, allow_post_probe: bool, ) -> Result, AuthError> { - let response = match self.discovery_get(url).await { - Ok(r) => r, - Err(e) => { - debug!("resource metadata probe failed: {}", e); - return Ok(None); - } - }; + let response = self + .discovery_get(url) + .await + .map_err(|error| Self::discovery_failed(url, error))?; match response.status() { StatusCode::OK => Ok(Some(url.clone())), @@ -2267,20 +2278,13 @@ impl AuthorizationManager { .header(CONTENT_TYPE, "application/json") .body(RESOURCE_METADATA_POST_PROBE_BODY.as_bytes().to_vec()) .map_err(|error| AuthError::InternalError(error.to_string()))?; - let response = match self - .http_client - .execute(OAuthHttpRequest::new( + let response = self + .discovery_request(OAuthHttpRequest::new( request, OAuthHttpRedirectPolicy::Stop, )) .await - { - Ok(response) => response, - Err(error) => { - debug!("resource metadata POST probe failed: {}", error); - return Ok(None); - } - }; + .map_err(|error| Self::discovery_failed(url, error))?; if response.status() != StatusCode::UNAUTHORIZED { debug!( @@ -2328,13 +2332,10 @@ impl AuthorizationManager { "resource metadata discovery url: {:?}", resource_metadata_url ); - let response = match self.discovery_get(resource_metadata_url).await { - Ok(r) => r, - Err(e) => { - debug!("resource metadata request failed: {}", e); - return Ok(None); - } - }; + let response = self + .discovery_get(resource_metadata_url) + .await + .map_err(|error| Self::discovery_failed(resource_metadata_url, error))?; if response.status() != StatusCode::OK { debug!( @@ -2354,6 +2355,28 @@ impl AuthorizationManager { Ok(Some(metadata)) } + fn discovery_failed(url: &Url, error: OAuthHttpClientError) -> AuthError { + let source = std::error::Error::source(&error).unwrap_or(&error); + AuthError::MetadataError(format!( + "OAuth metadata discovery failed for {url}\n Caused by: {}", + crate::error::ErrorChain(source) + )) + } + + async fn discovery_request( + &self, + request: OAuthHttpRequest, + ) -> Result { + let response = self.http_client.execute(request).await?; + if response.status().is_server_error() { + return Err(OAuthHttpClientError::new(format!( + "HTTP {}", + response.status() + ))); + } + Ok(response) + } + async fn discovery_get(&self, url: &Url) -> Result { let mut current_url = url.clone(); for _ in 0..MAX_OAUTH_DISCOVERY_REDIRECTS { @@ -2364,8 +2387,7 @@ impl AuthorizationManager { .body(Vec::new()) .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; let response = self - .http_client - .execute(OAuthHttpRequest::new( + .discovery_request(OAuthHttpRequest::new( request, OAuthHttpRedirectPolicy::Stop, )) @@ -3617,6 +3639,99 @@ mod tests { .unwrap() } + #[test] + fn oauth_http_client_error_preserves_source_chain() { + #[derive(Debug, thiserror::Error)] + #[error("request failed")] + struct RequestError(#[source] std::io::Error); + + let error = OAuthHttpClientError::from_error(RequestError(std::io::Error::other( + "certificate signed by unknown authority", + ))); + let source = std::error::Error::source(&error).unwrap(); + assert!(source.downcast_ref::().is_some()); + + let url = Url::parse("https://mcp.example.com/mcp").unwrap(); + let error = AuthorizationManager::discovery_failed(&url, error); + assert_eq!( + error.to_string(), + "Metadata error: OAuth metadata discovery failed for https://mcp.example.com/mcp\n Caused by: request failed\n Caused by: certificate signed by unknown authority" + ); + } + + #[tokio::test] + async fn default_http_client_preserves_connection_failure_cause() { + let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let url = format!("http://{}/mcp", listener.local_addr().unwrap()); + drop(listener); + + let manager = AuthorizationManager::new(&url).await.unwrap(); + let error = manager.discover_metadata().await.unwrap_err(); + + assert!( + matches!( + error, + AuthError::MetadataError(ref reason) + if reason.contains(&url) + && reason.contains("\n Caused by: error sending request for url") + && reason.matches("error sending request for url").count() == 1 + && reason.to_ascii_lowercase().contains("connection refused") + ), + "unexpected discovery error: {error}" + ); + } + + #[tokio::test] + async fn authorization_metadata_propagates_transport_failure() { + let responses = preregistered_discovery_responses() + .into_iter() + .take(2) + .collect(); + let client = RecordingOAuthHttpClient::with_responses(responses); + let manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + Arc::new(client.clone()), + ) + .await + .unwrap(); + + let error = manager.discover_metadata().await.unwrap_err(); + + assert!( + matches!( + error, + AuthError::MetadataError(ref reason) + if reason.contains("https://auth.example.com/.well-known/oauth-authorization-server") + && reason.contains("missing fake response") + ), + "unexpected discovery error: {error}" + ); + assert_eq!(client.requests().len(), 3); + } + + #[tokio::test] + async fn discovery_propagates_server_errors() { + let manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + Arc::new(RecordingOAuthHttpClient::with_responses(vec![ + empty_response(503), + ])), + ) + .await + .unwrap(); + + let error = manager.discover_metadata().await.unwrap_err(); + + assert!( + matches!( + error, + AuthError::MetadataError(ref reason) + if reason.contains("https://mcp.example.com/mcp") && reason.contains("503") + ), + "unexpected discovery error: {error}" + ); + } + #[tokio::test] async fn custom_http_client_handles_protected_resource_discovery() { let challenge = oauth2::http::Response::builder()