use std::{
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use dashmap::DashMap;
use rand::RngCore as _;
use sha2::{Digest, Sha256};
use thiserror::Error;
use crate::state_encryption::StateEncryptionService;
const MAX_PKCE_ENTRIES: usize = 10_000;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum PkceError {
#[error(
"state not found — the authorization flow may have already been completed or the state is invalid"
)]
StateNotFound,
#[error("state expired — please restart the authorization flow")]
StateExpired,
#[error("PKCE state store is full — too many concurrent authorization flows")]
StoreFull,
}
#[derive(Debug)]
pub struct ConsumedPkceState {
pub verifier: String,
pub redirect_uri: String,
}
struct PkceEntry {
verifier: String,
redirect_uri: String,
created_at: tokio::time::Instant,
ttl: Duration,
}
pub struct InMemoryPkceStateStore {
state_ttl_secs: u64,
entries: DashMap<String, PkceEntry>,
encryptor: Option<Arc<StateEncryptionService>>,
max_entries: usize,
size: AtomicUsize,
}
impl InMemoryPkceStateStore {
fn new(state_ttl_secs: u64, encryptor: Option<Arc<StateEncryptionService>>) -> Self {
Self {
state_ttl_secs,
entries: DashMap::new(),
encryptor,
max_entries: MAX_PKCE_ENTRIES,
size: AtomicUsize::new(0),
}
}
#[cfg(test)]
fn with_max_entries(
state_ttl_secs: u64,
encryptor: Option<Arc<StateEncryptionService>>,
max_entries: usize,
) -> Self {
Self {
state_ttl_secs,
entries: DashMap::new(),
encryptor,
max_entries,
size: AtomicUsize::new(0),
}
}
fn create_state_sync(&self, redirect_uri: &str) -> Result<(String, String), anyhow::Error> {
if self.size.fetch_add(1, Ordering::AcqRel) >= self.max_entries {
self.size.fetch_sub(1, Ordering::AcqRel);
return Err(PkceError::StoreFull.into());
}
let mut verifier_bytes = [0u8; 32];
rand::rng().fill_bytes(&mut verifier_bytes);
let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
let mut key_bytes = [0u8; 32];
rand::rng().fill_bytes(&mut key_bytes);
let internal_key = URL_SAFE_NO_PAD.encode(key_bytes);
self.entries.insert(
internal_key.clone(),
PkceEntry {
verifier: verifier.clone(),
redirect_uri: redirect_uri.to_owned(),
created_at: tokio::time::Instant::now(),
ttl: Duration::from_secs(self.state_ttl_secs),
},
);
let outbound_token = match &self.encryptor {
Some(enc) => enc.encrypt(internal_key.as_bytes())?,
None => internal_key,
};
Ok((outbound_token, verifier))
}
pub fn purge_expired(&self) {
let mut removed = 0usize;
self.entries.retain(|_, e| {
let keep = e.created_at.elapsed() <= e.ttl;
removed += usize::from(!keep);
keep
});
self.size.fetch_sub(removed, Ordering::AcqRel);
}
fn consume_state_sync(&self, outbound_token: &str) -> Result<ConsumedPkceState, PkceError> {
let internal_key = match &self.encryptor {
Some(enc) => {
let bytes = enc.decrypt(outbound_token).map_err(|_| PkceError::StateNotFound)?;
String::from_utf8(bytes).map_err(|_| PkceError::StateNotFound)?
},
None => outbound_token.to_owned(),
};
let (_, entry) = self.entries.remove(&internal_key).ok_or(PkceError::StateNotFound)?;
self.size.fetch_sub(1, Ordering::AcqRel);
if entry.created_at.elapsed() > entry.ttl {
return Err(PkceError::StateExpired);
}
Ok(ConsumedPkceState {
verifier: entry.verifier,
redirect_uri: entry.redirect_uri,
})
}
fn cleanup_expired_sync(&self) {
self.purge_expired();
}
fn len_sync(&self) -> usize {
self.entries.len()
}
}
#[cfg(feature = "redis-pkce")]
pub static REDIS_PKCE_ERRORS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
#[cfg(feature = "redis-pkce")]
fn note_redis_pkce_error() {
REDIS_PKCE_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
#[cfg(feature = "redis-pkce")]
pub fn redis_pkce_error_count_total() -> u64 {
REDIS_PKCE_ERRORS.load(std::sync::atomic::Ordering::Relaxed)
}
#[cfg(feature = "redis-pkce")]
pub struct RedisPkceStateStore {
pool: redis::aio::ConnectionManager,
state_ttl_secs: u64,
encryptor: Option<Arc<StateEncryptionService>>,
}
#[cfg(feature = "redis-pkce")]
impl RedisPkceStateStore {
pub async fn new(
url: &str,
state_ttl_secs: u64,
encryptor: Option<Arc<StateEncryptionService>>,
) -> Result<Self, redis::RedisError> {
let client = redis::Client::open(url)?;
let pool = redis::aio::ConnectionManager::new(client).await?;
Ok(Self {
pool,
state_ttl_secs,
encryptor,
})
}
async fn create_state_impl(
&self,
redirect_uri: &str,
) -> Result<(String, String), anyhow::Error> {
let mut verifier_bytes = [0u8; 32];
rand::rng().fill_bytes(&mut verifier_bytes);
let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
let mut key_bytes = [0u8; 32];
rand::rng().fill_bytes(&mut key_bytes);
let internal_key = URL_SAFE_NO_PAD.encode(key_bytes);
let redis_key = format!("fraiseql:pkce:{internal_key}");
let value = serde_json::json!({
"verifier": verifier,
"redirect_uri": redirect_uri,
})
.to_string();
let mut conn = self.pool.clone();
redis::cmd("SET")
.arg(&redis_key)
.arg(&value)
.arg("EX")
.arg(self.state_ttl_secs)
.query_async::<()>(&mut conn)
.await
.inspect_err(|_| note_redis_pkce_error())?;
let outbound_token = match &self.encryptor {
Some(enc) => enc.encrypt(internal_key.as_bytes())?,
None => internal_key,
};
Ok((outbound_token, verifier))
}
async fn consume_state_impl(
&self,
outbound_token: &str,
) -> Result<ConsumedPkceState, PkceError> {
#[derive(serde::Deserialize)]
struct StoredEntry {
verifier: String,
redirect_uri: String,
}
let internal_key = match &self.encryptor {
Some(enc) => {
let bytes = enc.decrypt(outbound_token).map_err(|_| PkceError::StateNotFound)?;
String::from_utf8(bytes).map_err(|_| PkceError::StateNotFound)?
},
None => outbound_token.to_owned(),
};
let redis_key = format!("fraiseql:pkce:{internal_key}");
let mut conn = self.pool.clone();
let raw: Option<String> = redis::cmd("GETDEL")
.arg(&redis_key)
.query_async(&mut conn)
.await
.map_err(|error| {
note_redis_pkce_error();
tracing::error!(%error, "Redis PKCE GETDEL failed — treating as state-not-found");
PkceError::StateNotFound
})?;
let json = raw.ok_or(PkceError::StateNotFound)?;
let entry: StoredEntry =
serde_json::from_str(&json).map_err(|_| PkceError::StateNotFound)?;
Ok(ConsumedPkceState {
verifier: entry.verifier,
redirect_uri: entry.redirect_uri,
})
}
}
#[non_exhaustive]
pub enum PkceStateStore {
InMemory(InMemoryPkceStateStore),
#[cfg(feature = "redis-pkce")]
Redis(RedisPkceStateStore),
}
impl PkceStateStore {
#[must_use]
pub fn new(state_ttl_secs: u64, encryptor: Option<Arc<StateEncryptionService>>) -> Self {
Self::InMemory(InMemoryPkceStateStore::new(state_ttl_secs, encryptor))
}
#[cfg(test)]
pub(crate) fn new_capped(
state_ttl_secs: u64,
encryptor: Option<Arc<StateEncryptionService>>,
max_entries: usize,
) -> Self {
Self::InMemory(InMemoryPkceStateStore::with_max_entries(
state_ttl_secs,
encryptor,
max_entries,
))
}
#[cfg(feature = "redis-pkce")]
pub async fn new_redis(
url: &str,
state_ttl_secs: u64,
encryptor: Option<Arc<StateEncryptionService>>,
) -> Result<Self, redis::RedisError> {
let inner = RedisPkceStateStore::new(url, state_ttl_secs, encryptor).await?;
Ok(Self::Redis(inner))
}
#[must_use]
pub const fn is_in_memory(&self) -> bool {
matches!(self, Self::InMemory(_))
}
pub async fn create_state(
&self,
redirect_uri: &str,
) -> Result<(String, String), anyhow::Error> {
match self {
Self::InMemory(s) => s.create_state_sync(redirect_uri),
#[cfg(feature = "redis-pkce")]
Self::Redis(s) => s.create_state_impl(redirect_uri).await,
}
}
pub async fn consume_state(
&self,
outbound_token: &str,
) -> Result<ConsumedPkceState, PkceError> {
match self {
Self::InMemory(s) => s.consume_state_sync(outbound_token),
#[cfg(feature = "redis-pkce")]
Self::Redis(s) => s.consume_state_impl(outbound_token).await,
}
}
#[must_use]
pub fn s256_challenge(verifier: &str) -> String {
URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()))
}
#[allow(unknown_lints, clippy::unused_async_trait_impl)]
pub async fn cleanup_expired(&self) {
match self {
Self::InMemory(s) => s.cleanup_expired_sync(),
#[cfg(feature = "redis-pkce")]
Self::Redis(_) => {}, }
}
pub fn purge_expired(&self) {
match self {
Self::InMemory(s) => s.purge_expired(),
#[cfg(feature = "redis-pkce")]
Self::Redis(_) => {}, }
}
#[must_use]
pub fn len(&self) -> usize {
match self {
Self::InMemory(s) => s.len_sync(),
#[cfg(feature = "redis-pkce")]
Self::Redis(_) => 0,
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}