use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use rand::Rng;
use sha2::{Digest, Sha256};
use crate::domain::errors::AuthError;
pub const VERIFIER_MIN_LEN: usize = 43;
pub const VERIFIER_MAX_LEN: usize = 128;
const VERIFIER_BYTES: usize = 32;
const STATE_BYTES: usize = 32;
#[derive(Debug, Clone)]
pub struct CodeVerifier(String);
impl CodeVerifier {
pub fn new() -> Self {
let mut bytes = [0u8; VERIFIER_BYTES];
rand::rng().fill_bytes(&mut bytes);
let encoded = URL_SAFE_NO_PAD.encode(bytes);
debug_assert!(
encoded.len() >= VERIFIER_MIN_LEN && encoded.len() <= VERIFIER_MAX_LEN,
"verifier length {} out of RFC 7636 range [{}, {}]",
encoded.len(),
VERIFIER_MIN_LEN,
VERIFIER_MAX_LEN
);
Self(encoded)
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn from_string(s: impl Into<String>) -> Result<Self, AuthError> {
let s = s.into();
if s.len() < VERIFIER_MIN_LEN || s.len() > VERIFIER_MAX_LEN {
return Err(AuthError::ValidationError(format!(
"code_verifier length {} not in [{}, {}]",
s.len(),
VERIFIER_MIN_LEN,
VERIFIER_MAX_LEN
)));
}
if !s.chars().all(|c| c.is_ascii_alphanumeric() || matches!(c, '-' | '.' | '_' | '~')) {
return Err(AuthError::ValidationError(
"code_verifier contains disallowed characters (RFC 7636 §4.1)".to_string(),
));
}
Ok(Self(s))
}
pub fn to_challenge(&self) -> CodeChallenge {
let hash = Sha256::digest(self.0.as_bytes());
let encoded = URL_SAFE_NO_PAD.encode(hash);
CodeChallenge(encoded)
}
}
impl Default for CodeVerifier {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CodeChallenge(String);
impl CodeChallenge {
pub fn as_str(&self) -> &str {
&self.0
}
pub fn verify(&self, verifier: &CodeVerifier) -> Result<(), AuthError> {
let derived = verifier.to_challenge();
if derived.0.len() != self.0.len() {
return Err(AuthError::ValidationError("code_challenge mismatch".to_string()));
}
let mismatch = derived
.0
.as_bytes()
.iter()
.zip(self.0.as_bytes().iter())
.fold(0u8, |acc, (a, b)| acc | (a ^ b));
if mismatch != 0 {
return Err(AuthError::ValidationError("code_challenge mismatch".to_string()));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OAuthState(String);
impl OAuthState {
pub fn new() -> Self {
let mut bytes = [0u8; STATE_BYTES];
rand::rng().fill_bytes(&mut bytes);
Self(URL_SAFE_NO_PAD.encode(bytes))
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn verify(&self, received: &str) -> Result<(), AuthError> {
if received.len() != self.0.len() {
return Err(AuthError::ValidationError(
"state parameter mismatch (CSRF check failed)".to_string(),
));
}
let mismatch = received
.as_bytes()
.iter()
.zip(self.0.as_bytes().iter())
.fold(0u8, |acc, (a, b)| acc | (a ^ b));
if mismatch != 0 {
return Err(AuthError::ValidationError(
"state parameter mismatch (CSRF check failed)".to_string(),
));
}
Ok(())
}
}
impl Default for OAuthState {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_verifier_length_is_rfc_compliant() {
let v = CodeVerifier::new();
let len = v.as_str().len();
assert!(
(VERIFIER_MIN_LEN..=VERIFIER_MAX_LEN).contains(&len),
"verifier length {len} not in [{VERIFIER_MIN_LEN}, {VERIFIER_MAX_LEN}]"
);
}
#[test]
fn test_verifier_charset_is_unreserved_ascii() {
let v = CodeVerifier::new();
for c in v.as_str().chars() {
assert!(
c.is_ascii_alphanumeric() || matches!(c, '-' | '.' | '_' | '~'),
"verifier contains disallowed char: {c:?}"
);
}
}
#[test]
fn test_verifier_uniqueness() {
let v1 = CodeVerifier::new();
let v2 = CodeVerifier::new();
assert_ne!(v1.as_str(), v2.as_str());
}
#[test]
fn test_verifier_from_string_valid() {
let s = "a".repeat(VERIFIER_MIN_LEN);
assert!(CodeVerifier::from_string(s).is_ok());
}
#[test]
fn test_verifier_from_string_too_short_rejected() {
let s = "a".repeat(VERIFIER_MIN_LEN - 1);
assert!(CodeVerifier::from_string(s).is_err());
}
#[test]
fn test_verifier_from_string_too_long_rejected() {
let s = "a".repeat(VERIFIER_MAX_LEN + 1);
assert!(CodeVerifier::from_string(s).is_err());
}
#[test]
fn test_verifier_from_string_bad_chars_rejected() {
let s = format!("{}+", "a".repeat(VERIFIER_MIN_LEN));
assert!(CodeVerifier::from_string(s).is_err());
}
#[test]
fn test_s256_challenge_derivation_known_vector() {
let verifier = CodeVerifier::from_string("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk")
.expect("valid RFC example verifier");
let challenge = verifier.to_challenge();
assert_eq!(
challenge.as_str(),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM",
"S256 challenge must match RFC 7636 Appendix B"
);
}
#[test]
fn test_challenge_verify_correct_verifier_succeeds() {
let verifier = CodeVerifier::new();
let challenge = verifier.to_challenge();
assert!(
challenge.verify(&verifier).is_ok(),
"valid verifier must pass challenge verification"
);
}
#[test]
fn test_challenge_verify_wrong_verifier_fails() {
let verifier1 = CodeVerifier::new();
let verifier2 = CodeVerifier::new();
let challenge = verifier1.to_challenge();
assert!(
challenge.verify(&verifier2).is_err(),
"wrong verifier must fail challenge verification"
);
}
#[test]
fn test_state_uniqueness() {
let s1 = OAuthState::new();
let s2 = OAuthState::new();
assert_ne!(s1.as_str(), s2.as_str());
}
#[test]
fn test_state_verify_matching_state_passes() {
let state = OAuthState::new();
assert!(state.verify(state.as_str()).is_ok());
}
#[test]
fn test_state_verify_wrong_state_rejected() {
let state = OAuthState::new();
let other = OAuthState::new();
assert!(state.verify(other.as_str()).is_err(), "mismatched state must be rejected (CSRF)");
}
#[test]
fn test_state_verify_empty_rejected() {
let state = OAuthState::new();
assert!(state.verify("").is_err());
}
}