Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
193 changes: 185 additions & 8 deletions crates/rmcp/src/transport/auth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -578,6 +578,28 @@ pub struct AuthorizationMetadata {
pub additional_fields: HashMap<String, serde_json::Value>,
}

/// Describes how authorization metadata was obtained during OAuth discovery.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum AuthorizationMetadataSource {
/// Protected resource metadata identified the authorization server.
ProtectedResourceMetadata,
/// Authorization server metadata was discovered directly from the base URL.
AuthorizationServerMetadata,
/// No metadata was discovered; endpoints were synthesized for legacy OAuth.
LegacyEndpointFallback,
}

/// Authorization metadata together with the discovery mechanism that produced it.
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct DiscoveredAuthorizationMetadata {
/// The discovered authorization metadata or synthesized legacy endpoints.
pub metadata: AuthorizationMetadata,
/// Whether the metadata was discovered or synthesized as a legacy fallback.
pub source: AuthorizationMetadataSource,
}

#[derive(Debug, Clone, Deserialize)]
struct ResourceServerMetadata {
resource: Option<String>,
Expand Down Expand Up @@ -1320,18 +1342,44 @@ impl AuthorizationManager {
Ok(())
}

/// discover oauth2 metadata (per SEP-985: Protected Resource Metadata first, then direct OAuth)
/// Discover OAuth metadata, preserving legacy endpoint fallback compatibility.
///
/// Use [`Self::discover_metadata_with_source`] when callers need to
/// distinguish discovered metadata from synthesized legacy endpoints.
pub async fn discover_metadata(&self) -> Result<AuthorizationMetadata, AuthError> {
self.discover_metadata_with_source()
.await
.map(|discovered| discovered.metadata)
}

/// Discover OAuth metadata and identify the mechanism that produced it.
///
/// Protected resource metadata is tried first, followed by direct
/// authorization server metadata discovery. If neither succeeds, legacy
/// endpoints are synthesized for compatibility and explicitly identified
/// as [`AuthorizationMetadataSource::LegacyEndpointFallback`].
pub async fn discover_metadata_with_source(
&self,
) -> Result<DiscoveredAuthorizationMetadata, AuthError> {
if let Some(metadata) = self.discover_oauth_server_via_resource_metadata().await? {
return Ok(metadata);
return Ok(DiscoveredAuthorizationMetadata {
metadata,
source: AuthorizationMetadataSource::ProtectedResourceMetadata,
});
}

if let Some(metadata) = self.try_discover_oauth_server(&self.base_url).await? {
return Ok(metadata);
return Ok(DiscoveredAuthorizationMetadata {
metadata,
source: AuthorizationMetadataSource::AuthorizationServerMetadata,
});
}

debug!("falling back to legacy OAuth endpoints derived from the base URL");
Ok(Self::legacy_authorization_metadata(&self.base_url))
Ok(DiscoveredAuthorizationMetadata {
metadata: Self::legacy_authorization_metadata(&self.base_url),
source: AuthorizationMetadataSource::LegacyEndpointFallback,
})
}

fn legacy_authorization_metadata(base_url: &Url) -> AuthorizationMetadata {
Expand Down Expand Up @@ -3695,10 +3743,10 @@ mod tests {

use super::{
AuthError, AuthorizationCallback, AuthorizationManager, AuthorizationMetadata,
AuthorizationRequest, AuthorizationSession, CredentialStore, InMemoryCredentialStore,
InMemoryStateStore, OAuthClientConfig, OAuthHttpClient, OAuthHttpClientError,
OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig,
StateStore, StoredAuthorizationState, is_https_url,
AuthorizationMetadataSource, AuthorizationRequest, AuthorizationSession, CredentialStore,
InMemoryCredentialStore, InMemoryStateStore, OAuthClientConfig, OAuthHttpClient,
OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest,
ScopeUpgradeConfig, StateStore, StoredAuthorizationState, is_https_url,
};
use crate::transport::auth::VendorExtraTokenFields;

Expand Down Expand Up @@ -3830,6 +3878,97 @@ mod tests {
);
}

#[tokio::test]
async fn discover_metadata_with_source_identifies_protected_resource_metadata() {
let challenge = oauth2::http::Response::builder()
.status(401)
.header(
"www-authenticate",
r#"Bearer resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource""#,
)
.body(Vec::new())
.unwrap();
let client = RecordingOAuthHttpClient::with_responses(vec![
challenge,
http_response(
200,
serde_json::json!({
"resource": "https://mcp.example.com/mcp",
"authorization_servers": ["https://auth.example.com"]
}),
),
http_response(
200,
serde_json::json!({
"issuer": "https://auth.example.com",
"authorization_endpoint": "https://auth.example.com/authorize",
"token_endpoint": "https://auth.example.com/token"
}),
),
]);
let manager = AuthorizationManager::new_with_oauth_http_client(
"https://mcp.example.com/mcp",
Arc::new(client),
)
.await
.unwrap();

let discovered = manager.discover_metadata_with_source().await.unwrap();

assert_eq!(
(
discovered.source,
discovered.metadata.issuer.as_deref(),
discovered.metadata.token_endpoint.as_str(),
),
(
AuthorizationMetadataSource::ProtectedResourceMetadata,
Some("https://auth.example.com"),
"https://auth.example.com/token",
)
);
}

#[tokio::test]
async fn discover_metadata_with_source_identifies_authorization_server_metadata() {
let client = RecordingOAuthHttpClient::with_responses(vec![
empty_response(404),
empty_response(404),
empty_response(404),
http_response(
200,
serde_json::json!({
"issuer": "https://mcp.example.com",
"authorization_endpoint": "https://mcp.example.com/authorize",
"token_endpoint": "https://mcp.example.com/token"
}),
),
]);
let manager = AuthorizationManager::new_with_oauth_http_client(
"https://mcp.example.com/",
Arc::new(client.clone()),
)
.await
.unwrap();

let discovered = manager.discover_metadata_with_source().await.unwrap();

assert_eq!(
(
discovered.source,
discovered.metadata.issuer.as_deref(),
discovered.metadata.token_endpoint.as_str(),
client.requests().len(),
),
(
AuthorizationMetadataSource::AuthorizationServerMetadata,
Some("https://mcp.example.com"),
"https://mcp.example.com/token",
4,
)
);
}

#[tokio::test]
async fn protected_resource_metadata_supports_authorization_server_path_insertion() {
let client = RecordingOAuthHttpClient::with_responses(vec![
Expand Down Expand Up @@ -4124,6 +4263,44 @@ mod tests {
);
}

#[tokio::test]
async fn discover_metadata_with_source_identifies_legacy_endpoint_fallback() {
let client = RecordingOAuthHttpClient::with_responses(vec![
empty_response(404),
empty_response(404),
empty_response(404),
empty_response(404),
empty_response(404),
]);
let manager = AuthorizationManager::new_with_oauth_http_client(
"https://public.example.com/",
Arc::new(client.clone()),
)
.await
.unwrap();

let discovered = manager.discover_metadata_with_source().await.unwrap();

assert_eq!(
(
discovered.source,
discovered.metadata.authorization_endpoint.as_str(),
discovered.metadata.token_endpoint.as_str(),
discovered.metadata.registration_endpoint.as_deref(),
discovered.metadata.issuer.as_deref(),
client.requests().len(),
),
(
AuthorizationMetadataSource::LegacyEndpointFallback,
"https://public.example.com/authorize",
"https://public.example.com/token",
Some("https://public.example.com/register"),
None,
5,
)
);
}

fn preregistered_as_metadata_response() -> HttpResponse {
http_response(
200,
Expand Down