use axum_security_oauth2::HttpClient;
use serde::Deserialize;
use url::Url;
use crate::error::DiscoveryError;
#[derive(Debug, Clone, Deserialize)]
pub struct ProviderMetadata {
pub issuer: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub jwks_uri: String,
#[serde(default)]
pub userinfo_endpoint: Option<String>,
#[serde(default)]
pub end_session_endpoint: Option<String>,
#[serde(default)]
pub id_token_signing_alg_values_supported: Option<Vec<String>>,
#[serde(flatten)]
pub extra: serde_json::Map<String, serde_json::Value>,
}
impl ProviderMetadata {
pub async fn discover(
issuer_url: &str,
http: &HttpClient,
) -> Result<ProviderMetadata, DiscoveryError> {
let config_url = discovery_url(issuer_url)?;
let response = http.get(&config_url).await.map_err(DiscoveryError::Http)?;
if !response.is_success() {
return Err(DiscoveryError::Status(response.status));
}
let metadata: ProviderMetadata =
serde_json::from_slice(&response.body).map_err(DiscoveryError::Parse)?;
if !issuer_matches(issuer_url, &metadata.issuer) {
return Err(DiscoveryError::IssuerMismatch);
}
Ok(metadata)
}
}
fn discovery_url(issuer_url: &str) -> Result<Url, DiscoveryError> {
let base = issuer_url.trim_end_matches('/');
Url::parse(&format!("{base}/.well-known/openid-configuration"))
.map_err(DiscoveryError::InvalidIssuerUrl)
}
fn issuer_matches(requested: &str, returned: &str) -> bool {
requested.trim_end_matches('/') == returned.trim_end_matches('/')
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn builds_discovery_url() {
assert_eq!(
discovery_url("https://accounts.google.com")
.unwrap()
.as_str(),
"https://accounts.google.com/.well-known/openid-configuration"
);
assert_eq!(
discovery_url("https://accounts.google.com/")
.unwrap()
.as_str(),
"https://accounts.google.com/.well-known/openid-configuration"
);
}
#[test]
fn issuer_match_tolerates_trailing_slash() {
assert!(issuer_matches(
"https://issuer.example/",
"https://issuer.example"
));
assert!(issuer_matches(
"https://issuer.example",
"https://issuer.example/"
));
assert!(!issuer_matches(
"https://issuer.example",
"https://evil.example"
));
}
#[test]
fn deserializes_metadata_with_extra_fields() {
let json = r#"{
"issuer": "https://issuer.example",
"authorization_endpoint": "https://issuer.example/auth",
"token_endpoint": "https://issuer.example/token",
"jwks_uri": "https://issuer.example/jwks",
"end_session_endpoint": "https://issuer.example/logout",
"id_token_signing_alg_values_supported": ["RS256", "ES256"],
"scopes_supported": ["openid", "email"]
}"#;
let m: ProviderMetadata = serde_json::from_str(json).unwrap();
assert_eq!(m.issuer, "https://issuer.example");
assert_eq!(m.jwks_uri, "https://issuer.example/jwks");
assert_eq!(
m.end_session_endpoint.as_deref(),
Some("https://issuer.example/logout")
);
assert_eq!(
m.id_token_signing_alg_values_supported.as_deref(),
Some(&["RS256".to_string(), "ES256".to_string()][..])
);
assert!(m.userinfo_endpoint.is_none());
assert!(m.extra.contains_key("scopes_supported"));
}
}