1use anyhow::{Context, Result};
7use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
8use ring::rand::{SecureRandom, SystemRandom};
9use sha2::{Digest, Sha256};
10use std::fmt;
11
12const CODE_VERIFIER_LENGTH: usize = 64;
14
15const CODE_VERIFIER_CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~";
17
18#[derive(Clone)]
24pub struct PkceChallenge {
25 pub(crate) code_verifier: String,
27 pub(crate) code_challenge: String,
29 pub(crate) code_challenge_method: String,
31}
32
33impl fmt::Debug for PkceChallenge {
34 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
35 f.debug_struct("PkceChallenge")
36 .field("code_verifier", &"<redacted>")
37 .field("code_challenge", &self.code_challenge)
38 .field("code_challenge_method", &self.code_challenge_method)
39 .finish()
40 }
41}
42
43impl PkceChallenge {
44 fn from_verifier(code_verifier: String) -> Result<Self> {
46 let code_challenge = compute_s256_challenge(&code_verifier)?;
47 Ok(Self {
48 code_verifier,
49 code_challenge,
50 code_challenge_method: "S256".to_string(),
51 })
52 }
53}
54
55pub fn generate_pkce_challenge() -> Result<PkceChallenge> {
69 let code_verifier = generate_code_verifier()?;
70 PkceChallenge::from_verifier(code_verifier)
71}
72
73fn generate_code_verifier() -> Result<String> {
78 let rng = SystemRandom::new();
79 let charset_len = u8::try_from(CODE_VERIFIER_CHARSET.len()).context("PKCE verifier alphabet is too large")?;
80 let max_valid =
81 u8::try_from(256u16 - 256u16 % u16::from(charset_len)).context("PKCE verifier rejection range is invalid")?;
82 let mut verifier = String::with_capacity(CODE_VERIFIER_LENGTH);
83 let mut buf = [0u8; 1];
84
85 while verifier.len() < CODE_VERIFIER_LENGTH {
86 rng.fill(&mut buf)
87 .map_err(|_| anyhow::anyhow!("failed to read from OS random source"))?;
88 if buf[0] < max_valid {
90 let idx = usize::from(buf[0] % charset_len);
91 if let Some(&character) = CODE_VERIFIER_CHARSET.get(idx) {
92 verifier.push(char::from(character));
93 }
94 }
95 }
96
97 Ok(verifier)
98}
99
100fn compute_s256_challenge(code_verifier: &str) -> Result<String> {
104 let mut hasher = Sha256::new();
105 hasher.update(code_verifier.as_bytes());
106 let hash = hasher.finalize();
107
108 Ok(URL_SAFE_NO_PAD.encode(hash))
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 #[test]
116 fn test_generate_pkce_challenge() {
117 let challenge = generate_pkce_challenge().unwrap();
118
119 assert_eq!(challenge.code_verifier.len(), CODE_VERIFIER_LENGTH);
121
122 for c in challenge.code_verifier.chars() {
124 assert!(CODE_VERIFIER_CHARSET.contains(&(c as u8)), "Invalid character in verifier: {c}");
125 }
126
127 assert_eq!(challenge.code_challenge_method, "S256");
129
130 assert_eq!(challenge.code_challenge.len(), 43);
132 }
133
134 #[test]
135 fn test_deterministic_challenge() {
136 let verifier = "test_verifier_string_for_deterministic_test";
138 let challenge1 = PkceChallenge::from_verifier(verifier.to_string()).unwrap();
139 let challenge2 = PkceChallenge::from_verifier(verifier.to_string()).unwrap();
140
141 assert_eq!(challenge1.code_challenge, challenge2.code_challenge);
142 }
143
144 #[test]
145 fn test_unique_verifiers() {
146 let c1 = generate_pkce_challenge().unwrap();
148 let c2 = generate_pkce_challenge().unwrap();
149
150 assert_ne!(c1.code_verifier, c2.code_verifier);
151 }
152}