use std::{sync::Arc, time::Duration as StdDuration};
pub(crate) const OAUTH_REQUEST_TIMEOUT: StdDuration = StdDuration::from_secs(30);
use std::fmt::Write as _;
use serde::{Deserialize, Serialize};
use zeroize::Zeroizing;
use super::{
super::jwks::{JwksCache, JwksError},
pkce::{PKCEChallenge, gen_random_token},
types::{IdTokenClaims, TokenResponse, UserInfo},
};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct OIDCProviderConfig {
pub issuer: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub userinfo_endpoint: Option<String>,
pub jwks_uri: String,
pub scopes_supported: Vec<String>,
pub response_types_supported: Vec<String>,
}
impl OIDCProviderConfig {
#[must_use]
pub fn new(
issuer: String,
authorization_endpoint: String,
token_endpoint: String,
jwks_uri: String,
) -> Self {
Self {
issuer,
authorization_endpoint,
token_endpoint,
userinfo_endpoint: None,
jwks_uri,
scopes_supported: vec![
"openid".to_string(),
"profile".to_string(),
"email".to_string(),
],
response_types_supported: vec!["code".to_string()],
}
}
}
#[derive(Debug, Clone)]
pub struct AuthorizationRequest {
pub url: String,
pub state: String,
pub pkce: Option<PKCEChallenge>,
pub nonce: Option<super::pkce::NonceParameter>,
}
#[derive(Clone)]
pub struct OAuth2Client {
pub client_id: String,
pub(crate) client_secret: Zeroizing<String>,
pub authorization_endpoint: String,
token_endpoint: String,
redirect_uri: Option<String>,
pub scopes: Vec<String>,
pub use_pkce: bool,
http_client: reqwest::Client,
}
#[allow(clippy::missing_fields_in_debug)] impl std::fmt::Debug for OAuth2Client {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuth2Client")
.field("client_id", &self.client_id)
.field("client_secret", &"[REDACTED]")
.field("authorization_endpoint", &self.authorization_endpoint)
.field("scopes", &self.scopes)
.field("use_pkce", &self.use_pkce)
.finish_non_exhaustive()
}
}
impl OAuth2Client {
pub(crate) const MAX_OAUTH_RESPONSE_BYTES: usize = 1024 * 1024;
pub fn new(
client_id: impl Into<String>,
client_secret: impl Into<String>,
authorization_endpoint: impl Into<String>,
token_endpoint: impl Into<String>,
) -> Self {
Self {
client_id: client_id.into(),
client_secret: Zeroizing::new(client_secret.into()),
authorization_endpoint: authorization_endpoint.into(),
token_endpoint: token_endpoint.into(),
redirect_uri: None,
scopes: vec![
"openid".to_string(),
"profile".to_string(),
"email".to_string(),
],
use_pkce: false,
http_client: reqwest::Client::builder()
.timeout(OAUTH_REQUEST_TIMEOUT)
.build()
.unwrap_or_default(),
}
}
pub fn with_redirect_uri(mut self, uri: impl Into<String>) -> Self {
self.redirect_uri = Some(uri.into());
self
}
#[must_use]
pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
self.scopes = scopes;
self
}
#[must_use]
pub const fn with_pkce(mut self, enabled: bool) -> Self {
self.use_pkce = enabled;
self
}
#[must_use]
pub fn authorization_url(&self, redirect_uri: &str) -> AuthorizationRequest {
let state = gen_random_token();
let scope = self.scopes.join(" ");
let mut url = format!(
"{}?client_id={}&redirect_uri={}&response_type=code&scope={}&state={}",
self.authorization_endpoint,
urlencoding::encode(&self.client_id),
urlencoding::encode(redirect_uri),
urlencoding::encode(&scope),
urlencoding::encode(&state),
);
let pkce = if self.use_pkce {
let challenge = PKCEChallenge::new();
let _ = write!(
url,
"&code_challenge={}&code_challenge_method=S256",
urlencoding::encode(&challenge.code_challenge),
);
Some(challenge)
} else {
None
};
AuthorizationRequest {
url,
state,
pkce,
nonce: None,
}
}
async fn post_token_request(&self, params: &[(&str, &str)]) -> Result<TokenResponse, String> {
let response = self
.http_client
.post(&self.token_endpoint)
.form(params)
.send()
.await
.map_err(|e| format!("Token request failed: {e}"))?;
let status = response.status();
let body_bytes = response
.bytes()
.await
.map_err(|e| format!("Failed to read token response body: {e}"))?;
if !status.is_success() {
let capped = &body_bytes[..body_bytes.len().min(Self::MAX_OAUTH_RESPONSE_BYTES)];
let body = String::from_utf8_lossy(capped);
return Err(format!("Token endpoint returned error: {body}"));
}
if body_bytes.len() > Self::MAX_OAUTH_RESPONSE_BYTES {
return Err(format!(
"Token response body too large ({} bytes, max {})",
body_bytes.len(),
Self::MAX_OAUTH_RESPONSE_BYTES
));
}
serde_json::from_slice::<TokenResponse>(&body_bytes)
.map_err(|e| format!("Failed to parse token response: {e}"))
}
pub async fn exchange_code(
&self,
code: &str,
redirect_uri: &str,
) -> Result<TokenResponse, String> {
if let Some(registered) = &self.redirect_uri {
if registered.trim_end_matches('/') != redirect_uri.trim_end_matches('/') {
return Err(format!(
"redirect_uri mismatch: supplied '{}' does not match registered '{}'",
redirect_uri, registered
));
}
}
let params = [
("grant_type", "authorization_code"),
("code", code),
("client_id", self.client_id.as_str()),
("client_secret", self.client_secret.as_str()),
("redirect_uri", redirect_uri),
];
self.post_token_request(¶ms).await
}
pub async fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse, String> {
let params = [
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
("client_id", self.client_id.as_str()),
("client_secret", self.client_secret.as_str()),
];
self.post_token_request(¶ms).await
}
}
pub const ALLOWED_OIDC_ALGORITHMS: &[jsonwebtoken::Algorithm] = &[
jsonwebtoken::Algorithm::RS256,
jsonwebtoken::Algorithm::RS384,
jsonwebtoken::Algorithm::RS512,
jsonwebtoken::Algorithm::ES256,
jsonwebtoken::Algorithm::ES384,
];
pub const FORBIDDEN_OIDC_ALGORITHMS: &[jsonwebtoken::Algorithm] = &[
jsonwebtoken::Algorithm::HS256,
jsonwebtoken::Algorithm::HS384,
jsonwebtoken::Algorithm::HS512,
];
pub const REQUIRED_JWT_TYP: &str = "JWT";
pub const FORBIDDEN_KEY_INJECTION_HEADERS: &[&str] = &["jku", "jwk", "x5u", "x5c"];
pub struct OIDCClient {
pub config: OIDCProviderConfig,
pub client_id: String,
#[allow(dead_code)] pub(crate) client_secret: Zeroizing<String>,
pub jwks_cache: Arc<JwksCache>,
http_client: reqwest::Client,
}
#[allow(clippy::missing_fields_in_debug)] impl std::fmt::Debug for OIDCClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OIDCClient")
.field("config", &self.config)
.field("client_id", &self.client_id)
.field("client_secret", &"[REDACTED]")
.finish_non_exhaustive()
}
}
impl OIDCClient {
pub(crate) const MAX_USERINFO_RESPONSE_BYTES: usize = 1024 * 1024;
pub fn new(
config: OIDCProviderConfig,
client_id: impl Into<String>,
client_secret: impl Into<String>,
) -> Result<Self, JwksError> {
let jwks_cache = Arc::new(JwksCache::new(&config.jwks_uri, StdDuration::from_secs(3600))?);
Ok(Self {
config,
client_id: client_id.into(),
client_secret: Zeroizing::new(client_secret.into()),
jwks_cache,
http_client: reqwest::Client::builder()
.timeout(OAUTH_REQUEST_TIMEOUT)
.build()
.unwrap_or_default(),
})
}
pub fn with_jwks_cache(
config: OIDCProviderConfig,
client_id: impl Into<String>,
client_secret: impl Into<String>,
jwks_cache: Arc<JwksCache>,
) -> Self {
Self {
config,
client_id: client_id.into(),
client_secret: Zeroizing::new(client_secret.into()),
jwks_cache,
http_client: reqwest::Client::builder()
.timeout(OAUTH_REQUEST_TIMEOUT)
.build()
.unwrap_or_default(),
}
}
#[must_use]
pub fn authorization_url(&self, redirect_uri: &str) -> AuthorizationRequest {
let state = gen_random_token();
let scope = self.config.scopes_supported.join(" ");
let nonce = super::pkce::NonceParameter::new();
let challenge = PKCEChallenge::new();
let url = format!(
"{}?client_id={}&redirect_uri={}&response_type=code&scope={}&state={}\
&nonce={}&code_challenge={}&code_challenge_method=S256",
self.config.authorization_endpoint,
urlencoding::encode(&self.client_id),
urlencoding::encode(redirect_uri),
urlencoding::encode(&scope),
urlencoding::encode(&state),
urlencoding::encode(&nonce.nonce),
urlencoding::encode(&challenge.code_challenge),
);
AuthorizationRequest {
url,
state,
pkce: Some(challenge),
nonce: Some(nonce),
}
}
pub async fn verify_id_token(
&self,
id_token: &str,
expected_nonce: Option<&str>,
max_age_secs: Option<u64>,
) -> Result<IdTokenClaims, String> {
let header = jsonwebtoken::decode_header(id_token)
.map_err(|e| format!("Invalid JWT header: {e}"))?;
if FORBIDDEN_OIDC_ALGORITHMS.contains(&header.alg) {
let alg_str = format!("{:?}", header.alg);
return Err(format!("Forbidden OIDC algorithm: {alg_str}"));
}
if !ALLOWED_OIDC_ALGORITHMS.contains(&header.alg) {
let alg_str = format!("{:?}", header.alg);
return Err(format!("OIDC algorithm not in allowlist: {alg_str}"));
}
if let Some(ref typ) = header.typ {
if typ.to_uppercase() != REQUIRED_JWT_TYP {
return Err(format!(
"Unexpected JWT typ header '{typ}': expected '{REQUIRED_JWT_TYP}'"
));
}
}
if header.jku.is_some()
|| header.jwk.is_some()
|| header.x5u.is_some()
|| header.x5c.is_some()
{
return Err("JWT header contains forbidden key-injection parameter".to_string());
}
let kid = header.kid.ok_or("JWT missing 'kid' in header")?;
let key = self
.jwks_cache
.get_key(&kid)
.await
.map_err(|e| format!("JWKS fetch error: {e}"))?
.ok_or_else(|| format!("No key found for kid '{kid}'"))?;
let mut validation = jsonwebtoken::Validation::new(header.alg);
validation.set_issuer(&[&self.config.issuer]);
validation.set_audience(&[&self.client_id]);
validation.set_required_spec_claims(&["exp", "iat", "iss", "aud", "sub"]);
let token_data = jsonwebtoken::decode::<IdTokenClaims>(id_token, &key, &validation)
.map_err(|e| format!("ID token validation failed: {e}"))?;
let claims = token_data.claims;
claims
.validate_temporal_claims()
.map_err(|e| format!("ID token temporal validation failed: {e}"))?;
if let Some(expected) = expected_nonce {
super::claims_validator::validate_nonce_claim(&claims, expected)
.map_err(|e| e.to_string())?;
}
if let Some(max_age) = max_age_secs {
let now_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(i64::MAX, |d| i64::try_from(d.as_secs()).unwrap_or(i64::MAX));
super::claims_validator::validate_auth_time_claim(&claims, max_age, now_secs)
.map_err(|e| e.to_string())?;
}
Ok(claims)
}
pub async fn get_userinfo(&self, access_token: &str) -> Result<UserInfo, String> {
let endpoint = self
.config
.userinfo_endpoint
.as_ref()
.ok_or("No userinfo endpoint configured for this provider")?;
let response = self
.http_client
.get(endpoint)
.bearer_auth(access_token)
.send()
.await
.map_err(|e| format!("Userinfo request failed: {e}"))?;
if !response.status().is_success() {
return Err(format!("Userinfo endpoint returned {}", response.status()));
}
let body = response
.bytes()
.await
.map_err(|e| format!("Failed to read userinfo response: {e}"))?;
if body.len() > Self::MAX_USERINFO_RESPONSE_BYTES {
return Err(format!(
"Userinfo response too large ({} bytes, max {})",
body.len(),
Self::MAX_USERINFO_RESPONSE_BYTES
));
}
serde_json::from_slice::<UserInfo>(&body)
.map_err(|e| format!("Failed to parse userinfo response: {e}"))
}
}