Skip to content

Commit e4cebf7

Browse files
committed
fix: keep post probe eager to preserve scope hint
1 parent 2042d10 commit e4cebf7

1 file changed

Lines changed: 62 additions & 82 deletions

File tree

crates/rmcp/src/transport/auth.rs

Lines changed: 62 additions & 82 deletions
Original file line numberDiff line numberDiff line change
@@ -1003,14 +1003,13 @@ pub struct AuthorizationManager {
10031003
}
10041004

10051005
/// Outcome of a resource metadata discovery GET issued by
1006-
/// [`AuthorizationManager::fetch_resource_metadata_url`].
1007-
enum ResourceMetadataUrlProbe {
1006+
/// [`AuthorizationManager::probe_resource_metadata_url`].
1007+
enum ResourceMetadataProbeOutcome {
10081008
/// The probe located a resource metadata URL: either a 200 on the probed
10091009
/// URL itself, or a 401 whose WWW-Authenticate header pointed at one.
10101010
Found(Url),
10111011
/// The probe got a 404/405, typical of streamable HTTP servers that
1012-
/// reject session-less GETs, so the POST probe fallback is worth trying
1013-
/// if well-known discovery also comes up empty.
1012+
/// reject session-less GETs, so the POST probe fallback is worth trying.
10141013
PostProbeEligible,
10151014
/// The probe yielded nothing usable.
10161015
Unavailable,
@@ -2441,13 +2440,24 @@ impl AuthorizationManager {
24412440
}
24422441

24432442
async fn discover_resource_metadata_url(&self) -> Result<Option<Url>, AuthError> {
2444-
let should_post_probe = match self.fetch_resource_metadata_url(&self.base_url).await {
2445-
ResourceMetadataUrlProbe::Found(resource_metadata_url) => {
2443+
match self.probe_resource_metadata_url(&self.base_url).await {
2444+
ResourceMetadataProbeOutcome::Found(resource_metadata_url) => {
24462445
return Ok(Some(resource_metadata_url));
24472446
}
2448-
ResourceMetadataUrlProbe::PostProbeEligible => true,
2449-
ResourceMetadataUrlProbe::Unavailable => false,
2450-
};
2447+
// The POST probe's 401 carries both the resource metadata URL and
2448+
// the scope hint from WWW-Authenticate, so it must run before
2449+
// well-known discovery: a well-known hit would otherwise hide the
2450+
// scope the server advertises only on that challenge.
2451+
ResourceMetadataProbeOutcome::PostProbeEligible => {
2452+
if let Some(resource_metadata_url) = self
2453+
.post_probe_resource_metadata_url(&self.base_url)
2454+
.await?
2455+
{
2456+
return Ok(Some(resource_metadata_url));
2457+
}
2458+
}
2459+
ResourceMetadataProbeOutcome::Unavailable => {}
2460+
}
24512461

24522462
// If the primary URL doesn't use WWW-Authenticate, try oauth-protected-resource discovery.
24532463
// https://www.rfc-editor.org/rfc/rfc9728.html#name-obtaining-protected-resourc
@@ -2458,56 +2468,47 @@ impl AuthorizationManager {
24582468
discovery_url.set_query(None);
24592469
discovery_url.set_fragment(None);
24602470
discovery_url.set_path(&candidate_path);
2461-
if let ResourceMetadataUrlProbe::Found(resource_metadata_url) =
2462-
self.fetch_resource_metadata_url(&discovery_url).await
2471+
if let ResourceMetadataProbeOutcome::Found(resource_metadata_url) =
2472+
self.probe_resource_metadata_url(&discovery_url).await
24632473
{
24642474
return Ok(Some(resource_metadata_url));
24652475
}
24662476
}
24672477

2468-
if should_post_probe {
2469-
return self
2470-
.fetch_resource_metadata_url_with_post_probe(&self.base_url)
2471-
.await;
2472-
}
2473-
24742478
Ok(None)
24752479
}
24762480

24772481
/// Extract the resource metadata url from the WWW-Authenticate header value.
24782482
/// https://www.rfc-editor.org/rfc/rfc9728.html#name-use-of-www-authenticate-for
2479-
async fn fetch_resource_metadata_url(&self, url: &Url) -> ResourceMetadataUrlProbe {
2483+
async fn probe_resource_metadata_url(&self, url: &Url) -> ResourceMetadataProbeOutcome {
24802484
let response = match self.discovery_get(url).await {
24812485
Ok(r) => r,
24822486
Err(e) => {
24832487
debug!("resource metadata probe failed: {}", e);
2484-
return ResourceMetadataUrlProbe::Unavailable;
2488+
return ResourceMetadataProbeOutcome::Unavailable;
24852489
}
24862490
};
24872491

24882492
match response.status() {
2489-
StatusCode::OK => ResourceMetadataUrlProbe::Found(url.clone()),
2493+
StatusCode::OK => ResourceMetadataProbeOutcome::Found(url.clone()),
24902494
StatusCode::UNAUTHORIZED => self
24912495
.extract_resource_metadata_url_from_www_authenticate(&response)
24922496
.await
24932497
.map_or(
2494-
ResourceMetadataUrlProbe::Unavailable,
2495-
ResourceMetadataUrlProbe::Found,
2498+
ResourceMetadataProbeOutcome::Unavailable,
2499+
ResourceMetadataProbeOutcome::Found,
24962500
),
24972501
StatusCode::NOT_FOUND | StatusCode::METHOD_NOT_ALLOWED => {
2498-
ResourceMetadataUrlProbe::PostProbeEligible
2502+
ResourceMetadataProbeOutcome::PostProbeEligible
24992503
}
25002504
status => {
25012505
debug!("resource metadata probe returned unexpected status: {status}");
2502-
ResourceMetadataUrlProbe::Unavailable
2506+
ResourceMetadataProbeOutcome::Unavailable
25032507
}
25042508
}
25052509
}
25062510

2507-
async fn fetch_resource_metadata_url_with_post_probe(
2508-
&self,
2509-
url: &Url,
2510-
) -> Result<Option<Url>, AuthError> {
2511+
async fn post_probe_resource_metadata_url(&self, url: &Url) -> Result<Option<Url>, AuthError> {
25112512
let request = oauth2::http::Request::builder()
25122513
.method("POST")
25132514
.uri(url.as_str())
@@ -2546,7 +2547,7 @@ impl AuthorizationManager {
25462547
}
25472548

25482549
/// Best-effort cleanup of a session the server may have created for the
2549-
/// POST probe's synthetic `initialize` request (see issue #1048).
2550+
/// POST probe's synthetic `initialize` request.
25502551
/// Failures are logged and swallowed: cleanup must never fail discovery.
25512552
async fn delete_post_probe_session(&self, url: &Url, response: &HttpResponse) {
25522553
let Some(session_id) = response.headers().get(HEADER_SESSION_ID) else {
@@ -3479,7 +3480,7 @@ impl OAuthState {
34793480
)
34803481
}
34813482

3482-
async fn placeholder(&self) -> Result<Self, AuthError> {
3483+
async fn placeholder_state(&self) -> Result<Self, AuthError> {
34833484
let (http_client, refresh_redirect_policy) = self.oauth_http_client_config();
34843485
Ok(OAuthState::Unauthorized(
34853486
AuthorizationManager::new_inner(
@@ -3595,7 +3596,7 @@ impl OAuthState {
35953596
&mut self,
35963597
request: AuthorizationRequest,
35973598
) -> Result<(), AuthError> {
3598-
let placeholder = self.placeholder().await?;
3599+
let placeholder = self.placeholder_state().await?;
35993600
let old = std::mem::replace(self, placeholder);
36003601
let OAuthState::Unauthorized(mut manager) = old else {
36013602
*self = old;
@@ -3627,7 +3628,7 @@ impl OAuthState {
36273628

36283629
/// complete authorization
36293630
pub async fn complete_authorization(&mut self) -> Result<(), AuthError> {
3630-
let placeholder = self.placeholder().await?;
3631+
let placeholder = self.placeholder_state().await?;
36313632
if let OAuthState::Session(session) = std::mem::replace(self, placeholder) {
36323633
*self = OAuthState::Authorized(session.auth_manager);
36333634
Ok(())
@@ -3637,7 +3638,7 @@ impl OAuthState {
36373638
}
36383639
/// covert to authorized http client
36393640
pub async fn to_authorized_http_client(&mut self) -> Result<(), AuthError> {
3640-
let placeholder = self.placeholder().await?;
3641+
let placeholder = self.placeholder_state().await?;
36413642
if let OAuthState::Authorized(manager) = std::mem::replace(self, placeholder) {
36423643
*self = OAuthState::AuthorizedHttpClient(AuthorizedHttpClient::new(
36433644
Arc::new(manager),
@@ -3657,7 +3658,7 @@ impl OAuthState {
36573658
required_scope: &str,
36583659
redirect_uri: &str,
36593660
) -> Result<String, AuthError> {
3660-
let placeholder = self.placeholder().await?;
3661+
let placeholder = self.placeholder_state().await?;
36613662
let old = std::mem::replace(self, placeholder);
36623663
let OAuthState::Authorized(manager) = old else {
36633664
*self = old;
@@ -3789,7 +3790,7 @@ impl OAuthState {
37893790
&mut self,
37903791
config: ClientCredentialsConfig,
37913792
) -> Result<(), AuthError> {
3792-
let placeholder = self.placeholder().await?;
3793+
let placeholder = self.placeholder_state().await?;
37933794
let OAuthState::Unauthorized(mut manager) = std::mem::replace(self, placeholder) else {
37943795
return Err(AuthError::InternalError(
37953796
"Client credentials flow requires Unauthorized state".to_string(),
@@ -3971,12 +3972,20 @@ mod tests {
39713972
);
39723973
}
39733974

3975+
// The POST probe must run before well-known discovery: its 401 challenge
3976+
// is the only channel for the WWW-Authenticate scope hint, which a
3977+
// well-known hit would otherwise hide.
39743978
#[tokio::test]
3975-
async fn well_known_resource_metadata_precedes_post_probe() {
3976-
let client = RecordingOAuthHttpClient::with_responses(vec![
3977-
empty_response(405),
3978-
empty_response(200),
3979-
]);
3979+
async fn post_probe_precedes_well_known_discovery() {
3980+
let challenge = oauth2::http::Response::builder()
3981+
.status(401)
3982+
.header(
3983+
"www-authenticate",
3984+
r#"Bearer resource_metadata="https://mcp.example.com/custom/metadata/location.json""#,
3985+
)
3986+
.body(Vec::new())
3987+
.unwrap();
3988+
let client = RecordingOAuthHttpClient::with_responses(vec![empty_response(405), challenge]);
39803989
let manager = AuthorizationManager::new_with_oauth_http_client(
39813990
"https://mcp.example.com/mcp",
39823991
Arc::new(client.clone()),
@@ -3996,15 +4005,10 @@ mod tests {
39964005
.collect::<Vec<_>>(),
39974006
),
39984007
(
3999-
Some(
4000-
"https://mcp.example.com/.well-known/oauth-protected-resource/mcp".to_string()
4001-
),
4008+
Some("https://mcp.example.com/custom/metadata/location.json".to_string()),
40024009
vec![
40034010
("GET", "https://mcp.example.com/mcp"),
4004-
(
4005-
"GET",
4006-
"https://mcp.example.com/.well-known/oauth-protected-resource/mcp"
4007-
),
4011+
("POST", "https://mcp.example.com/mcp"),
40084012
],
40094013
)
40104014
);
@@ -4019,11 +4023,11 @@ mod tests {
40194023
.unwrap();
40204024
let client = RecordingOAuthHttpClient::with_responses(vec![
40214025
empty_response(405),
4026+
post_response,
4027+
empty_response(202),
40224028
empty_response(404),
40234029
empty_response(404),
40244030
empty_response(404),
4025-
post_response,
4026-
empty_response(202),
40274031
]);
40284032
let manager = AuthorizationManager::new_with_oauth_http_client(
40294033
"https://mcp.example.com/mcp",
@@ -4053,6 +4057,12 @@ mod tests {
40534057
None,
40544058
vec![
40554059
("GET", "https://mcp.example.com/mcp", None),
4060+
("POST", "https://mcp.example.com/mcp", None),
4061+
(
4062+
"DELETE",
4063+
"https://mcp.example.com/mcp",
4064+
Some("probe-session"),
4065+
),
40564066
(
40574067
"GET",
40584068
"https://mcp.example.com/.well-known/oauth-protected-resource/mcp",
@@ -4068,12 +4078,6 @@ mod tests {
40684078
"https://mcp.example.com/.well-known/oauth-protected-resource",
40694079
None,
40704080
),
4071-
("POST", "https://mcp.example.com/mcp", None),
4072-
(
4073-
"DELETE",
4074-
"https://mcp.example.com/mcp",
4075-
Some("probe-session"),
4076-
),
40774081
],
40784082
)
40794083
);
@@ -4092,9 +4096,6 @@ mod tests {
40924096
.unwrap();
40934097
let client = RecordingOAuthHttpClient::with_responses(vec![
40944098
empty_response(405),
4095-
empty_response(404),
4096-
empty_response(404),
4097-
empty_response(404),
40984099
post_response,
40994100
empty_response(202),
41004101
]);
@@ -4126,21 +4127,6 @@ mod tests {
41264127
Some("https://mcp.example.com/custom/metadata/location.json".to_string()),
41274128
vec![
41284129
("GET", "https://mcp.example.com/mcp", None),
4129-
(
4130-
"GET",
4131-
"https://mcp.example.com/.well-known/oauth-protected-resource/mcp",
4132-
None,
4133-
),
4134-
(
4135-
"GET",
4136-
"https://mcp.example.com/mcp/.well-known/oauth-protected-resource",
4137-
None,
4138-
),
4139-
(
4140-
"GET",
4141-
"https://mcp.example.com/.well-known/oauth-protected-resource",
4142-
None,
4143-
),
41444130
("POST", "https://mcp.example.com/mcp", None),
41454131
(
41464132
"DELETE",
@@ -4341,9 +4327,6 @@ mod tests {
43414327
.body(Vec::new())
43424328
.unwrap();
43434329
let client = RecordingOAuthHttpClient::with_responses(vec![
4344-
empty_response(404),
4345-
empty_response(404),
4346-
empty_response(404),
43474330
empty_response(404),
43484331
challenge,
43494332
http_response(
@@ -4386,9 +4369,6 @@ mod tests {
43864369
"https://auth.example.com/tenant1/token",
43874370
vec![
43884371
"https://mcp.example.com/mcp",
4389-
"https://mcp.example.com/.well-known/oauth-protected-resource/mcp",
4390-
"https://mcp.example.com/mcp/.well-known/oauth-protected-resource",
4391-
"https://mcp.example.com/.well-known/oauth-protected-resource",
43924372
"https://mcp.example.com/mcp",
43934373
"https://mcp.example.com/custom/metadata/location.json",
43944374
"https://auth.example.com/.well-known/oauth-authorization-server/tenant1",
@@ -4401,10 +4381,10 @@ mod tests {
44014381
client
44024382
.requests()
44034383
.iter()
4404-
.take(5)
4384+
.take(2)
44054385
.map(|request| request.method.as_str())
44064386
.collect::<Vec<_>>(),
4407-
vec!["GET", "GET", "GET", "GET", "POST"]
4387+
vec!["GET", "POST"]
44084388
);
44094389
}
44104390

@@ -4445,8 +4425,8 @@ mod tests {
44454425
Some("https://legacy.example.com/register"),
44464426
vec![
44474427
"https://legacy.example.com/",
4448-
"https://legacy.example.com/.well-known/oauth-protected-resource",
44494428
"https://legacy.example.com/",
4429+
"https://legacy.example.com/.well-known/oauth-protected-resource",
44504430
"https://legacy.example.com/.well-known/oauth-authorization-server",
44514431
"https://legacy.example.com/.well-known/openid-configuration",
44524432
],

0 commit comments

Comments
 (0)