use serde::{Deserialize, Serialize};
use url::Url;
use crate::error::{DiscoveryError, SsrfError};
use crate::identity::IdentityResolver;
use crate::ssrf::{SsrfFilter, MAX_OAUTH_RESPONSE_BYTES};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProtectedResourceMetadata {
pub resource: String,
pub authorization_servers: Vec<String>,
#[serde(default)]
pub scopes_supported: Vec<String>,
#[serde(default)]
pub bearer_methods_supported: Vec<String>,
#[serde(default)]
pub resource_documentation: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct AuthorizationServerMetadata {
pub issuer: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
#[serde(default)]
pub pushed_authorization_request_endpoint: String,
#[serde(default)]
pub require_pushed_authorization_requests: bool,
#[serde(default)]
pub dpop_signing_alg_values_supported: Vec<String>,
#[serde(default)]
pub code_challenge_methods_supported: Vec<String>,
#[serde(default)]
pub response_types_supported: Vec<String>,
#[serde(default)]
pub grant_types_supported: Vec<String>,
#[serde(default)]
pub token_endpoint_auth_methods_supported: Vec<String>,
#[serde(default)]
pub token_endpoint_auth_signing_alg_values_supported: Vec<String>,
#[serde(default)]
pub scopes_supported: Vec<String>,
#[serde(default)]
pub authorization_response_iss_parameter_supported: bool,
#[serde(default)]
pub client_id_metadata_document_supported: bool,
#[serde(default = "crate::discovery::default_true")]
pub require_request_uri_registration: bool,
}
fn default_true() -> bool {
true
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DiscoveredAuthEndpoints {
pub did: String,
pub handle: Option<String>,
pub pds_endpoint: String,
pub auth_server_issuer: String,
pub par_endpoint: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub dpop_algs: Vec<String>,
pub scopes: Vec<String>,
pub protected_resource_metadata: ProtectedResourceMetadata,
pub auth_server_metadata: AuthorizationServerMetadata,
}
fn is_origin_only(url_str: &str) -> bool {
if let Ok(parsed) = Url::parse(url_str) {
let scheme = parsed.scheme();
let host = parsed.host_str().unwrap_or("");
let is_loopback =
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]";
if scheme != "https" && !(scheme == "http" && is_loopback) {
return false;
}
if !parsed.username().is_empty() || parsed.password().is_some() {
return false;
}
let after_scheme = url_str
.split_once("://")
.map(|(_, rest)| rest)
.unwrap_or(url_str);
let authority = after_scheme
.split(['/', '?', '#'])
.next()
.unwrap_or(after_scheme);
let host_port = authority.rsplit('@').next().unwrap_or(authority);
let has_explicit_port = match host_port.rfind(']') {
Some(bracket_end) => host_port[bracket_end..].contains(':'),
None => host_port.contains(':'),
};
if has_explicit_port {
let port_str = host_port.rsplit(':').next().unwrap_or("");
match port_str.parse::<u16>() {
Ok(443) if scheme == "https" => return false,
Ok(80) if scheme == "http" => return false,
Ok(_) => {}
Err(_) => return false,
}
}
if host_port.contains('\\') {
return false;
}
let path = parsed.path();
(path.is_empty() || path == "/") && parsed.query().is_none() && parsed.fragment().is_none()
} else {
false
}
}
fn normalize_origin(url_str: &str) -> String {
if let Ok(parsed) = Url::parse(url_str) {
parsed.origin().ascii_serialization()
} else {
url_str.trim().trim_end_matches('/').to_string()
}
}
pub async fn fetch_protected_resource_metadata(
ssrf_filter: &SsrfFilter,
pds_endpoint: &str,
) -> Result<ProtectedResourceMetadata, DiscoveryError> {
let url = format!(
"{}/.well-known/oauth-protected-resource",
pds_endpoint.trim_end_matches('/')
);
let meta: ProtectedResourceMetadata = ssrf_filter
.safe_get_json(&url, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| match e {
SsrfError::HttpStatus(status, msg) => DiscoveryError::ProtectedResourceDiscoveryFailed(
format!("HTTP {status} fetching protected resource metadata from {url}: {msg}"),
),
SsrfError::Json(err) => DiscoveryError::ProtectedResourceDiscoveryFailed(format!(
"Invalid JSON in protected resource metadata from {url}: {err}"
)),
other => DiscoveryError::Ssrf(other),
})?;
if meta.authorization_servers.is_empty() {
return Err(DiscoveryError::MissingAuthorizationServers(
pds_endpoint.to_string(),
));
}
if meta.authorization_servers.len() > 1 {
return Err(DiscoveryError::MultipleAuthorizationServers(
meta.authorization_servers.len(),
));
}
let as_url = &meta.authorization_servers[0];
if !is_origin_only(as_url) {
return Err(DiscoveryError::InvalidAuthorizationServerUrl(
as_url.clone(),
));
}
let expected_origin = normalize_origin(pds_endpoint);
let actual_origin = normalize_origin(&meta.resource);
if expected_origin != actual_origin {
return Err(DiscoveryError::ResourceMismatch {
expected: expected_origin,
actual: meta.resource.clone(),
});
}
if !is_origin_only(&meta.resource) {
return Err(DiscoveryError::ResourceMismatch {
expected: expected_origin,
actual: meta.resource.clone(),
});
}
Ok(meta)
}
pub async fn fetch_auth_server_metadata(
ssrf_filter: &SsrfFilter,
auth_server_url: &str,
) -> Result<AuthorizationServerMetadata, DiscoveryError> {
let base = auth_server_url.trim_end_matches('/');
let primary_url = format!("{base}/.well-known/oauth-authorization-server");
let meta: AuthorizationServerMetadata = match ssrf_filter
.safe_get_json(&primary_url, MAX_OAUTH_RESPONSE_BYTES)
.await
{
Ok(m) => m,
Err(SsrfError::HttpStatus(404, _)) => {
let fallback_url = format!("{base}/.well-known/openid-configuration");
ssrf_filter
.safe_get_json(&fallback_url, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| match e {
SsrfError::HttpStatus(status, msg) => {
DiscoveryError::AuthServerDiscoveryFailed(format!(
"HTTP {status} fetching fallback authorization server metadata from {fallback_url}: {msg}"
))
}
SsrfError::Json(err) => DiscoveryError::AuthServerDiscoveryFailed(format!(
"Invalid JSON in fallback authorization server metadata from {fallback_url}: {err}"
)),
other => DiscoveryError::Ssrf(other),
})?
}
Err(SsrfError::HttpStatus(status, msg)) => {
return Err(DiscoveryError::AuthServerDiscoveryFailed(format!(
"HTTP {status} fetching authorization server metadata from {primary_url}: {msg}"
)));
}
Err(SsrfError::Json(err)) => {
return Err(DiscoveryError::AuthServerDiscoveryFailed(format!(
"Invalid JSON in authorization server metadata from {primary_url}: {err}"
)));
}
Err(other) => return Err(DiscoveryError::Ssrf(other)),
};
validate_auth_server_capabilities(&meta, auth_server_url)?;
Ok(meta)
}
pub fn validate_auth_server_capabilities(
meta: &AuthorizationServerMetadata,
auth_server_url: &str,
) -> Result<(), DiscoveryError> {
if !is_origin_only(auth_server_url) {
return Err(DiscoveryError::InvalidAuthorizationServerUrl(
auth_server_url.to_string(),
));
}
if !is_origin_only(&meta.issuer) {
return Err(DiscoveryError::InvalidAuthorizationServerUrl(
meta.issuer.clone(),
));
}
let expected_origin = normalize_origin(auth_server_url);
let actual_origin = normalize_origin(&meta.issuer);
if expected_origin != actual_origin {
return Err(DiscoveryError::IssuerMismatch {
expected: auth_server_url.to_string(),
actual: meta.issuer.clone(),
});
}
if meta.pushed_authorization_request_endpoint.trim().is_empty() {
return Err(DiscoveryError::MissingParEndpoint(
auth_server_url.to_string(),
));
}
validate_https_endpoint(
&meta.pushed_authorization_request_endpoint,
"PAR",
auth_server_url,
)?;
if !meta.require_pushed_authorization_requests {
return Err(DiscoveryError::ParNotRequired(auth_server_url.to_string()));
}
if meta.token_endpoint.trim().is_empty() {
return Err(DiscoveryError::MissingTokenEndpoint(
auth_server_url.to_string(),
));
}
validate_https_endpoint(&meta.token_endpoint, "token", auth_server_url)?;
if meta.authorization_endpoint.trim().is_empty() {
return Err(DiscoveryError::MissingAuthorizationEndpoint(
auth_server_url.to_string(),
));
}
validate_https_endpoint(
&meta.authorization_endpoint,
"authorization",
auth_server_url,
)?;
if !meta.response_types_supported.iter().any(|r| r == "code") {
return Err(DiscoveryError::MissingResponseType(
auth_server_url.to_string(),
));
}
if !meta
.grant_types_supported
.iter()
.any(|g| g == "authorization_code")
{
return Err(DiscoveryError::MissingGrantType {
auth_server: auth_server_url.to_string(),
missing: "authorization_code".to_string(),
});
}
if !meta
.grant_types_supported
.iter()
.any(|g| g == "refresh_token")
{
return Err(DiscoveryError::MissingGrantType {
auth_server: auth_server_url.to_string(),
missing: "refresh_token".to_string(),
});
}
let has_none = meta
.token_endpoint_auth_methods_supported
.iter()
.any(|m| m == "none");
let has_private_key_jwt = meta
.token_endpoint_auth_methods_supported
.iter()
.any(|m| m == "private_key_jwt");
if !has_none || !has_private_key_jwt {
return Err(DiscoveryError::MissingTokenAuthMethod(
auth_server_url.to_string(),
));
}
if !meta
.token_endpoint_auth_signing_alg_values_supported
.iter()
.any(|alg| alg == "ES256")
{
return Err(DiscoveryError::MissingTokenAuthSigningAlg(
auth_server_url.to_string(),
));
}
if meta
.token_endpoint_auth_signing_alg_values_supported
.iter()
.any(|alg| alg == "none")
{
return Err(DiscoveryError::InvalidTokenAuthSigningAlg(
auth_server_url.to_string(),
));
}
if !meta.scopes_supported.iter().any(|s| s == "atproto") {
return Err(DiscoveryError::MissingAtprotoScope(
auth_server_url.to_string(),
));
}
if !meta.authorization_response_iss_parameter_supported {
return Err(DiscoveryError::MissingIssParameterSupport(
auth_server_url.to_string(),
));
}
if !meta.client_id_metadata_document_supported {
return Err(DiscoveryError::MissingClientMetadataSupport(
auth_server_url.to_string(),
));
}
if !meta.require_request_uri_registration {
return Err(DiscoveryError::MissingRequestUriRegistration(
auth_server_url.to_string(),
));
}
if !meta
.dpop_signing_alg_values_supported
.iter()
.any(|alg| alg == "ES256")
{
return Err(DiscoveryError::MissingDpopAlgorithm(
auth_server_url.to_string(),
));
}
if !meta
.code_challenge_methods_supported
.iter()
.any(|method| method == "S256")
{
return Err(DiscoveryError::MissingPkceMethod(
auth_server_url.to_string(),
));
}
Ok(())
}
pub async fn discover_oauth_endpoints(
resolver: &IdentityResolver,
did_or_handle: &str,
) -> Result<DiscoveredAuthEndpoints, DiscoveryError> {
let identity = resolver.resolve_ident(did_or_handle).await?;
let pds_meta =
fetch_protected_resource_metadata(resolver.ssrf_filter(), &identity.pds_endpoint).await?;
let auth_server_url = &pds_meta.authorization_servers[0];
let as_meta = fetch_auth_server_metadata(resolver.ssrf_filter(), auth_server_url).await?;
Ok(DiscoveredAuthEndpoints {
did: identity.did,
handle: identity.handle,
pds_endpoint: identity.pds_endpoint,
auth_server_issuer: as_meta.issuer.clone(),
par_endpoint: as_meta.pushed_authorization_request_endpoint.clone(),
authorization_endpoint: as_meta.authorization_endpoint.clone(),
token_endpoint: as_meta.token_endpoint.clone(),
dpop_algs: as_meta.dpop_signing_alg_values_supported.clone(),
scopes: if !as_meta.scopes_supported.is_empty() {
as_meta.scopes_supported.clone()
} else if !pds_meta.scopes_supported.is_empty() {
pds_meta.scopes_supported.clone()
} else {
vec!["atproto".to_string()]
},
protected_resource_metadata: pds_meta,
auth_server_metadata: as_meta,
})
}
impl IdentityResolver {
pub async fn discover_oauth_endpoints(
&self,
did_or_handle: &str,
) -> Result<DiscoveredAuthEndpoints, DiscoveryError> {
discover_oauth_endpoints(self, did_or_handle).await
}
}
fn validate_https_endpoint(
endpoint: &str,
label: &str,
issuer: &str,
) -> Result<(), DiscoveryError> {
let _ = issuer;
let parsed = Url::parse(endpoint).map_err(|e| {
DiscoveryError::InvalidEndpointUrl(format!(
"Invalid {label} endpoint URL '{endpoint}': {e}"
))
})?;
if parsed.scheme() != "https" {
let host = parsed.host_str().unwrap_or_default();
let host_is_local =
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]";
let scheme_is_http = parsed.scheme() == "http";
if !(scheme_is_http && host_is_local) {
return Err(DiscoveryError::InvalidEndpointUrl(format!(
"{label} endpoint URL must be HTTPS (got '{}'; only loopback/localhost may use HTTP)",
parsed.scheme()
)));
}
}
if !parsed.username().is_empty() || parsed.password().is_some() {
return Err(DiscoveryError::InvalidEndpointUrl(format!(
"{label} endpoint URL must not contain embedded credentials: '{endpoint}'"
)));
}
if parsed.host_str().is_none() {
return Err(DiscoveryError::InvalidEndpointUrl(format!(
"{label} endpoint URL is missing a host: '{endpoint}'"
)));
}
Ok(())
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod tests {
use super::*;
fn valid_test_metadata() -> AuthorizationServerMetadata {
AuthorizationServerMetadata {
issuer: "https://auth.example.com".to_string(),
authorization_endpoint: "https://auth.example.com/oauth/authorize".to_string(),
token_endpoint: "https://auth.example.com/oauth/token".to_string(),
pushed_authorization_request_endpoint: "https://auth.example.com/oauth/par".to_string(),
require_pushed_authorization_requests: true,
dpop_signing_alg_values_supported: vec!["ES256".to_string()],
code_challenge_methods_supported: vec!["S256".to_string()],
response_types_supported: vec!["code".to_string()],
grant_types_supported: vec![
"authorization_code".to_string(),
"refresh_token".to_string(),
],
token_endpoint_auth_methods_supported: vec![
"none".to_string(),
"private_key_jwt".to_string(),
],
token_endpoint_auth_signing_alg_values_supported: vec!["ES256".to_string()],
scopes_supported: vec!["atproto".to_string()],
authorization_response_iss_parameter_supported: true,
client_id_metadata_document_supported: true,
require_request_uri_registration: true,
}
}
#[test]
fn test_validate_auth_server_capabilities_valid() {
let meta = valid_test_metadata();
assert!(validate_auth_server_capabilities(&meta, "https://auth.example.com").is_ok());
}
#[test]
fn test_validate_auth_server_capabilities_missing_es256() {
let meta = AuthorizationServerMetadata {
dpop_signing_alg_values_supported: vec!["RS256".to_string()],
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::MissingDpopAlgorithm(_))
));
}
#[test]
fn test_validate_auth_server_capabilities_missing_s256_pkce() {
let meta = AuthorizationServerMetadata {
code_challenge_methods_supported: vec!["plain".to_string()],
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::MissingPkceMethod(_))
));
}
#[test]
fn test_validate_auth_server_capabilities_explicit_false_request_uri_registration() {
let meta = AuthorizationServerMetadata {
require_request_uri_registration: false,
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::MissingRequestUriRegistration(_))
));
}
#[test]
fn test_auth_server_metadata_omitted_request_uri_registration_defaults_true() {
let json = serde_json::json!({
"issuer": "https://auth.example.com",
"authorization_endpoint": "https://auth.example.com/oauth/authorize",
"token_endpoint": "https://auth.example.com/oauth/token",
"pushed_authorization_request_endpoint": "https://auth.example.com/oauth/par",
"require_pushed_authorization_requests": true,
"dpop_signing_alg_values_supported": ["ES256"],
"code_challenge_methods_supported": ["S256"],
"response_types_supported": ["code"],
"grant_types_supported": ["authorization_code", "refresh_token"],
"token_endpoint_auth_methods_supported": ["none", "private_key_jwt"],
"token_endpoint_auth_signing_alg_values_supported": ["ES256"],
"scopes_supported": ["atproto"],
"authorization_response_iss_parameter_supported": true,
"client_id_metadata_document_supported": true
});
let meta: AuthorizationServerMetadata = serde_json::from_value(json).unwrap();
assert!(meta.require_request_uri_registration);
}
#[test]
fn test_auth_server_metadata_explicit_false_request_uri_registration_deserializes() {
let json = serde_json::json!({
"issuer": "https://auth.example.com",
"authorization_endpoint": "https://auth.example.com/oauth/authorize",
"token_endpoint": "https://auth.example.com/oauth/token",
"require_request_uri_registration": false
});
let meta: AuthorizationServerMetadata = serde_json::from_value(json).unwrap();
assert!(!meta.require_request_uri_registration);
}
#[test]
fn test_validate_auth_server_capabilities_issuer_mismatch() {
let meta = AuthorizationServerMetadata {
issuer: "https://attacker.example.com".to_string(),
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::IssuerMismatch { .. })
));
}
#[test]
fn test_validate_auth_server_capabilities_missing_token_auth_signing_alg() {
let meta = AuthorizationServerMetadata {
token_endpoint_auth_signing_alg_values_supported: vec!["RS256".to_string()],
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::MissingTokenAuthSigningAlg(_))
));
}
#[test]
fn test_validate_auth_server_capabilities_invalid_token_auth_signing_alg_none() {
let meta = AuthorizationServerMetadata {
token_endpoint_auth_signing_alg_values_supported: vec![
"ES256".to_string(),
"none".to_string(),
],
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::InvalidTokenAuthSigningAlg(_))
));
}
#[test]
fn test_validate_auth_server_capabilities_explicit_443_rejected() {
let meta = AuthorizationServerMetadata {
issuer: "https://auth.example.com:443".to_string(),
..valid_test_metadata()
};
assert!(matches!(
validate_auth_server_capabilities(&meta, "https://auth.example.com"),
Err(DiscoveryError::InvalidAuthorizationServerUrl(_))
));
let valid_meta = valid_test_metadata();
assert!(matches!(
validate_auth_server_capabilities(&valid_meta, "https://auth.example.com:443"),
Err(DiscoveryError::InvalidAuthorizationServerUrl(_))
));
}
#[test]
fn test_is_origin_only_rejects_explicit_443() {
assert!(is_origin_only("https://auth.example.com"));
assert!(!is_origin_only("https://auth.example.com:443"));
assert!(!is_origin_only("https://auth.example.com:443/"));
assert!(is_origin_only("https://auth.example.com:8443"));
}
#[test]
fn test_is_origin_only_rejects_leading_zero_and_malformed_ports() {
assert!(!is_origin_only("https://auth.example.com:0443"));
assert!(!is_origin_only("http://auth.example.com:0080"));
assert!(!is_origin_only("https://auth.example.com:not_a_port"));
assert!(!is_origin_only("https://auth.example.com:443\\"));
assert!(is_origin_only("https://auth.example.com:44371"));
assert!(is_origin_only("http://127.0.0.1:8080"));
}
#[test]
fn test_is_origin_only_loopback_http_acceptance_boundaries() {
assert!(is_origin_only("http://localhost"));
assert!(is_origin_only("http://127.0.0.1"));
assert!(is_origin_only("http://127.0.0.1:8080"));
assert!(is_origin_only("http://[::1]:8080"));
assert!(!is_origin_only("http://auth.example.com"));
assert!(!is_origin_only("http://auth.example.com:80"));
assert!(!is_origin_only("http://auth.example.com:8080"));
assert!(!is_origin_only("http://127.0.0.2")); assert!(!is_origin_only("http://user@127.0.0.1"));
assert!(!is_origin_only("http://127.0.0.1/xrpc"));
assert!(!is_origin_only("https://auth.example.com/?a=b"));
assert!(!is_origin_only("https://auth.example.com/#frag"));
assert!(is_origin_only("https://auth.example.com:8080"));
assert!(is_origin_only("https://auth.example.com:44371"));
}
}