use rand::RngCore;
use serde::{Deserialize, Serialize};
use crate::crypto::{base64url_encode, constant_time_eq, sha256_digest};
use crate::error::PkceError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
pub enum PkceMethod {
#[default]
#[serde(rename = "S256")]
S256,
}
impl PkceMethod {
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::S256 => "S256",
}
}
}
impl std::fmt::Display for PkceMethod {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct PkcePair {
pub verifier: String,
pub challenge: String,
pub method: PkceMethod,
}
impl PkcePair {
#[must_use]
pub fn generate() -> Self {
let mut entropy = [0u8; 32];
rand::thread_rng().fill_bytes(&mut entropy);
let verifier = base64url_encode(&entropy);
let challenge = derive_s256_challenge(&verifier);
Self {
verifier,
challenge,
method: PkceMethod::S256,
}
}
pub fn generate_with_entropy_size(entropy_bytes: usize) -> Result<Self, PkceError> {
if !(32..=96).contains(&entropy_bytes) {
return Err(PkceError::InvalidVerifierLength {
len: (entropy_bytes * 4).div_ceil(3),
min: 43,
max: 128,
});
}
let mut bytes = vec![0u8; entropy_bytes];
rand::thread_rng().fill_bytes(&mut bytes);
let verifier = base64url_encode(&bytes);
Self::from_verifier(verifier)
}
pub fn from_verifier(verifier: String) -> Result<Self, PkceError> {
validate_verifier(&verifier)?;
let challenge = derive_s256_challenge(&verifier);
Ok(Self {
verifier,
challenge,
method: PkceMethod::S256,
})
}
pub fn verify(&self, candidate_verifier: &str) -> Result<(), PkceError> {
verify_pkce(candidate_verifier, &self.challenge)
}
}
#[must_use]
pub fn derive_s256_challenge(verifier: &str) -> String {
let digest = sha256_digest(verifier.as_bytes());
base64url_encode(&digest)
}
pub fn validate_verifier(verifier: &str) -> Result<(), PkceError> {
use crate::kernels::pkce_bytes::{validate_verifier_bytes, VerifierByteError};
match validate_verifier_bytes(verifier.as_bytes()) {
Ok(()) => Ok(()),
Err(VerifierByteError::InvalidLength(len)) => Err(PkceError::InvalidVerifierLength {
len,
min: crate::kernels::pkce_bytes::VERIFIER_MIN_LENGTH,
max: crate::kernels::pkce_bytes::VERIFIER_MAX_LENGTH,
}),
Err(VerifierByteError::InvalidCharacter { byte, position }) => {
Err(PkceError::InvalidVerifierCharacter {
char: byte as char,
position,
})
}
}
}
pub fn verify_pkce(verifier: &str, expected_challenge: &str) -> Result<(), PkceError> {
validate_verifier(verifier)?;
if expected_challenge.len() != 43 {
return Err(PkceError::InvalidChallengeLength {
len: expected_challenge.len(),
});
}
let derived = derive_s256_challenge(verifier);
if constant_time_eq(derived.as_bytes(), expected_challenge.as_bytes()) {
Ok(())
} else {
Err(PkceError::ChallengeMismatch)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod tests {
use super::*;
use proptest::prelude::*;
#[test]
fn test_rfc7636_appendix_b_vector() {
let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
let expected_challenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM";
let pkce = PkcePair::from_verifier(verifier.to_string()).unwrap();
assert_eq!(pkce.challenge, expected_challenge);
assert_eq!(pkce.method, PkceMethod::S256);
assert!(verify_pkce(verifier, expected_challenge).is_ok());
}
#[test]
fn test_pkce_generate_roundtrip() {
let pkce = PkcePair::generate();
assert_eq!(pkce.verifier.len(), 43);
assert_eq!(pkce.challenge.len(), 43);
assert!(pkce.verify(&pkce.verifier).is_ok());
assert!(verify_pkce(&pkce.verifier, &pkce.challenge).is_ok());
}
#[test]
fn test_invalid_verifier_length() {
let short = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjX";
assert_eq!(short.len(), 42);
assert!(matches!(
validate_verifier(short),
Err(PkceError::InvalidVerifierLength { len: 42, .. })
));
let long = "a".repeat(129);
assert!(matches!(
validate_verifier(&long),
Err(PkceError::InvalidVerifierLength { len: 129, .. })
));
}
#[test]
fn test_invalid_verifier_characters() {
let with_space = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEj k";
assert!(matches!(
validate_verifier(with_space),
Err(PkceError::InvalidVerifierCharacter {
char: ' ',
position: 41
})
));
let with_plus = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEj+k";
assert!(matches!(
validate_verifier(with_plus),
Err(PkceError::InvalidVerifierCharacter {
char: '+',
position: 41
})
));
let with_eq = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEj=k";
assert!(matches!(
validate_verifier(with_eq),
Err(PkceError::InvalidVerifierCharacter {
char: '=',
position: 41
})
));
}
#[test]
fn test_challenge_mismatch() {
let pkce = PkcePair::generate();
let wrong_challenge = "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA";
assert!(matches!(
verify_pkce(&pkce.verifier, wrong_challenge),
Err(PkceError::ChallengeMismatch)
));
}
proptest! {
#[test]
fn prop_valid_verifier_always_verifies(
verifier in "[A-Za-z0-9\\-._~]{43,128}"
) {
let challenge = derive_s256_challenge(&verifier);
prop_assert!(verify_pkce(&verifier, &challenge).is_ok());
}
#[test]
fn prop_mutated_challenge_fails_verification(
verifier in "[A-Za-z0-9\\-._~]{43,128}",
mutate_idx in 0usize..43
) {
let mut challenge_chars: Vec<char> = derive_s256_challenge(&verifier).chars().collect();
let original_char = challenge_chars[mutate_idx];
let replacement = if original_char == 'A' { 'B' } else { 'A' };
challenge_chars[mutate_idx] = replacement;
let mutated_challenge: String = challenge_chars.into_iter().collect();
prop_assert!(verify_pkce(&verifier, &mutated_challenge).is_err());
}
}
}