use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use rand::Rng;
use sha2::{Digest, Sha256};
#[derive(Debug, Clone)]
pub(crate) struct Pkce {
pub(crate) verifier: String,
pub(crate) challenge: String,
pub(crate) state: String,
}
impl Pkce {
pub(crate) fn generate() -> Self {
let verifier = random_token(32);
let challenge = URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()));
let state = random_token(16);
Self {
verifier,
challenge,
state,
}
}
}
fn random_token(bytes: usize) -> String {
let mut buf = vec![0u8; bytes];
rand::rng().fill_bytes(&mut buf);
URL_SAFE_NO_PAD.encode(buf)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn verifier_is_within_the_spec_length_range() {
let len = Pkce::generate().verifier.len();
assert!(
(43..=128).contains(&len),
"verifier length {len} out of range"
);
}
#[test]
fn challenge_is_the_s256_of_the_verifier() {
let pkce = Pkce::generate();
let expected = URL_SAFE_NO_PAD.encode(Sha256::digest(pkce.verifier.as_bytes()));
assert_eq!(pkce.challenge, expected);
}
#[test]
fn tokens_are_url_safe_with_no_padding() {
let url_safe = |b: u8| b.is_ascii_alphanumeric() || b == b'-' || b == b'_';
assert!(url_safe(b'A'));
assert!(url_safe(b'9'));
assert!(url_safe(b'-'));
assert!(url_safe(b'_'));
assert!(!url_safe(b'!'));
let pkce = Pkce::generate();
for token in [&pkce.verifier, &pkce.challenge, &pkce.state] {
assert!(
token.bytes().all(url_safe),
"token is not url-safe: {token}"
);
}
}
#[test]
fn each_generation_is_unique() {
let a = Pkce::generate();
let b = Pkce::generate();
assert_ne!(a.verifier, b.verifier);
assert_ne!(a.state, b.state);
assert_ne!(a.challenge, b.challenge);
}
#[test]
fn random_token_length_scales_with_entropy() {
assert!(random_token(48).len() > random_token(16).len());
}
}