use std::sync::Arc;
use async_trait::async_trait;
use dashmap::DashMap;
use crate::error::Result;
#[async_trait]
pub trait StateStore: Send + Sync {
async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()>;
async fn retrieve(&self, state: &str) -> Result<(String, u64)>;
}
#[derive(Debug)]
pub struct InMemoryStateStore {
pub(crate) states: Arc<DashMap<String, (String, u64)>>,
max_states: usize,
}
impl InMemoryStateStore {
const MAX_STATES: usize = 10_000;
#[must_use]
pub fn new() -> Self {
Self {
states: Arc::new(DashMap::new()),
max_states: Self::MAX_STATES,
}
}
#[must_use]
pub fn with_max_states(max_states: usize) -> Self {
Self {
states: Arc::new(DashMap::new()),
max_states: max_states.max(1), }
}
fn cleanup_expired(&self) -> Result<bool> {
let Ok(now) = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
else {
return Err(crate::error::AuthError::ConfigError {
message: "system clock error: cannot validate state TTLs".to_string(),
});
};
self.states.retain(|_key, (_provider, expiry)| *expiry > now);
Ok(self.states.len() >= self.max_states)
}
fn evict_oldest(&self) -> bool {
let oldest = self.states.iter().min_by_key(|e| e.value().1).map(|e| e.key().clone());
match oldest {
Some(key) => self.states.remove(&key).is_some(),
None => false,
}
}
}
impl Default for InMemoryStateStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl StateStore for InMemoryStateStore {
async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
if self.cleanup_expired()? {
self.evict_oldest();
}
self.states.insert(state, (provider, expiry_secs));
Ok(())
}
async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
let (_key, value) =
self.states.remove(state).ok_or_else(|| crate::error::AuthError::InvalidState)?;
Ok(value)
}
}
#[cfg(feature = "redis-rate-limiting")]
#[derive(Clone)]
pub struct RedisStateStore {
client: redis::aio::ConnectionManager,
}
#[cfg(feature = "redis-rate-limiting")]
impl RedisStateStore {
pub async fn new(redis_url: &str) -> Result<Self> {
let client =
redis::Client::open(redis_url).map_err(|e| crate::error::AuthError::ConfigError {
message: e.to_string(),
})?;
let connection_manager = client.get_connection_manager().await.map_err(|e| {
crate::error::AuthError::ConfigError {
message: e.to_string(),
}
})?;
Ok(Self {
client: connection_manager,
})
}
fn state_key(state: &str) -> String {
format!("oauth:state:{}", state)
}
}
#[cfg(feature = "redis-rate-limiting")]
#[async_trait]
impl StateStore for RedisStateStore {
async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
use redis::AsyncCommands;
let key = Self::state_key(&state);
let ttl = expiry_secs
.saturating_sub(
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
)
.max(1);
let mut conn = self.client.clone();
let _: () = conn.set_ex(&key, &provider, ttl).await.map_err(|e| {
crate::error::AuthError::ConfigError {
message: e.to_string(),
}
})?;
Ok(())
}
async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
use redis::AsyncCommands;
let key = Self::state_key(state);
let mut conn = self.client.clone();
let provider: Option<String> =
conn.get_del(&key).await.map_err(|e| crate::error::AuthError::ConfigError {
message: e.to_string(),
})?;
let provider = provider.ok_or(crate::error::AuthError::InvalidState)?;
let expiry_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Ok((provider, expiry_secs))
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod lru_eviction_tests {
use super::*;
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("clock")
.as_secs()
}
#[tokio::test]
async fn store_evicts_oldest_at_capacity_instead_of_rejecting() {
let store = InMemoryStateStore::with_max_states(2);
let now = now_secs();
store.store("s1".into(), "p".into(), now + 100).await.unwrap();
store.store("s2".into(), "p".into(), now + 200).await.unwrap();
store.store("s3".into(), "p".into(), now + 300).await.unwrap();
assert_eq!(store.states.len(), 2, "store should stay at capacity, not grow");
assert!(store.retrieve("s1").await.is_err(), "oldest (s1) should have been evicted");
assert!(store.retrieve("s3").await.is_ok(), "newest (s3) should be present");
}
}