Skip to content

Commit f2e5f39

Browse files
committed
fix(auth): preserve OAuth discovery transport errors
1 parent 8803d39 commit f2e5f39

2 files changed

Lines changed: 174 additions & 45 deletions

File tree

crates/rmcp/src/error.rs

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,20 @@ impl Display for ErrorData {
1717

1818
impl std::error::Error for ErrorData {}
1919

20+
#[cfg(all(feature = "auth", any(feature = "client", feature = "server")))]
21+
pub(crate) struct ErrorChain<'a>(pub(crate) &'a (dyn std::error::Error + 'static));
22+
23+
#[cfg(all(feature = "auth", any(feature = "client", feature = "server")))]
24+
impl Display for ErrorChain<'_> {
25+
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26+
write!(f, "{}", self.0)?;
27+
for source in std::iter::successors(self.0.source(), |source| source.source()) {
28+
write!(f, "\n Caused by: {source}")?;
29+
}
30+
Ok(())
31+
}
32+
}
33+
2034
/// This is an unified error type for the errors could be returned by the service.
2135
#[derive(Debug, thiserror::Error)]
2236
#[allow(clippy::large_enum_variant)]

crates/rmcp/src/transport/auth.rs

Lines changed: 160 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -78,16 +78,32 @@ impl OAuthHttpRequest {
7878

7979
/// Error returned by a custom OAuth HTTP client.
8080
#[derive(Debug, Error)]
81-
#[error("{message}")]
81+
#[error(transparent)]
8282
pub struct OAuthHttpClientError {
83-
message: String,
83+
inner: OAuthHttpClientErrorKind,
84+
}
85+
86+
#[derive(Debug, Error)]
87+
enum OAuthHttpClientErrorKind {
88+
#[error("{0}")]
89+
Message(String),
90+
91+
#[error("{0}")]
92+
Source(#[source] Box<dyn std::error::Error + Send + Sync>),
8493
}
8594

8695
impl OAuthHttpClientError {
87-
/// Create an error from a transport-provided message.
96+
/// Create an error from a message.
8897
pub fn new(message: impl Into<String>) -> Self {
8998
Self {
90-
message: message.into(),
99+
inner: OAuthHttpClientErrorKind::Message(message.into()),
100+
}
101+
}
102+
103+
/// Create an error from its underlying cause.
104+
pub fn from_error(source: impl Into<Box<dyn std::error::Error + Send + Sync>>) -> Self {
105+
Self {
106+
inner: OAuthHttpClientErrorKind::Source(source.into()),
91107
}
92108
}
93109
}
@@ -137,12 +153,12 @@ impl OAuthHttpClient for ReqwestOAuthHttpClient {
137153
OAuthHttpRedirectPolicy::Follow => &self.follow_redirects,
138154
OAuthHttpRedirectPolicy::Stop => &self.stop_redirects,
139155
};
140-
let request = reqwest::Request::try_from(request)
141-
.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
156+
let request =
157+
reqwest::Request::try_from(request).map_err(OAuthHttpClientError::from_error)?;
142158
let response = client
143159
.execute(request)
144160
.await
145-
.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
161+
.map_err(OAuthHttpClientError::from_error)?;
146162

147163
let mut builder = oauth2::http::Response::builder()
148164
.status(response.status())
@@ -153,7 +169,7 @@ impl OAuthHttpClient for ReqwestOAuthHttpClient {
153169
let mut body = Vec::new();
154170
let mut body_stream = response.bytes_stream();
155171
while let Some(chunk) = body_stream.next().await {
156-
let chunk = chunk.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
172+
let chunk = chunk.map_err(OAuthHttpClientError::from_error)?;
157173
if chunk.len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES - body.len() {
158174
return Err(OAuthHttpClientError::new(format!(
159175
"OAuth HTTP response body exceeds {MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES} bytes"
@@ -1984,13 +2000,10 @@ impl AuthorizationManager {
19842000
discovery_url: &Url,
19852001
) -> Result<Option<AuthorizationMetadata>, AuthError> {
19862002
debug!("discovery url: {:?}", discovery_url);
1987-
let response = match self.discovery_get(discovery_url).await {
1988-
Ok(r) => r,
1989-
Err(e) => {
1990-
debug!("discovery request failed: {}", e);
1991-
return Ok(None);
1992-
}
1993-
};
2003+
let response = self
2004+
.discovery_get(discovery_url)
2005+
.await
2006+
.map_err(|error| Self::discovery_failed(discovery_url, error))?;
19942007

19952008
if response.status() != StatusCode::OK {
19962009
debug!("discovery returned non-200: {}", response.status());
@@ -2200,8 +2213,9 @@ impl AuthorizationManager {
22002213
}
22012214

22022215
async fn discover_resource_metadata_url(&self) -> Result<Option<Url>, AuthError> {
2203-
if let Ok(Some(resource_metadata_url)) =
2204-
self.fetch_resource_metadata_url(&self.base_url, true).await
2216+
if let Some(resource_metadata_url) = self
2217+
.fetch_resource_metadata_url(&self.base_url, true)
2218+
.await?
22052219
{
22062220
return Ok(Some(resource_metadata_url));
22072221
}
@@ -2215,9 +2229,9 @@ impl AuthorizationManager {
22152229
discovery_url.set_query(None);
22162230
discovery_url.set_fragment(None);
22172231
discovery_url.set_path(&candidate_path);
2218-
if let Ok(Some(resource_metadata_url)) = self
2232+
if let Some(resource_metadata_url) = self
22192233
.fetch_resource_metadata_url(&discovery_url, false)
2220-
.await
2234+
.await?
22212235
{
22222236
return Ok(Some(resource_metadata_url));
22232237
}
@@ -2233,13 +2247,10 @@ impl AuthorizationManager {
22332247
url: &Url,
22342248
allow_post_probe: bool,
22352249
) -> Result<Option<Url>, AuthError> {
2236-
let response = match self.discovery_get(url).await {
2237-
Ok(r) => r,
2238-
Err(e) => {
2239-
debug!("resource metadata probe failed: {}", e);
2240-
return Ok(None);
2241-
}
2242-
};
2250+
let response = self
2251+
.discovery_get(url)
2252+
.await
2253+
.map_err(|error| Self::discovery_failed(url, error))?;
22432254

22442255
match response.status() {
22452256
StatusCode::OK => Ok(Some(url.clone())),
@@ -2267,20 +2278,13 @@ impl AuthorizationManager {
22672278
.header(CONTENT_TYPE, "application/json")
22682279
.body(RESOURCE_METADATA_POST_PROBE_BODY.as_bytes().to_vec())
22692280
.map_err(|error| AuthError::InternalError(error.to_string()))?;
2270-
let response = match self
2271-
.http_client
2272-
.execute(OAuthHttpRequest::new(
2281+
let response = self
2282+
.discovery_request(OAuthHttpRequest::new(
22732283
request,
22742284
OAuthHttpRedirectPolicy::Stop,
22752285
))
22762286
.await
2277-
{
2278-
Ok(response) => response,
2279-
Err(error) => {
2280-
debug!("resource metadata POST probe failed: {}", error);
2281-
return Ok(None);
2282-
}
2283-
};
2287+
.map_err(|error| Self::discovery_failed(url, error))?;
22842288

22852289
if response.status() != StatusCode::UNAUTHORIZED {
22862290
debug!(
@@ -2328,13 +2332,10 @@ impl AuthorizationManager {
23282332
"resource metadata discovery url: {:?}",
23292333
resource_metadata_url
23302334
);
2331-
let response = match self.discovery_get(resource_metadata_url).await {
2332-
Ok(r) => r,
2333-
Err(e) => {
2334-
debug!("resource metadata request failed: {}", e);
2335-
return Ok(None);
2336-
}
2337-
};
2335+
let response = self
2336+
.discovery_get(resource_metadata_url)
2337+
.await
2338+
.map_err(|error| Self::discovery_failed(resource_metadata_url, error))?;
23382339

23392340
if response.status() != StatusCode::OK {
23402341
debug!(
@@ -2354,6 +2355,28 @@ impl AuthorizationManager {
23542355
Ok(Some(metadata))
23552356
}
23562357

2358+
fn discovery_failed(url: &Url, error: OAuthHttpClientError) -> AuthError {
2359+
let source = std::error::Error::source(&error).unwrap_or(&error);
2360+
AuthError::MetadataError(format!(
2361+
"OAuth metadata discovery failed for {url}\n Caused by: {}",
2362+
crate::error::ErrorChain(source)
2363+
))
2364+
}
2365+
2366+
async fn discovery_request(
2367+
&self,
2368+
request: OAuthHttpRequest,
2369+
) -> Result<HttpResponse, OAuthHttpClientError> {
2370+
let response = self.http_client.execute(request).await?;
2371+
if response.status().is_server_error() {
2372+
return Err(OAuthHttpClientError::new(format!(
2373+
"HTTP {}",
2374+
response.status()
2375+
)));
2376+
}
2377+
Ok(response)
2378+
}
2379+
23572380
async fn discovery_get(&self, url: &Url) -> Result<HttpResponse, OAuthHttpClientError> {
23582381
let mut current_url = url.clone();
23592382
for _ in 0..MAX_OAUTH_DISCOVERY_REDIRECTS {
@@ -2364,8 +2387,7 @@ impl AuthorizationManager {
23642387
.body(Vec::new())
23652388
.map_err(|error| OAuthHttpClientError::new(error.to_string()))?;
23662389
let response = self
2367-
.http_client
2368-
.execute(OAuthHttpRequest::new(
2390+
.discovery_request(OAuthHttpRequest::new(
23692391
request,
23702392
OAuthHttpRedirectPolicy::Stop,
23712393
))
@@ -3617,6 +3639,99 @@ mod tests {
36173639
.unwrap()
36183640
}
36193641

3642+
#[test]
3643+
fn oauth_http_client_error_preserves_source_chain() {
3644+
#[derive(Debug, thiserror::Error)]
3645+
#[error("request failed")]
3646+
struct RequestError(#[source] std::io::Error);
3647+
3648+
let error = OAuthHttpClientError::from_error(RequestError(std::io::Error::other(
3649+
"certificate signed by unknown authority",
3650+
)));
3651+
let source = std::error::Error::source(&error).unwrap();
3652+
assert!(source.downcast_ref::<RequestError>().is_some());
3653+
3654+
let url = Url::parse("https://mcp.example.com/mcp").unwrap();
3655+
let error = AuthorizationManager::discovery_failed(&url, error);
3656+
assert_eq!(
3657+
error.to_string(),
3658+
"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"
3659+
);
3660+
}
3661+
3662+
#[tokio::test]
3663+
async fn default_http_client_preserves_connection_failure_cause() {
3664+
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
3665+
let url = format!("http://{}/mcp", listener.local_addr().unwrap());
3666+
drop(listener);
3667+
3668+
let manager = AuthorizationManager::new(&url).await.unwrap();
3669+
let error = manager.discover_metadata().await.unwrap_err();
3670+
3671+
assert!(
3672+
matches!(
3673+
error,
3674+
AuthError::MetadataError(ref reason)
3675+
if reason.contains(&url)
3676+
&& reason.contains("\n Caused by: error sending request for url")
3677+
&& reason.matches("error sending request for url").count() == 1
3678+
&& reason.to_ascii_lowercase().contains("connection refused")
3679+
),
3680+
"unexpected discovery error: {error}"
3681+
);
3682+
}
3683+
3684+
#[tokio::test]
3685+
async fn authorization_metadata_propagates_transport_failure() {
3686+
let responses = preregistered_discovery_responses()
3687+
.into_iter()
3688+
.take(2)
3689+
.collect();
3690+
let client = RecordingOAuthHttpClient::with_responses(responses);
3691+
let manager = AuthorizationManager::new_with_oauth_http_client(
3692+
"https://mcp.example.com/mcp",
3693+
Arc::new(client.clone()),
3694+
)
3695+
.await
3696+
.unwrap();
3697+
3698+
let error = manager.discover_metadata().await.unwrap_err();
3699+
3700+
assert!(
3701+
matches!(
3702+
error,
3703+
AuthError::MetadataError(ref reason)
3704+
if reason.contains("https://auth.example.com/.well-known/oauth-authorization-server")
3705+
&& reason.contains("missing fake response")
3706+
),
3707+
"unexpected discovery error: {error}"
3708+
);
3709+
assert_eq!(client.requests().len(), 3);
3710+
}
3711+
3712+
#[tokio::test]
3713+
async fn discovery_propagates_server_errors() {
3714+
let manager = AuthorizationManager::new_with_oauth_http_client(
3715+
"https://mcp.example.com/mcp",
3716+
Arc::new(RecordingOAuthHttpClient::with_responses(vec![
3717+
empty_response(503),
3718+
])),
3719+
)
3720+
.await
3721+
.unwrap();
3722+
3723+
let error = manager.discover_metadata().await.unwrap_err();
3724+
3725+
assert!(
3726+
matches!(
3727+
error,
3728+
AuthError::MetadataError(ref reason)
3729+
if reason.contains("https://mcp.example.com/mcp") && reason.contains("503")
3730+
),
3731+
"unexpected discovery error: {error}"
3732+
);
3733+
}
3734+
36203735
#[tokio::test]
36213736
async fn custom_http_client_handles_protected_resource_discovery() {
36223737
let challenge = oauth2::http::Response::builder()

0 commit comments

Comments
 (0)