use std::fmt;
use crate::oauth::error::OauthError;
const STATE_BYTES: usize = 32;
pub struct PkceVerifier(oauth2::PkceCodeVerifier);
impl PkceVerifier {
#[must_use]
pub fn from_secret(secret: String) -> Self {
Self(oauth2::PkceCodeVerifier::new(secret))
}
#[must_use]
pub fn secret(&self) -> &str {
self.0.secret()
}
pub(crate) fn into_inner(self) -> oauth2::PkceCodeVerifier {
self.0
}
}
impl fmt::Debug for PkceVerifier {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("PkceVerifier([redacted])")
}
}
pub struct OauthState(String);
impl OauthState {
pub fn generate() -> Result<Self, OauthError> {
let mut bytes = [0u8; STATE_BYTES];
getrandom::fill(&mut bytes).map_err(|_| OauthError::Entropy)?;
Ok(Self(hex_encode(&bytes)))
}
#[must_use]
pub fn from_stored(state: String) -> Self {
Self(state)
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
#[must_use]
pub fn verify(&self, returned: &str) -> bool {
constant_time_eq(self.0.as_bytes(), returned.as_bytes())
}
}
impl fmt::Debug for OauthState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("OauthState([redacted])")
}
}
#[must_use]
pub fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff = 0u8;
for (x, y) in a.iter().zip(b.iter()) {
diff |= x ^ y;
}
std::hint::black_box(diff) == 0
}
fn hex_encode(bytes: &[u8]) -> String {
const DIGITS: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push(DIGITS[usize::from(byte >> 4)] as char);
out.push(DIGITS[usize::from(byte & 0x0f)] as char);
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_generated_state_is_hex_and_full_length() {
let state = OauthState::generate().expect("entropy");
assert_eq!(state.as_str().len(), STATE_BYTES * 2);
assert!(state.as_str().bytes().all(|b| b.is_ascii_hexdigit()));
}
#[test]
fn two_generated_states_differ() {
let a = OauthState::generate().expect("entropy");
let b = OauthState::generate().expect("entropy");
assert_ne!(a.as_str(), b.as_str());
}
#[test]
fn a_state_verifies_against_itself_and_nothing_else() {
let state = OauthState::generate().expect("entropy");
let echoed = state.as_str().to_string();
assert!(state.verify(&echoed));
assert!(!state.verify(""));
assert!(!state.verify(&format!("{echoed}x")));
let mut tampered = echoed.clone();
tampered.replace_range(0..1, "z");
assert!(!state.verify(&tampered));
}
#[test]
fn constant_time_eq_agrees_with_equality_wherever_the_difference_sits() {
let base = vec![7u8; 64];
assert!(constant_time_eq(&base, &base));
for index in [0usize, 1, 31, 63] {
let mut other = base.clone();
other[index] ^= 0x01;
assert!(!constant_time_eq(&base, &other), "missed byte {index}");
}
assert!(!constant_time_eq(&base, &base[..63]));
}
}