use std::sync::Arc;
use thiserror::Error;
use super::{AuthorizationServerMetadata, DiscoveryFetcher, FetcherError};
use crate::config::ProtectedResourceMetadata;
use crate::ssrf::SsrfValidator;
#[derive(Debug, Error)]
pub enum ClientDiscoveryError {
#[error("could not determine a Protected Resource Metadata URL: {0}")]
NoResourceMetadataCandidate(String),
#[error("failed to fetch Protected Resource Metadata from any candidate URL: {0:?}")]
ResourceMetadataFetchFailed(Vec<(String, String)>),
#[error("Protected Resource Metadata response was invalid: {0}")]
ResourceMetadataInvalid(String),
#[error("Protected Resource Metadata named no authorization server")]
NoAuthorizationServer,
#[error("authorization server metadata discovery failed: {0}")]
AuthorizationServerDiscovery(#[from] FetcherError),
#[error(
"authorization server does not advertise S256 PKCE support (code_challenge_methods_supported); refusing to proceed"
)]
PkceNotSupported,
}
#[derive(Debug, Clone)]
pub struct DiscoveredAuthorizationServer {
pub issuer: String,
pub authorization_endpoint: String,
pub token_endpoint: Option<String>,
pub registration_endpoint: Option<String>,
pub jwks_uri: Option<String>,
pub revocation_endpoint: Option<String>,
pub introspection_endpoint: Option<String>,
pub scopes_supported: Option<Vec<String>>,
pub code_challenge_methods_supported: Vec<String>,
}
impl DiscoveredAuthorizationServer {
fn from_metadata(metadata: &AuthorizationServerMetadata) -> Result<Self, ClientDiscoveryError> {
if !metadata.supports_pkce_method("S256") {
return Err(ClientDiscoveryError::PkceNotSupported);
}
Ok(Self {
issuer: metadata.issuer.clone(),
authorization_endpoint: metadata.authorization_endpoint.clone(),
token_endpoint: metadata.token_endpoint.clone(),
registration_endpoint: metadata.registration_endpoint.clone(),
jwks_uri: metadata.jwks_uri.clone(),
revocation_endpoint: metadata.revocation_endpoint.clone(),
introspection_endpoint: metadata.introspection_endpoint.clone(),
scopes_supported: metadata.scopes_supported.clone(),
code_challenge_methods_supported: metadata
.code_challenge_methods_supported
.clone()
.unwrap_or_default(),
})
}
}
pub struct ClientDiscovery {
ssrf_validator: Arc<SsrfValidator>,
as_fetcher: DiscoveryFetcher,
}
impl ClientDiscovery {
pub fn new(ssrf_validator: SsrfValidator) -> Result<Self, FetcherError> {
let ssrf_validator = Arc::new(ssrf_validator);
let as_fetcher = DiscoveryFetcher::new((*ssrf_validator).clone())?;
Ok(Self {
ssrf_validator,
as_fetcher,
})
}
pub async fn discover_resource_metadata(
&self,
www_authenticate: Option<&str>,
server_url: &str,
) -> Result<ProtectedResourceMetadata, ClientDiscoveryError> {
if let Some(header) = www_authenticate
&& let Some(url) = parse_resource_metadata_param(header)
{
return self.fetch_resource_metadata_json(&url).await;
}
let candidates = well_known_resource_metadata_urls(server_url)
.map_err(ClientDiscoveryError::NoResourceMetadataCandidate)?;
let mut attempts = Vec::with_capacity(candidates.len());
for url in candidates {
match self.fetch_resource_metadata_json(&url).await {
Ok(metadata) => return Ok(metadata),
Err(e) => attempts.push((url, e.to_string())),
}
}
Err(ClientDiscoveryError::ResourceMetadataFetchFailed(attempts))
}
pub async fn discover_authorization_server(
&self,
resource: &ProtectedResourceMetadata,
) -> Result<DiscoveredAuthorizationServer, ClientDiscoveryError> {
let issuer = resource
.authorization_servers
.first()
.ok_or(ClientDiscoveryError::NoAuthorizationServer)?;
let metadata = self.as_fetcher.fetch(issuer).await?;
DiscoveredAuthorizationServer::from_metadata(metadata.oauth2())
}
pub async fn discover(
&self,
www_authenticate: Option<&str>,
server_url: &str,
) -> Result<DiscoveredAuthorizationServer, ClientDiscoveryError> {
let resource = self
.discover_resource_metadata(www_authenticate, server_url)
.await?;
self.discover_authorization_server(&resource).await
}
async fn fetch_resource_metadata_json(
&self,
url: &str,
) -> Result<ProtectedResourceMetadata, ClientDiscoveryError> {
let bytes = self
.ssrf_validator
.fetch(url)
.await
.map_err(|e| ClientDiscoveryError::ResourceMetadataInvalid(e.to_string()))?;
serde_json::from_slice(&bytes)
.map_err(|e| ClientDiscoveryError::ResourceMetadataInvalid(e.to_string()))
}
}
fn parse_resource_metadata_param(header_value: &str) -> Option<String> {
const KEY: &str = "resource_metadata=\"";
let start = header_value.find(KEY)? + KEY.len();
let rest = &header_value[start..];
let end = rest.find('"')?;
Some(rest[..end].to_string())
}
fn well_known_resource_metadata_urls(server_url: &str) -> Result<Vec<String>, String> {
let parsed = url::Url::parse(server_url)
.map_err(|e| format!("invalid server URL '{server_url}': {e}"))?;
let mut urls = Vec::with_capacity(2);
let path = parsed.path().trim_matches('/');
if !path.is_empty() {
let mut with_path = parsed.clone();
with_path.set_path(&format!("/.well-known/oauth-protected-resource/{path}"));
urls.push(with_path.to_string());
}
let mut root = parsed;
root.set_path("/.well-known/oauth-protected-resource");
urls.push(root.to_string());
Ok(urls)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_resource_metadata_param() {
let header = "Bearer resource_metadata=\"https://mcp.example.com/.well-known/oauth-protected-resource\", scope=\"files:read\"";
assert_eq!(
parse_resource_metadata_param(header),
Some("https://mcp.example.com/.well-known/oauth-protected-resource".to_string())
);
}
#[test]
fn test_parse_resource_metadata_param_absent() {
let header = "Bearer error=\"invalid_token\"";
assert_eq!(parse_resource_metadata_param(header), None);
}
#[test]
fn test_well_known_urls_no_path() {
let urls = well_known_resource_metadata_urls("https://mcp.example.com").unwrap();
assert_eq!(
urls,
vec!["https://mcp.example.com/.well-known/oauth-protected-resource"]
);
}
#[test]
fn test_well_known_urls_with_path_priority_order() {
let urls = well_known_resource_metadata_urls("https://example.com/public/mcp").unwrap();
assert_eq!(
urls,
vec![
"https://example.com/.well-known/oauth-protected-resource/public/mcp",
"https://example.com/.well-known/oauth-protected-resource",
]
);
}
#[test]
fn test_discovered_authorization_server_requires_s256_pkce() {
let mut metadata = sample_as_metadata();
metadata.code_challenge_methods_supported = None;
assert!(matches!(
DiscoveredAuthorizationServer::from_metadata(&metadata),
Err(ClientDiscoveryError::PkceNotSupported)
));
metadata.code_challenge_methods_supported = Some(vec!["plain".to_string()]);
assert!(matches!(
DiscoveredAuthorizationServer::from_metadata(&metadata),
Err(ClientDiscoveryError::PkceNotSupported)
));
metadata.code_challenge_methods_supported = Some(vec!["S256".to_string()]);
assert!(DiscoveredAuthorizationServer::from_metadata(&metadata).is_ok());
}
fn sample_as_metadata() -> AuthorizationServerMetadata {
AuthorizationServerMetadata {
issuer: "https://auth.example.com".to_string(),
authorization_endpoint: "https://auth.example.com/authorize".to_string(),
token_endpoint: Some("https://auth.example.com/token".to_string()),
jwks_uri: None,
registration_endpoint: None,
scopes_supported: None,
response_types_supported: vec!["code".to_string()],
response_modes_supported: None,
grant_types_supported: None,
token_endpoint_auth_methods_supported: None,
token_endpoint_auth_signing_alg_values_supported: None,
service_documentation: None,
ui_locales_supported: None,
op_policy_uri: None,
op_tos_uri: None,
revocation_endpoint: None,
revocation_endpoint_auth_methods_supported: None,
introspection_endpoint: None,
introspection_endpoint_auth_methods_supported: None,
code_challenge_methods_supported: Some(vec!["S256".to_string()]),
additional_fields: std::collections::HashMap::new(),
}
}
}