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) -> bool {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.states.retain(|_key, (_provider, expiry)| *expiry > now);
self.states.len() >= self.max_states
}
}
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() {
return Err(crate::error::AuthError::ConfigError {
message: "State store at capacity, cannot store new state".to_string(),
});
}
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))
}
}