From 4a4e04e80934dac9695dd3d90ce54c32b3e4eb79 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 1/2] fix(auth): preserve OAuth discovery transport errors --- crates/rmcp/src/error.rs | 14 ++ crates/rmcp/src/transport/auth.rs | 211 +++++++++++++++++++++++------- 2 files changed, 180 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 c6751d88f..ce55f8e79 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -72,16 +72,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()), } } } @@ -131,12 +147,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()) @@ -147,7 +163,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" @@ -2283,13 +2299,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()); @@ -2387,7 +2400,7 @@ impl AuthorizationManager { async fn discover_oauth_server_via_resource_metadata( &self, ) -> Result, AuthError> { - let Some(resource_metadata_url) = self.discover_resource_metadata_url().await else { + let Some(resource_metadata_url) = self.discover_resource_metadata_url().await? else { return Ok(None); }; self.discover_oauth_server_from_resource_metadata_url(&resource_metadata_url) @@ -2509,10 +2522,11 @@ impl AuthorizationManager { && Self::is_same_origin(&root_resource, &path_resource) } - async fn discover_resource_metadata_url(&self) -> Option { - if let Some(resource_metadata_url) = self.probe_resource_metadata_url(&self.base_url).await + async fn discover_resource_metadata_url(&self) -> Result, AuthError> { + if let Some(resource_metadata_url) = + self.probe_resource_metadata_url(&self.base_url).await? { - return Some(resource_metadata_url); + return Ok(Some(resource_metadata_url)); } // If the primary URL doesn't use WWW-Authenticate, try oauth-protected-resource discovery. @@ -2525,37 +2539,33 @@ impl AuthorizationManager { discovery_url.set_fragment(None); discovery_url.set_path(&candidate_path); if let Some(resource_metadata_url) = - self.probe_resource_metadata_url(&discovery_url).await + self.probe_resource_metadata_url(&discovery_url).await? { - return Some(resource_metadata_url); + return Ok(Some(resource_metadata_url)); } } - None + Ok(None) } /// Probe `url` with a GET, extracting the resource metadata url from a /// 200 (the url itself is the metadata document) or from a 401's /// WWW-Authenticate header value. /// https://www.rfc-editor.org/rfc/rfc9728.html#name-use-of-www-authenticate-for - async fn probe_resource_metadata_url(&self, url: &Url) -> Option { - let response = match self.discovery_get(url).await { - Ok(r) => r, - Err(e) => { - debug!("resource metadata probe failed: {}", e); - return None; - } - }; + async fn probe_resource_metadata_url(&self, url: &Url) -> Result, AuthError> { + let response = self + .discovery_get(url) + .await + .map_err(|error| Self::discovery_failed(url, error))?; match response.status() { - StatusCode::OK => Some(url.clone()), - StatusCode::UNAUTHORIZED => { - self.extract_resource_metadata_url_from_www_authenticate(&response) - .await - } + StatusCode::OK => Ok(Some(url.clone())), + StatusCode::UNAUTHORIZED => Ok(self + .extract_resource_metadata_url_from_www_authenticate(&response) + .await), status => { debug!("resource metadata probe returned unexpected status: {status}"); - None + Ok(None) } } } @@ -2588,13 +2598,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!( @@ -2614,6 +2621,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 { @@ -2624,8 +2653,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, )) @@ -3870,6 +3898,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() From 86a3c638ce18a6f8b0e4356dab2fd7117cd1105e Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 28 Jul 2026 10:01:43 -0400 Subject: [PATCH 2/2] refactor!: expose domain-specific boxed OAuth HTTP errors Let custom OAuth HTTP clients preserve native error types and source chains while keeping the type-erased boundary explicit in the public API. BREAKING CHANGE: OAuthHttpClientError is now a boxed error alias; custom clients should return native errors with .into() or box them directly. --- crates/rmcp/src/transport/auth.rs | 120 +++++++++++++++--------------- docs/OAUTH_SUPPORT.md | 4 +- 2 files changed, 61 insertions(+), 63 deletions(-) diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index ce55f8e79..ef562ffda 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -70,36 +70,19 @@ impl OAuthHttpRequest { } } -/// Error returned by a custom OAuth HTTP client. -#[derive(Debug, Error)] -#[error(transparent)] -pub struct OAuthHttpClientError { - inner: OAuthHttpClientErrorKind, -} +/// Type-erased error returned by an [`OAuthHttpClient`]. +pub type OAuthHttpClientError = Box; #[derive(Debug, Error)] -enum OAuthHttpClientErrorKind { - #[error("{0}")] - Message(String), - - #[error("{0}")] - Source(#[source] Box), -} - -impl OAuthHttpClientError { - /// Create an error from a message. - pub fn new(message: impl Into) -> Self { - Self { - 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()), - } - } +enum OAuthHttpError { + #[error("OAuth HTTP response body exceeds {0} bytes")] + ResponseBodyTooLarge(usize), + #[error("unexpected HTTP status {0}")] + UnexpectedStatus(StatusCode), + #[error("OAuth discovery redirect to non-same-origin URL rejected: {0}")] + CrossOriginRedirect(Url), + #[error("OAuth discovery exceeded {0} redirects")] + TooManyRedirects(usize), } /// Future returned by [`OAuthHttpClient::execute`]. @@ -147,12 +130,12 @@ impl OAuthHttpClient for ReqwestOAuthHttpClient { OAuthHttpRedirectPolicy::Follow => &self.follow_redirects, OAuthHttpRedirectPolicy::Stop => &self.stop_redirects, }; - let request = - reqwest::Request::try_from(request).map_err(OAuthHttpClientError::from_error)?; + let request = reqwest::Request::try_from(request) + .map_err(|error| Box::new(error) as OAuthHttpClientError)?; let response = client .execute(request) .await - .map_err(OAuthHttpClientError::from_error)?; + .map_err(|error| Box::new(error) as OAuthHttpClientError)?; let mut builder = oauth2::http::Response::builder() .status(response.status()) @@ -163,17 +146,17 @@ 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(OAuthHttpClientError::from_error)?; + let chunk = chunk.map_err(|error| Box::new(error) as OAuthHttpClientError)?; 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" - ))); + return Err(Box::new(OAuthHttpError::ResponseBodyTooLarge( + MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES, + )) as OAuthHttpClientError); } body.extend_from_slice(&chunk); } builder .body(body) - .map_err(|error| OAuthHttpClientError::new(error.to_string())) + .map_err(|error| Box::new(error) as OAuthHttpClientError) }) } } @@ -183,16 +166,35 @@ struct OAuth2HttpClient<'a> { redirect_policy: OAuthHttpRedirectPolicy, } +#[derive(Debug)] +struct OAuth2HttpClientError(OAuthHttpClientError); + +impl std::fmt::Display for OAuth2HttpClientError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("OAuth HTTP request failed") + } +} + +impl std::error::Error for OAuth2HttpClientError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(self.0.as_ref()) + } +} + impl<'c> AsyncHttpClient<'c> for OAuth2HttpClient<'_> { - type Error = OAuthHttpClientError; + type Error = OAuth2HttpClientError; type Future = std::pin::Pin< Box> + Send + 'c>, >; fn call(&'c self, request: HttpRequest) -> Self::Future { - self.client - .execute(OAuthHttpRequest::new(request, self.redirect_policy)) + Box::pin(async move { + self.client + .execute(OAuthHttpRequest::new(request, self.redirect_policy)) + .await + .map_err(OAuth2HttpClientError) + }) } } @@ -2622,10 +2624,9 @@ impl AuthorizationManager { } 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) + crate::error::ErrorChain(error.as_ref()) )) } @@ -2635,9 +2636,8 @@ impl AuthorizationManager { ) -> Result { let response = self.http_client.execute(request).await?; if response.status().is_server_error() { - return Err(OAuthHttpClientError::new(format!( - "HTTP {}", - response.status() + return Err(Box::new(OAuthHttpError::UnexpectedStatus( + response.status(), ))); } Ok(response) @@ -2651,7 +2651,7 @@ impl AuthorizationManager { .uri(current_url.as_str()) .header(HEADER_MCP_PROTOCOL_VERSION, "2024-11-05") .body(Vec::new()) - .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + .map_err(|error| Box::new(error) as OAuthHttpClientError)?; let response = self .discovery_request(OAuthHttpRequest::new( request, @@ -2668,23 +2668,21 @@ impl AuthorizationManager { }; let location = location .to_str() - .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + .map_err(|error| Box::new(error) as OAuthHttpClientError)?; let next_url = current_url .join(location) - .map_err(|error| OAuthHttpClientError::new(error.to_string()))?; + .map_err(|error| Box::new(error) as OAuthHttpClientError)?; if Self::is_http_url(&next_url) && Self::is_same_origin(¤t_url, &next_url) { current_url = next_url; continue; } - return Err(OAuthHttpClientError::new(format!( - "OAuth discovery redirect to non-same-origin URL rejected: {next_url}" - ))); + return Err(Box::new(OAuthHttpError::CrossOriginRedirect(next_url))); } - Err(OAuthHttpClientError::new(format!( - "OAuth discovery exceeded {MAX_OAUTH_DISCOVERY_REDIRECTS} redirects" + Err(Box::new(OAuthHttpError::TooManyRedirects( + MAX_OAUTH_DISCOVERY_REDIRECTS, ))) } @@ -3870,9 +3868,7 @@ mod tests { body: request.request.body().clone(), }); let response = self.responses.lock().unwrap().pop_front(); - Box::pin(async move { - response.ok_or_else(|| OAuthHttpClientError::new("missing fake response")) - }) + Box::pin(async move { response.ok_or_else(|| "missing fake response".into()) }) } } @@ -3904,11 +3900,11 @@ mod tests { #[error("request failed")] struct RequestError(#[source] std::io::Error); - let error = OAuthHttpClientError::from_error(RequestError(std::io::Error::other( + let error: OAuthHttpClientError = RequestError(std::io::Error::other( "certificate signed by unknown authority", - ))); - let source = std::error::Error::source(&error).unwrap(); - assert!(source.downcast_ref::().is_some()); + )) + .into(); + assert!(error.downcast_ref::().is_some()); let url = Url::parse("https://mcp.example.com/mcp").unwrap(); let error = AuthorizationManager::discovery_failed(&url, error); @@ -3925,7 +3921,7 @@ mod tests { drop(listener); let manager = AuthorizationManager::new(&url).await.unwrap(); - let error = manager.discover_metadata().await.unwrap_err(); + let error = manager.resolve_metadata().await.unwrap_err(); assert!( matches!( @@ -3954,7 +3950,7 @@ mod tests { .await .unwrap(); - let error = manager.discover_metadata().await.unwrap_err(); + let error = manager.resolve_metadata().await.unwrap_err(); assert!( matches!( @@ -3979,7 +3975,7 @@ mod tests { .await .unwrap(); - let error = manager.discover_metadata().await.unwrap_err(); + let error = manager.resolve_metadata().await.unwrap_err(); assert!( matches!( diff --git a/docs/OAUTH_SUPPORT.md b/docs/OAUTH_SUPPORT.md index 1d09d7402..b35478739 100644 --- a/docs/OAUTH_SUPPORT.md +++ b/docs/OAUTH_SUPPORT.md @@ -61,7 +61,9 @@ have been obtained. If OAuth requests must run outside reqwest, implement `OAuthHttpClient` and use `OAuthState::new_with_oauth_http_client`. The SDK passes each OAuth request to your implementation with the raw HTTP request, a suggested timeout, and an -`OAuthHttpRedirectPolicy`. +`OAuthHttpRedirectPolicy`. `OAuthHttpClientFuture` returns +`OAuthHttpClientError`, so implementations can propagate their native error +types with `?` without flattening their source chains into strings. ```rust ignore use std::sync::Arc;