use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::error::{SaTokenError, SaTokenResult};
use crate::http_basic::ct_eq;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum CodeChallengeMethod {
S256,
Plain,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PkceChallenge {
pub code_challenge: String,
pub code_challenge_method: CodeChallengeMethod,
}
impl PkceChallenge {
pub fn from_verifier_s256(code_verifier: &str) -> SaTokenResult<Self> {
Self::validate_verifier_len(code_verifier)?;
let digest = Sha256::digest(code_verifier.as_bytes());
Ok(Self {
code_challenge: URL_SAFE_NO_PAD.encode(digest),
code_challenge_method: CodeChallengeMethod::S256,
})
}
fn validate_verifier_len(code_verifier: &str) -> SaTokenResult<()> {
if !(43..=128).contains(&code_verifier.len()) {
return Err(SaTokenError::OAuth2PkceMismatch);
}
Ok(())
}
pub fn verify(&self, code_verifier: &str) -> SaTokenResult<()> {
Self::validate_verifier_len(code_verifier)?;
let computed = match self.code_challenge_method {
CodeChallengeMethod::S256 => {
let digest = Sha256::digest(code_verifier.as_bytes());
URL_SAFE_NO_PAD.encode(digest)
}
CodeChallengeMethod::Plain => code_verifier.to_string(),
};
if ct_eq(computed.as_bytes(), self.code_challenge.as_bytes()) {
Ok(())
} else {
Err(SaTokenError::OAuth2PkceMismatch)
}
}
}