use std::fmt;
use crate::crypto::constant_time::constant_time_eq;
use crate::crypto::zeroize::Zeroizing;
use crate::crypto::{RandomError, fill_random};
use crate::encoding::base64url_encode;
use crate::util::log::debug;
#[doc(alias = "csrf_state")]
pub struct OAuthState {
value: Zeroizing<String>,
}
impl OAuthState {
pub fn generate() -> Result<Self, RandomError> {
let mut buf = [0u8; 32];
fill_random(&mut buf)?;
let value = base64url_encode(&buf);
crate::crypto::zeroize::zeroize(&mut buf);
debug!("oauth: state parameter generated");
Ok(Self {
value: Zeroizing::new(value),
})
}
#[must_use]
#[inline]
pub fn value(&self) -> &str {
&self.value
}
#[must_use]
#[inline]
pub fn verify(&self, received: &str) -> bool {
constant_time_eq(self.value.as_bytes(), received.as_bytes())
}
}
impl fmt::Debug for OAuthState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OAuthState")
.field("value", &"[REDACTED]")
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn generate_produces_nonempty_value() {
let state = OAuthState::generate().unwrap();
assert!(!state.value().is_empty(), "state should not be empty");
}
#[test]
fn generate_produces_base64url_value() {
let state = OAuthState::generate().unwrap();
assert!(
state
.value()
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'),
"state should contain only base64url characters: {}",
state.value(),
);
}
#[test]
fn generate_produces_43_char_value() {
let state = OAuthState::generate().unwrap();
assert_eq!(
state.value().len(),
43,
"32 bytes base64url-encoded = 43 characters",
);
}
#[test]
fn verify_correct_state() {
let state = OAuthState::generate().unwrap();
let value = state.value().to_owned();
assert!(
state.verify(&value),
"verify should succeed for matching state"
);
}
#[test]
fn verify_wrong_state_fails() {
let state = OAuthState::generate().unwrap();
assert!(
!state.verify("wrong-state-value"),
"verify should fail for non-matching state",
);
}
#[test]
fn verify_empty_state_fails() {
let state = OAuthState::generate().unwrap();
assert!(
!state.verify(""),
"verify should fail for empty received state",
);
}
#[test]
fn verify_similar_state_fails() {
let state = OAuthState::generate().unwrap();
let mut tampered = state.value().to_owned();
tampered.push('X');
assert!(
!state.verify(&tampered),
"verify should fail for tampered state",
);
}
#[test]
fn debug_redacts_state_value() {
let state = OAuthState::generate().unwrap();
let debug = format!("{state:?}");
assert!(
debug.contains("[REDACTED]"),
"Debug should redact the state value: {debug}",
);
assert!(
!debug.contains(state.value()),
"Debug must not contain the raw state value: {debug}",
);
}
#[test]
fn two_states_differ() {
let a = OAuthState::generate().unwrap();
let b = OAuthState::generate().unwrap();
assert_ne!(a.value(), b.value(), "two generated states should differ");
}
#[test]
fn different_states_do_not_verify() {
let a = OAuthState::generate().unwrap();
let b = OAuthState::generate().unwrap();
assert!(
!a.verify(b.value()),
"state A should not verify state B's value",
);
}
}