use std::collections::HashMap;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use std::time::{Duration, Instant};
use ahash::AHasher;
use parking_lot::RwLock;
use tokio_util::sync::CancellationToken;
use crate::client::StoredStateEntry;
use crate::error::StoreError;
pub const SHARD_ADMISSION_CAPACITY: usize = 1024;
pub const NUM_SHARDS: usize = 64;
pub const DEFAULT_STATE_TTL: Duration = Duration::from_secs(300);
pub trait OAuthStore: Send + Sync + 'static {
fn insert_state(
&self,
state: String,
entry: StoredStateEntry,
ttl: Duration,
) -> impl std::future::Future<Output = Result<(), StoreError>> + Send;
fn take_state(
&self,
state: &str,
) -> impl std::future::Future<Output = Result<Option<StoredStateEntry>, StoreError>> + Send;
fn contains_state(
&self,
state: &str,
) -> impl std::future::Future<Output = Result<bool, StoreError>> + Send;
fn prune_expired(&self) -> impl std::future::Future<Output = Result<usize, StoreError>> + Send;
}
#[derive(Debug, Clone)]
struct StoredStateRecord {
entry: StoredStateEntry,
created_at: Instant,
expires_at: Instant,
}
impl StoredStateRecord {
fn new(entry: StoredStateEntry, ttl: Duration) -> Self {
let now = Instant::now();
let expires_at = now
.checked_add(ttl)
.unwrap_or(now + Duration::from_secs(86400 * 365));
Self {
entry,
created_at: now,
expires_at,
}
}
fn is_expired(&self, now: Instant) -> bool {
if now >= self.expires_at {
true
} else {
let elapsed = now.saturating_duration_since(self.created_at);
let ttl = self.expires_at.saturating_duration_since(self.created_at);
elapsed >= ttl
}
}
}
pub struct OAuthStateStore {
shards: [RwLock<HashMap<String, StoredStateRecord>>; NUM_SHARDS],
default_ttl: Duration,
}
impl std::fmt::Debug for OAuthStateStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuthStateStore")
.field("num_shards", &NUM_SHARDS)
.field("default_ttl", &self.default_ttl)
.finish()
}
}
impl Default for OAuthStateStore {
fn default() -> Self {
Self::new(DEFAULT_STATE_TTL)
}
}
impl OAuthStateStore {
#[must_use]
pub fn new(default_ttl: Duration) -> Self {
let shards = std::array::from_fn(|_| RwLock::new(HashMap::new()));
Self {
shards,
default_ttl,
}
}
#[must_use]
#[inline]
pub fn shard_index(&self, key: &str) -> usize {
let mut hasher = AHasher::default();
key.hash(&mut hasher);
(hasher.finish() as usize) % NUM_SHARDS
}
#[must_use]
pub fn default_ttl(&self) -> Duration {
self.default_ttl
}
pub fn insert_state_sync(
&self,
state: String,
entry: StoredStateEntry,
ttl: Duration,
) -> Result<(), StoreError> {
let idx = self.shard_index(&state);
let record = StoredStateRecord::new(entry, ttl);
let mut shard = self.shards[idx].write();
if shard.len() >= SHARD_ADMISSION_CAPACITY {
let now = Instant::now();
shard.retain(|_, r| !r.is_expired(now));
if shard.len() >= SHARD_ADMISSION_CAPACITY {
return Err(StoreError::CapacityExceeded(SHARD_ADMISSION_CAPACITY));
}
}
shard.insert(state, record);
Ok(())
}
pub fn insert_default_sync(
&self,
state: String,
entry: StoredStateEntry,
) -> Result<(), StoreError> {
self.insert_state_sync(state, entry, self.default_ttl)
}
#[must_use]
pub fn take_state_sync(&self, state: &str) -> Option<StoredStateEntry> {
let idx = self.shard_index(state);
let mut shard = self.shards[idx].write();
if let Some(record) = shard.remove(state) {
let now = Instant::now();
if !record.is_expired(now) {
Some(record.entry)
} else {
None
}
} else {
None
}
}
#[must_use]
pub fn contains_state_sync(&self, state: &str) -> bool {
let idx = self.shard_index(state);
let shard = self.shards[idx].read();
if let Some(record) = shard.get(state) {
let now = Instant::now();
!record.is_expired(now)
} else {
false
}
}
pub fn prune_expired_sync(&self) -> usize {
let now = Instant::now();
let mut total_pruned = 0;
for shard in &self.shards {
let mut guard = shard.write();
let initial_count = guard.len();
guard.retain(|_, record| !record.is_expired(now));
total_pruned += initial_count.saturating_sub(guard.len());
}
total_pruned
}
#[must_use]
pub fn total_entries(&self) -> usize {
let mut sum = 0;
for shard in &self.shards {
sum += shard.read().len();
}
sum
}
#[must_use]
pub fn shard_len(&self, shard_idx: usize) -> usize {
self.shards[shard_idx % NUM_SHARDS].read().len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.total_entries() == 0
}
pub fn clear(&self) {
for shard in &self.shards {
shard.write().clear();
}
}
pub fn spawn_pruning_task(
self: &Arc<Self>,
interval_duration: Duration,
cancellation_token: CancellationToken,
) -> tokio::task::JoinHandle<()> {
assert!(
!interval_duration.is_zero(),
"spawn_pruning_task: interval_duration must be non-zero"
);
assert!(
tokio::runtime::Handle::try_current().is_ok(),
"spawn_pruning_task: must be called within a tokio runtime"
);
let store = Arc::clone(self);
tokio::spawn(async move {
let mut interval = tokio::time::interval(interval_duration);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = cancellation_token.cancelled() => {
tracing::debug!("OAuthStateStore background pruning worker received cancellation signal; exiting.");
break;
}
_ = interval.tick() => {
let pruned = store.prune_expired_sync();
if pruned > 0 {
tracing::trace!("OAuthStateStore background pruner evicted {} expired states", pruned);
}
}
}
}
})
}
}
impl OAuthStore for OAuthStateStore {
async fn insert_state(
&self,
state: String,
entry: StoredStateEntry,
ttl: Duration,
) -> Result<(), StoreError> {
self.insert_state_sync(state, entry, ttl)
}
async fn take_state(&self, state: &str) -> Result<Option<StoredStateEntry>, StoreError> {
Ok(self.take_state_sync(state))
}
async fn contains_state(&self, state: &str) -> Result<bool, StoreError> {
Ok(self.contains_state_sync(state))
}
async fn prune_expired(&self) -> Result<usize, StoreError> {
Ok(self.prune_expired_sync())
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod tests {
use super::*;
use crate::dpop::DPoPKey;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::SystemTime;
fn mock_stored_state(state: &str) -> StoredStateEntry {
StoredStateEntry {
state: state.to_string(),
client_id: "https://app.example.com/client-metadata.json".to_string(),
code_verifier: "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk".to_string(),
dpop_key: DPoPKey::generate(),
issuer: "https://auth.example.com".to_string(),
did: Some("did:plc:alice123".to_string()),
handle: Some("alice.bsky.social".to_string()),
redirect_uri: "https://app.example.com/callback".to_string(),
pds_endpoint: "https://pds.example.com".to_string(),
token_endpoint: "https://auth.example.com/oauth/token".to_string(),
scopes: "atproto".to_string(),
created_at: SystemTime::now(),
expires_in_secs: 300,
}
}
#[test]
fn test_store_initialization_and_num_shards() {
let store = OAuthStateStore::default();
assert_eq!(store.shards.len(), 64);
assert_eq!(store.total_entries(), 0);
assert!(store.is_empty());
}
#[test]
fn test_store_insert_and_single_use_consumption() {
let store = OAuthStateStore::default();
let state = "csrf_token_test_123";
let entry = mock_stored_state(state);
assert!(!store.contains_state_sync(state));
store
.insert_state_sync(state.to_string(), entry.clone(), Duration::from_secs(60))
.unwrap();
assert!(store.contains_state_sync(state));
assert_eq!(store.total_entries(), 1);
let consumed = store.take_state_sync(state);
assert!(consumed.is_some());
let consumed = consumed.unwrap();
assert_eq!(consumed.state, state);
assert_eq!(consumed.client_id, entry.client_id);
let second_take = store.take_state_sync(state);
assert!(second_take.is_none());
assert!(!store.contains_state_sync(state));
assert_eq!(store.total_entries(), 0);
}
#[test]
fn test_shard_distribution_uniformity() {
let store = OAuthStateStore::default();
let mut hit_shards = std::collections::HashSet::new();
for i in 0..1000 {
let key = format!("state_entropy_sample_token_{i}");
hit_shards.insert(store.shard_index(&key));
}
assert!(
hit_shards.len() >= 55,
"Shard distribution too sparse: hit {} shards out of 64",
hit_shards.len()
);
}
#[test]
fn test_concurrent_single_use_50_threads() {
let store = Arc::new(OAuthStateStore::default());
let state = "race_condition_state_token";
let entry = mock_stored_state(state);
store
.insert_state_sync(state.to_string(), entry, Duration::from_secs(60))
.unwrap();
let winner_count = Arc::new(AtomicUsize::new(0));
let barrier = Arc::new(std::sync::Barrier::new(50));
let mut handles = Vec::new();
for _ in 0..50 {
let s = Arc::clone(&store);
let w = Arc::clone(&winner_count);
let b = Arc::clone(&barrier);
let state_str = state.to_string();
handles.push(std::thread::spawn(move || {
b.wait();
if s.take_state_sync(&state_str).is_some() {
w.fetch_add(1, Ordering::SeqCst);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(
winner_count.load(Ordering::SeqCst),
1,
"Exactly one thread must successfully consume state among 50 racers"
);
assert_eq!(store.total_entries(), 0);
}
#[test]
fn test_ttl_expiration_and_pruning() {
let store = OAuthStateStore::default();
let state_active = "active_state";
let state_expired = "expired_state";
store
.insert_state_sync(
state_active.to_string(),
mock_stored_state(state_active),
Duration::from_secs(60),
)
.unwrap();
store
.insert_state_sync(
state_expired.to_string(),
mock_stored_state(state_expired),
Duration::ZERO,
)
.unwrap();
assert_eq!(store.total_entries(), 2);
assert!(store.contains_state_sync(state_active));
assert!(!store.contains_state_sync(state_expired));
assert!(store.take_state_sync(state_expired).is_none());
store
.insert_state_sync(
state_expired.to_string(),
mock_stored_state(state_expired),
Duration::ZERO,
)
.unwrap();
let pruned = store.prune_expired_sync();
assert_eq!(pruned, 1);
assert_eq!(store.total_entries(), 1);
assert!(store.contains_state_sync(state_active));
}
#[tokio::test]
async fn test_oauth_store_trait_async_operations() {
let store = OAuthStateStore::default();
let state = "async_trait_test_state";
let entry = mock_stored_state(state);
store
.insert_state(state.to_string(), entry.clone(), Duration::from_secs(300))
.await
.unwrap();
assert!(store.contains_state(state).await.unwrap());
let taken = store.take_state(state).await.unwrap();
assert!(taken.is_some());
assert_eq!(taken.unwrap().state, state);
assert!(!store.contains_state(state).await.unwrap());
assert!(store.take_state(state).await.unwrap().is_none());
}
#[tokio::test]
async fn test_background_pruning_task_cancellation() {
let store = Arc::new(OAuthStateStore::default());
let cancel_token = CancellationToken::new();
store
.insert_state_sync(
"exp1".to_string(),
mock_stored_state("exp1"),
Duration::ZERO,
)
.unwrap();
store
.insert_state_sync(
"exp2".to_string(),
mock_stored_state("exp2"),
Duration::ZERO,
)
.unwrap();
assert_eq!(store.total_entries(), 2);
let handle = store.spawn_pruning_task(Duration::from_millis(20), cancel_token.clone());
tokio::time::sleep(Duration::from_millis(60)).await;
assert_eq!(store.total_entries(), 0);
cancel_token.cancel();
let res = handle.await;
assert!(res.is_ok());
}
}