1use anyhow::Result;
7use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
8use ring::rand::{SecureRandom, SystemRandom};
9use sha2::{Digest, Sha256};
10
11const CODE_VERIFIER_LENGTH: usize = 64;
13
14const CODE_VERIFIER_CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~";
16
17#[derive(Debug, Clone)]
19pub struct PkceChallenge {
20 pub code_verifier: String,
22 pub code_challenge: String,
24 pub code_challenge_method: String,
26}
27
28impl PkceChallenge {
29 pub fn from_verifier(code_verifier: String) -> Result<Self> {
31 let code_challenge = compute_s256_challenge(&code_verifier)?;
32 Ok(Self {
33 code_verifier,
34 code_challenge,
35 code_challenge_method: "S256".to_string(),
36 })
37 }
38}
39
40pub fn generate_pkce_challenge() -> Result<PkceChallenge> {
54 let code_verifier = generate_code_verifier()?;
55 PkceChallenge::from_verifier(code_verifier)
56}
57
58fn generate_code_verifier() -> Result<String> {
63 let rng = SystemRandom::new();
64 let charset_len = CODE_VERIFIER_CHARSET.len() as u8;
65 let max_valid = (256u16 - 256u16 % charset_len as u16) as u8;
66 let mut verifier = String::with_capacity(CODE_VERIFIER_LENGTH);
67 let mut buf = [0u8; 1];
68
69 while verifier.len() < CODE_VERIFIER_LENGTH {
70 rng.fill(&mut buf)
71 .map_err(|_| anyhow::anyhow!("failed to read from OS random source"))?;
72 if buf[0] < max_valid {
74 let idx = (buf[0] % charset_len) as usize;
75 verifier.push(CODE_VERIFIER_CHARSET[idx] as char);
76 }
77 }
78
79 Ok(verifier)
80}
81
82fn compute_s256_challenge(code_verifier: &str) -> Result<String> {
86 let mut hasher = Sha256::new();
87 hasher.update(code_verifier.as_bytes());
88 let hash = hasher.finalize();
89
90 Ok(URL_SAFE_NO_PAD.encode(hash))
91}
92
93#[cfg(test)]
94mod tests {
95 use super::*;
96
97 #[test]
98 fn test_generate_pkce_challenge() {
99 let challenge = generate_pkce_challenge().unwrap();
100
101 assert_eq!(challenge.code_verifier.len(), CODE_VERIFIER_LENGTH);
103
104 for c in challenge.code_verifier.chars() {
106 assert!(CODE_VERIFIER_CHARSET.contains(&(c as u8)), "Invalid character in verifier: {c}");
107 }
108
109 assert_eq!(challenge.code_challenge_method, "S256");
111
112 assert_eq!(challenge.code_challenge.len(), 43);
114 }
115
116 #[test]
117 fn test_deterministic_challenge() {
118 let verifier = "test_verifier_string_for_deterministic_test";
120 let challenge1 = PkceChallenge::from_verifier(verifier.to_string()).unwrap();
121 let challenge2 = PkceChallenge::from_verifier(verifier.to_string()).unwrap();
122
123 assert_eq!(challenge1.code_challenge, challenge2.code_challenge);
124 }
125
126 #[test]
127 fn test_unique_verifiers() {
128 let c1 = generate_pkce_challenge().unwrap();
130 let c2 = generate_pkce_challenge().unwrap();
131
132 assert_ne!(c1.code_verifier, c2.code_verifier);
133 }
134}