use std::sync::OnceLock; use openidconnect::{ AccessTokenHash, AuthorizationCode, ClientId, ClientSecret, CsrfToken, EmptyExtraTokenFields, IdTokenFields, IssuerUrl, Nonce, OAuth2TokenResponse, PkceCodeChallenge, PkceCodeVerifier, RedirectUrl, Scope, StandardTokenResponse, TokenResponse, core::{ CoreAuthenticationFlow, CoreGenderClaim, CoreJweContentEncryptionAlgorithm, CoreJwsSigningAlgorithm, CoreProviderMetadata, CoreTokenType, }, reqwest, }; use serde::{Deserialize, Serialize}; use thiserror::Error; use tokio::sync::OnceCell; use tracing::{debug, error}; use crate::utils::config; #[derive(Error, Debug)] pub enum WhiskeyError { #[error("Internal error")] InternalError, #[error("Protocol error")] ProtocolError, } fn get_http_client() -> &'static reqwest::Client { static HTTP_CLIENT: OnceLock = OnceLock::new(); HTTP_CLIENT.get_or_init(|| { reqwest::ClientBuilder::new() .redirect(reqwest::redirect::Policy::none()) .build() .expect("Unable to build async http client") }) } async fn get_client() -> &'static Client { static CLIENT: OnceCell = OnceCell::const_new(); CLIENT .get_or_init(|| async { let http_client = get_http_client(); let config = config::get(); let redirect_url = format!("{}/whiskey/callback", config.get_base_url()); let config = config.clone(); let provider_metadata = CoreProviderMetadata::discover_async( IssuerUrl::new(config.oidc.issuer_url).unwrap(), http_client, ) .await .unwrap(); Client::from_provider_metadata( provider_metadata, ClientId::new(config.oidc.client_id), Some(ClientSecret::new(config.oidc.client_secret)), ) .set_redirect_uri(RedirectUrl::new(redirect_url).unwrap()) }) .await } #[derive(Debug, Serialize, Deserialize)] pub struct AuthorizeBackendData { pub csrf_token: String, pub pkce_verifier: PkceCodeVerifier, pub nonce: Nonce, } impl AuthorizeBackendData { pub fn csrf_token(&self) -> String { self.csrf_token.clone() } } pub async fn authorize() -> Result<(String, AuthorizeBackendData), WhiskeyError> { let client = get_client().await; let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256(); let (redirect_to, csrf_token, nonce) = client .authorize_url( CoreAuthenticationFlow::AuthorizationCode, CsrfToken::new_random, Nonce::new_random, ) .add_scope(Scope::new("profile".to_owned())) .add_scope(Scope::new("email".to_owned())) .set_pkce_challenge(pkce_challenge) .url(); Ok(( redirect_to.into(), AuthorizeBackendData { csrf_token: csrf_token.secret().to_owned(), pkce_verifier, nonce, }, )) } #[derive(Debug, Serialize, Clone)] pub struct UserInfoData { pub sub: String, pub sciper: String, pub name: String, pub firstname: String, pub email: String, pub groups: Option>, } pub async fn callback( code: String, state: String, backend_data: AuthorizeBackendData, ) -> Result { let client = get_client().await; if state != backend_data.csrf_token { error!("Wrong csrf token"); return Err(WhiskeyError::ProtocolError); } let token_response = client .exchange_code(AuthorizationCode::new(code)) .map_err(|err| { error!("Unable to build exchange code url: {err:?}"); WhiskeyError::InternalError })? .set_pkce_verifier(backend_data.pkce_verifier) .request_async(get_http_client()) .await .map_err(|err| match err { openidconnect::RequestTokenError::ServerResponse(err) => { error!("oidc protocol error: {err:?}"); WhiskeyError::ProtocolError } _ => { error!("other oidc server error: {err:?}"); WhiskeyError::InternalError } })?; let id_token = token_response.id_token().ok_or_else(|| { error!("missing id_token"); WhiskeyError::InternalError })?; let id_token_verifier = client.id_token_verifier(); let claims = id_token .claims(&id_token_verifier, &backend_data.nonce) .map_err(|err| { error!("unable to verify claims: {err:?}"); WhiskeyError::InternalError })?; if let Some(expected_access_token_hash) = claims.access_token_hash() { let actual_access_token_hash = AccessTokenHash::from_token( token_response.access_token(), id_token.signing_alg().map_err(|err| { error!("invalid signature algorithm in id_token: {err:?}"); WhiskeyError::InternalError })?, id_token.signing_key(&id_token_verifier).map_err(|err| { error!("invalid signature key in id_token: {err:?}"); WhiskeyError::InternalError })?, ) .map_err(|err| { error!("unable to compute access token hash: {err:?}"); WhiskeyError::InternalError })?; if actual_access_token_hash != *expected_access_token_hash { error!("actual_access_token_hash != expected_access_token_hash"); return Err(WhiskeyError::ProtocolError); } } else { error!("no hash in claims"); return Err(WhiskeyError::ProtocolError); } debug!("Whiskey id_token: {id_token:?}"); let firstname = claims .given_name() .ok_or_else(|| { error!("missing given_name claim from claims"); WhiskeyError::InternalError })? .get(None) .ok_or_else(|| { error!("missing given_name value from claims"); WhiskeyError::InternalError })? .to_string(); let name = claims .family_name() .ok_or_else(|| { error!("missing family_name claim from claims"); WhiskeyError::InternalError })? .get(None) .ok_or_else(|| { error!("missing family_name value from claims"); WhiskeyError::InternalError })? .to_string(); let email = claims .email() .ok_or_else(|| { error!("missing email claim from claims"); WhiskeyError::InternalError })? .to_string(); let sub = claims.subject().to_string(); let sciper = claims.additional_claims().sciper.clone(); let groups = match claims.additional_claims().groups.clone() { Some(groups) => Some(groups), None => { debug!("no groups claim in the id_token, trying the userinfo endpoint"); groups_from_userinfo(client, token_response.access_token()).await } }; debug!("Whiskey groups for {sciper}: {groups:?}"); Ok(UserInfoData { firstname, name, sub, sciper, email, groups, }) } async fn groups_from_userinfo( client: &Client, access_token: &openidconnect::AccessToken, ) -> Option> { let request = match client.user_info(access_token.to_owned(), None) { Ok(request) => request, Err(err) => { debug!("no userinfo endpoint advertised: {err:?}"); return None; } }; match request.request_async(get_http_client()).await.map( |claims: openidconnect::UserInfoClaims| { claims.additional_claims().groups.clone() }, ) { Ok(groups) => groups, Err(err) => { debug!("userinfo request failed: {err:?}"); None } } } #[derive(Serialize, Deserialize, Debug, PartialEq, Clone)] pub struct WhiskeyClaims { pub sciper: String, #[serde(default)] pub groups: Option>, } impl openidconnect::AdditionalClaims for WhiskeyClaims {} pub type WhiskeyFields = IdTokenFields< WhiskeyClaims, EmptyExtraTokenFields, CoreGenderClaim, CoreJweContentEncryptionAlgorithm, CoreJwsSigningAlgorithm, >; pub type WhiskeyTokenResponse = StandardTokenResponse; pub type Client = openidconnect::Client< WhiskeyClaims, openidconnect::core::CoreAuthDisplay, openidconnect::core::CoreGenderClaim, openidconnect::core::CoreJweContentEncryptionAlgorithm, openidconnect::core::CoreJsonWebKey, openidconnect::core::CoreAuthPrompt, openidconnect::StandardErrorResponse, WhiskeyTokenResponse, openidconnect::core::CoreTokenIntrospectionResponse, openidconnect::core::CoreRevocableToken, openidconnect::core::CoreRevocationErrorResponse, openidconnect::EndpointSet, openidconnect::EndpointNotSet, openidconnect::EndpointNotSet, openidconnect::EndpointNotSet, openidconnect::EndpointMaybeSet, openidconnect::EndpointMaybeSet, >; #[cfg(test)] mod tests { use super::WhiskeyClaims; fn parse(json: &str) -> WhiskeyClaims { serde_json::from_str(json).expect("claims should parse") } /// An absent claim must not be read as "belongs to no group": that would /// wipe the units of every user at every login. #[test] fn absent_groups_claim_is_none() { assert_eq!(parse(r#"{"sciper":"123456"}"#).groups, None); } #[test] fn groups_are_read() { assert_eq!( parse(r#"{"sciper":"123456","groups":["agepoly","balelec"]}"#).groups, Some(vec!["agepoly".to_owned(), "balelec".to_owned()]) ); } /// Distinct from the absent case: here Whiskey did answer, and the answer /// is that the user is in nothing. #[test] fn empty_groups_claim_is_some_empty() { assert_eq!( parse(r#"{"sciper":"123456","groups":[]}"#).groups, Some(vec![]) ); } /// The id_token carries plenty of claims we do not model #[test] fn unknown_claims_are_ignored() { assert_eq!( parse(r#"{"sciper":"123456","groups":["agepoly"],"uid":"x","other":42}"#).groups, Some(vec!["agepoly".to_owned()]) ); } #[test] fn a_null_groups_claim_is_none() { assert_eq!(parse(r#"{"sciper":"123456","groups":null}"#).groups, None); } }