use std::{fmt, time::Duration};
use axum_security_oauth2::{CsrfToken, LoginOptions, OAuth2Client, random_token};
use url::Url;
use crate::{
claims::OidcClaims,
error::{ClaimsError, OidcError, VerifyError},
jwks::JwksCache,
logout::LogoutUrl,
verifier::VerifiedIdToken,
};
pub struct OidcClient {
oauth2: OAuth2Client,
verifier: JwksCache,
end_session_endpoint: Option<Url>,
}
impl OidcClient {
pub(crate) fn from_parts(
oauth2: OAuth2Client,
verifier: JwksCache,
end_session_endpoint: Option<Url>,
) -> Self {
Self {
oauth2,
verifier,
end_session_endpoint,
}
}
pub fn start_login(&self) -> OidcLogin {
let nonce = random_token();
let login = self
.oauth2
.start_login_with(LoginOptions::new().param("nonce", &nonce));
OidcLogin {
url: login.url,
csrf_token: login.csrf_token,
pkce_verifier: login.pkce_verifier,
nonce,
}
}
pub async fn finish_login(
&self,
code: &str,
pkce_verifier: &str,
nonce: &str,
) -> Result<OidcTokens, OidcError> {
let tokens = self
.oauth2
.finish_login(code, pkce_verifier)
.await
.map_err(OidcError::Exchange)?;
let id_token = tokens
.extra_field::<String>("id_token")
.ok_or(OidcError::NoIdToken)?;
let verified = self
.verifier
.verify(&id_token, nonce)
.await
.map_err(OidcError::Verify)?;
Ok(OidcTokens {
id_token: verified,
access_token: tokens.access_token,
refresh_token: tokens.refresh_token,
expires_in: tokens.expires_in,
})
}
pub async fn warm_jwks(&self) -> Result<(), VerifyError> {
self.verifier.warm().await
}
pub fn client_id(&self) -> &str {
self.oauth2.client_id()
}
pub fn redirect_url(&self) -> Option<&Url> {
self.oauth2.redirect_url()
}
pub fn scopes(&self) -> &[String] {
self.oauth2.scopes()
}
pub fn end_session_endpoint(&self) -> Option<&Url> {
self.end_session_endpoint.as_ref()
}
pub fn logout_url(&self) -> Option<LogoutUrl> {
self.end_session_endpoint
.as_ref()
.map(|endpoint| LogoutUrl::new(endpoint.clone()).client_id(self.oauth2.client_id()))
}
}
impl fmt::Debug for OidcClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OidcClient")
.field("oauth2", &self.oauth2)
.field(
"end_session_endpoint",
&self.end_session_endpoint.as_ref().map(Url::as_str),
)
.finish_non_exhaustive()
}
}
#[non_exhaustive]
pub struct OidcLogin {
pub url: Url,
pub csrf_token: CsrfToken,
pub pkce_verifier: String,
pub nonce: String,
}
impl fmt::Debug for OidcLogin {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let base = &self.url[..url::Position::AfterPath];
f.debug_struct("OidcLogin")
.field("url", &format_args!("{base}?[redacted]"))
.field("csrf_token", &"[redacted]")
.field("pkce_verifier", &"[redacted]")
.field("nonce", &"[redacted]")
.finish()
}
}
pub struct OidcTokens {
id_token: VerifiedIdToken,
access_token: String,
refresh_token: Option<String>,
expires_in: Option<Duration>,
}
impl OidcTokens {
pub fn claims(&self) -> Result<OidcClaims<'_>, ClaimsError> {
self.id_token.claims()
}
pub fn id_token(&self) -> &str {
self.id_token.id_token()
}
pub fn verified_id_token(&self) -> &VerifiedIdToken {
&self.id_token
}
pub fn access_token(&self) -> &str {
&self.access_token
}
pub fn refresh_token(&self) -> Option<&str> {
self.refresh_token.as_deref()
}
pub fn expires_in(&self) -> Option<Duration> {
self.expires_in
}
}
impl fmt::Debug for OidcTokens {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OidcTokens")
.field("id_token", &self.id_token)
.field("access_token", &"[redacted]")
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "[redacted]"),
)
.field("expires_in", &self.expires_in)
.finish()
}
}