use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use base64::Engine;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use super::client::OidcClient;
type HmacSha256 = Hmac<Sha256>;
const DEFAULT_TTL: Duration = Duration::from_secs(600);
#[derive(Clone)]
pub struct PkceCookieManager {
secret: Arc<Vec<u8>>,
cookie_name: Arc<String>,
ttl: Duration,
}
pub struct PkceSession {
pub state: String,
pub verifier: String,
pub cookie_value: String,
}
impl PkceCookieManager {
pub fn new(secret: &[u8], cookie_name: &str, ttl: Duration) -> Self {
assert!(!secret.is_empty(), "PKCE cookie secret must not be empty");
Self {
secret: Arc::new(secret.to_vec()),
cookie_name: Arc::new(cookie_name.to_string()),
ttl,
}
}
pub fn with_default_ttl(secret: &[u8], cookie_name: &str) -> Self {
Self::new(secret, cookie_name, DEFAULT_TTL)
}
pub fn cookie_name(&self) -> &str {
&self.cookie_name
}
pub fn ttl(&self) -> Duration {
self.ttl
}
pub fn create(&self) -> PkceSession {
let state = OidcClient::generate_state();
let verifier = OidcClient::generate_code_verifier();
let expiry = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs()
+ self.ttl.as_secs();
let cookie_value = self.sign(&state, &verifier, expiry);
PkceSession {
state,
verifier,
cookie_value,
}
}
pub fn verify(&self, cookie_value: &str, expected_state: &str) -> Option<String> {
let (state, verifier, expiry) = self.verify_internal(cookie_value)?;
if state != expected_state {
tracing::warn!("PKCE cookie state mismatch: expected {}, got {}", expected_state, state);
return None;
}
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()?
.as_secs();
if now >= expiry {
tracing::debug!("PKCE cookie expired (expiry={}, now={})", expiry, now);
return None;
}
Some(verifier)
}
fn sign(&self, state: &str, verifier: &str, expiry: u64) -> String {
let b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD;
let state_b64 = b64.encode(state);
let verifier_b64 = b64.encode(verifier);
let expiry_b64 = b64.encode(expiry.to_string());
let message = format!("{}.{}.{}", state_b64, verifier_b64, expiry_b64);
let mut mac = HmacSha256::new_from_slice(&self.secret)
.expect("HMAC accepts any key length");
mac.update(message.as_bytes());
let mac_bytes = mac.finalize().into_bytes();
let mac_b64 = b64.encode(mac_bytes);
format!("{}.{}", message, mac_b64)
}
fn verify_internal(&self, cookie_value: &str) -> Option<(String, String, u64)> {
let b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD;
let parts: Vec<&str> = cookie_value.split('.').collect();
if parts.len() != 4 {
tracing::debug!("PKCE cookie has {} segments, expected 4", parts.len());
return None;
}
let [state_b64, verifier_b64, expiry_b64, mac_b64] = [parts[0], parts[1], parts[2], parts[3]];
let message = format!("{}.{}.{}", state_b64, verifier_b64, expiry_b64);
let mut mac = HmacSha256::new_from_slice(&self.secret).ok()?;
mac.update(message.as_bytes());
let expected_mac = mac.finalize().into_bytes();
let provided_mac = b64.decode(mac_b64).ok()?;
if expected_mac.as_slice() != provided_mac.as_slice() {
tracing::warn!("PKCE cookie HMAC verification failed");
return None;
}
let state = String::from_utf8(b64.decode(state_b64).ok()?).ok()?;
let verifier = String::from_utf8(b64.decode(verifier_b64).ok()?).ok()?;
let expiry_str = String::from_utf8(b64.decode(expiry_b64).ok()?).ok()?;
let expiry: u64 = expiry_str.parse().ok()?;
Some((state, verifier, expiry))
}
}
#[cfg(test)]
mod tests {
use super::*;
const TEST_SECRET: &[u8] = b"test-secret-key-at-least-32-bytes!!";
fn make_manager() -> PkceCookieManager {
PkceCookieManager::with_default_ttl(TEST_SECRET, "test_pkce")
}
#[test]
fn test_create_and_verify() {
let mgr = make_manager();
let session = mgr.create();
let verifier = mgr.verify(&session.cookie_value, &session.state);
assert!(verifier.is_some());
assert_eq!(verifier.unwrap(), session.verifier);
}
#[test]
fn test_state_mismatch() {
let mgr = make_manager();
let session = mgr.create();
let verifier = mgr.verify(&session.cookie_value, "wrong-state");
assert!(verifier.is_none());
}
#[test]
fn test_tampered_cookie() {
let mgr = make_manager();
let session = mgr.create();
let tampered = format!("{}.tampered", &session.cookie_value);
let verifier = mgr.verify(&tampered, &session.state);
assert!(verifier.is_none());
}
#[test]
fn test_tampered_state() {
let mgr = make_manager();
let session = mgr.create();
let evil_mgr = PkceCookieManager::with_default_ttl(b"evil-different-secret-key!!!!", "test_pkce");
let evil_session = evil_mgr.create();
let verifier = mgr.verify(&evil_session.cookie_value, &evil_session.state);
assert!(verifier.is_none());
}
#[test]
fn test_expired() {
let mgr = PkceCookieManager::new(TEST_SECRET, "test_pkce", Duration::from_secs(0));
let session = mgr.create();
std::thread::sleep(Duration::from_millis(10));
let verifier = mgr.verify(&session.cookie_value, &session.state);
assert!(verifier.is_none());
}
#[test]
fn test_malformed_cookie() {
let mgr = make_manager();
assert!(mgr.verify("abc", "state").is_none());
assert!(mgr.verify("a.b.c.d.e", "state").is_none());
assert!(mgr.verify("", "state").is_none());
}
#[test]
fn test_different_managers_same_secret() {
let mgr1 = PkceCookieManager::with_default_ttl(TEST_SECRET, "test_pkce");
let mgr2 = PkceCookieManager::with_default_ttl(TEST_SECRET, "test_pkce");
let session = mgr1.create();
let verifier = mgr2.verify(&session.cookie_value, &session.state);
assert!(verifier.is_some());
}
#[test]
fn test_different_managers_different_secret() {
let mgr1 = PkceCookieManager::with_default_ttl(b"secret-one-aaaaaaaaaaaaaaaaaaaa", "test_pkce");
let mgr2 = PkceCookieManager::with_default_ttl(b"secret-two-bbbbbbbbbbbbbbbbbbbb", "test_pkce");
let session = mgr1.create();
let verifier = mgr2.verify(&session.cookie_value, &session.state);
assert!(verifier.is_none());
}
#[test]
fn test_cookie_name_accessor() {
let mgr = PkceCookieManager::with_default_ttl(TEST_SECRET, "my_cookie_name");
assert_eq!(mgr.cookie_name(), "my_cookie_name");
}
#[test]
fn test_ttl_accessor() {
let mgr = PkceCookieManager::new(TEST_SECRET, "test_pkce", Duration::from_secs(120));
assert_eq!(mgr.ttl(), Duration::from_secs(120));
}
#[test]
fn test_unique_sessions() {
let mgr = make_manager();
let s1 = mgr.create();
let s2 = mgr.create();
assert_ne!(s1.state, s2.state);
assert_ne!(s1.verifier, s2.verifier);
assert_ne!(s1.cookie_value, s2.cookie_value);
}
}