Skip to main content

fraiseql_auth/
state_store.rs

1//! CSRF state store — trait definition and backends.
2//!
3//! Stores OAuth `state` parameters for the duration of an authorization flow and
4//! removes them on first retrieval, preventing state-replay attacks.
5
6use std::sync::Arc;
7
8use async_trait::async_trait;
9use dashmap::DashMap;
10
11use crate::error::Result;
12
13/// StateStore trait - implement this for different storage backends
14///
15/// Stores OAuth state parameters with expiration for CSRF protection.
16/// In distributed deployments, use a persistent backend (Redis) instead of in-memory.
17///
18/// # Examples
19///
20/// Use in-memory store for single-instance deployments:
21/// ```rust
22/// use std::sync::Arc;
23/// use fraiseql_auth::state_store::InMemoryStateStore;
24/// let state_store = Arc::new(InMemoryStateStore::new());
25/// ```
26///
27/// Use Redis for distributed deployments:
28/// ```no_run
29/// // Requires: live Redis server.
30/// use std::sync::Arc;
31/// # async fn example() -> fraiseql_auth::error::Result<()> {
32/// # #[cfg(feature = "redis-rate-limiting")] {
33/// use fraiseql_auth::state_store::RedisStateStore;
34/// let state_store = Arc::new(RedisStateStore::new("redis://localhost:6379").await?);
35/// # }
36/// # Ok(())
37/// # }
38/// ```
39// Reason: used as dyn Trait (Arc<dyn StateStore>); async_trait ensures Send bounds and
40// dyn-compatibility async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
41#[async_trait]
42pub trait StateStore: Send + Sync {
43    /// Store a state value with provider and expiration
44    ///
45    /// # Arguments
46    /// * `state` - The state parameter value
47    /// * `provider` - OAuth provider name
48    /// * `expiry_secs` - Unix timestamp when this state expires
49    async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()>;
50
51    /// Retrieve and remove a state value
52    ///
53    /// Returns (provider, expiry_secs) if state exists and is valid
54    /// Returns error if state doesn't exist or is invalid
55    async fn retrieve(&self, state: &str) -> Result<(String, u64)>;
56}
57
58/// In-memory state store using DashMap
59///
60/// **Warning**: Only suitable for single-instance deployments!
61/// For distributed systems, use RedisStateStore instead.
62///
63/// # SECURITY
64/// - Bounded to MAX_STATES entries to prevent unbounded memory growth
65/// - Expired states are automatically cleaned up on store operations
66/// - Implements LRU-like eviction when max capacity is reached
67#[derive(Debug)]
68pub struct InMemoryStateStore {
69    // Map of state -> (provider, expiry_secs)
70    pub(crate) states: Arc<DashMap<String, (String, u64)>>,
71    // Maximum number of states to store (prevents memory exhaustion)
72    max_states:        usize,
73}
74
75impl InMemoryStateStore {
76    /// Default maximum number of states to store (10,000 states)
77    /// At ~100 bytes per state, this limits memory to ~1 MB
78    const MAX_STATES: usize = 10_000;
79
80    /// Create a new in-memory state store with default limits
81    #[must_use]
82    pub fn new() -> Self {
83        Self {
84            states:     Arc::new(DashMap::new()),
85            max_states: Self::MAX_STATES,
86        }
87    }
88
89    /// Create a new in-memory state store with custom max size
90    ///
91    /// # Arguments
92    /// * `max_states` - Maximum number of states to store
93    #[must_use]
94    pub fn with_max_states(max_states: usize) -> Self {
95        Self {
96            states:     Arc::new(DashMap::new()),
97            max_states: max_states.max(1), // Ensure at least 1 state
98        }
99    }
100
101    /// Remove expired states and report whether the store is still at capacity.
102    ///
103    /// # Errors
104    ///
105    /// Returns [`AuthError::ConfigError`](crate::error::AuthError::ConfigError) if the
106    /// system clock cannot be read. This fails closed: a new state cannot be admitted
107    /// when state TTLs cannot be validated, and existing (possibly valid) in-flight
108    /// states are left intact rather than purged.
109    fn cleanup_expired(&self) -> Result<bool> {
110        let Ok(now) = std::time::SystemTime::now()
111            .duration_since(std::time::UNIX_EPOCH)
112            .map(|d| d.as_secs())
113        else {
114            return Err(crate::error::AuthError::ConfigError {
115                message: "system clock error: cannot validate state TTLs".to_string(),
116            });
117        };
118
119        // Remove all expired states.
120        self.states.retain(|_key, (_provider, expiry)| *expiry > now);
121
122        // Report whether we're still at capacity after cleanup.
123        Ok(self.states.len() >= self.max_states)
124    }
125
126    /// Remove the oldest (smallest-expiry) state. Returns `true` if one was removed.
127    ///
128    /// The iterator reference is dropped (the key is cloned) before `remove` is called,
129    /// so this never deadlocks the `DashMap`.
130    fn evict_oldest(&self) -> bool {
131        let oldest = self.states.iter().min_by_key(|e| e.value().1).map(|e| e.key().clone());
132        match oldest {
133            Some(key) => self.states.remove(&key).is_some(),
134            None => false,
135        }
136    }
137}
138
139impl Default for InMemoryStateStore {
140    fn default() -> Self {
141        Self::new()
142    }
143}
144
145// Reason: StateStore is defined with #[async_trait]; all implementations must match
146// its transformed method signatures to satisfy the trait contract
147// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
148#[async_trait]
149impl StateStore for InMemoryStateStore {
150    async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
151        // Remove expired states first; fail closed if the clock cannot be read.
152        if self.cleanup_expired()? {
153            // Still at capacity after cleanup — evict the oldest (smallest-expiry) state
154            // to admit the new authorization flow rather than rejecting it with a 500.
155            // The map stays bounded (one out, one in), so the memory bound holds while
156            // new logins keep working under load (L-state-store-doc: the struct docs
157            // already promised LRU-style eviction).
158            self.evict_oldest();
159        }
160
161        self.states.insert(state, (provider, expiry_secs));
162        Ok(())
163    }
164
165    async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
166        let (_key, value) =
167            self.states.remove(state).ok_or_else(|| crate::error::AuthError::InvalidState)?;
168        Ok(value)
169    }
170}
171
172/// Redis-backed state store for distributed deployments
173///
174/// Uses Redis to store OAuth state parameters, allowing state validation
175/// across multiple server instances. Automatically expires states after TTL.
176#[cfg(feature = "redis-rate-limiting")]
177#[derive(Clone)]
178pub struct RedisStateStore {
179    client: redis::aio::ConnectionManager,
180}
181
182#[cfg(feature = "redis-rate-limiting")]
183impl RedisStateStore {
184    /// Create a new Redis state store
185    ///
186    /// # Arguments
187    /// * `redis_url` - Connection string (e.g., "redis://localhost:6379")
188    ///
189    /// # Example
190    /// ```no_run
191    /// // Requires: live Redis server.
192    /// # async fn example() -> fraiseql_auth::error::Result<()> {
193    /// use fraiseql_auth::state_store::RedisStateStore;
194    /// let store = RedisStateStore::new("redis://localhost:6379").await?;
195    /// # Ok(())
196    /// # }
197    /// ```
198    /// # Errors
199    ///
200    /// Returns [`AuthError::ConfigError`](crate::error::AuthError::ConfigError) if the Redis URL is
201    /// invalid or if the connection manager cannot be established.
202    pub async fn new(redis_url: &str) -> Result<Self> {
203        let client =
204            redis::Client::open(redis_url).map_err(|e| crate::error::AuthError::ConfigError {
205                message: e.to_string(),
206            })?;
207
208        let connection_manager = client.get_connection_manager().await.map_err(|e| {
209            crate::error::AuthError::ConfigError {
210                message: e.to_string(),
211            }
212        })?;
213
214        Ok(Self {
215            client: connection_manager,
216        })
217    }
218
219    /// Get Redis key for state
220    fn state_key(state: &str) -> String {
221        format!("oauth:state:{}", state)
222    }
223}
224
225#[cfg(feature = "redis-rate-limiting")]
226// Reason: StateStore is defined with #[async_trait]; all implementations must match
227// its transformed method signatures to satisfy the trait contract
228// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
229#[async_trait]
230impl StateStore for RedisStateStore {
231    async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
232        use redis::AsyncCommands;
233
234        let key = Self::state_key(&state);
235        let ttl = expiry_secs
236            .saturating_sub(
237                std::time::SystemTime::now()
238                    .duration_since(std::time::UNIX_EPOCH)
239                    .unwrap_or_default()
240                    .as_secs(),
241            )
242            .max(1); // Minimum 1 second TTL
243
244        let mut conn = self.client.clone();
245        let _: () = conn.set_ex(&key, &provider, ttl).await.map_err(|e| {
246            crate::error::AuthError::ConfigError {
247                message: e.to_string(),
248            }
249        })?;
250
251        Ok(())
252    }
253
254    async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
255        use redis::AsyncCommands;
256
257        let key = Self::state_key(state);
258        let mut conn = self.client.clone();
259
260        // SECURITY: Use GETDEL (atomic get-and-delete, Redis ≥6.2) to prevent the
261        // GET+DEL race condition where two concurrent requests could both read the
262        // same state token before either deletes it, enabling replay attacks.
263        let provider: Option<String> =
264            conn.get_del(&key).await.map_err(|e| crate::error::AuthError::ConfigError {
265                message: e.to_string(),
266            })?;
267
268        let provider = provider.ok_or(crate::error::AuthError::InvalidState)?;
269
270        // Return current time as expiry (it was already validated by Redis TTL)
271        let expiry_secs = std::time::SystemTime::now()
272            .duration_since(std::time::UNIX_EPOCH)
273            .unwrap_or_default()
274            .as_secs();
275
276        Ok((provider, expiry_secs))
277    }
278}
279
280#[cfg(test)]
281#[allow(clippy::unwrap_used)] // Reason: test code, panics are acceptable
282mod lru_eviction_tests {
283    use super::*;
284
285    fn now_secs() -> u64 {
286        std::time::SystemTime::now()
287            .duration_since(std::time::UNIX_EPOCH)
288            .expect("clock")
289            .as_secs()
290    }
291
292    // L-state-store-doc: at capacity the store must evict the oldest entry (LRU), not
293    // reject new authorization flows with a 500.
294    #[tokio::test]
295    async fn store_evicts_oldest_at_capacity_instead_of_rejecting() {
296        let store = InMemoryStateStore::with_max_states(2);
297        let now = now_secs();
298        store.store("s1".into(), "p".into(), now + 100).await.unwrap();
299        store.store("s2".into(), "p".into(), now + 200).await.unwrap();
300        // Third insert at capacity must SUCCEED by evicting the oldest (s1).
301        store.store("s3".into(), "p".into(), now + 300).await.unwrap();
302
303        assert_eq!(store.states.len(), 2, "store should stay at capacity, not grow");
304        assert!(store.retrieve("s1").await.is_err(), "oldest (s1) should have been evicted");
305        assert!(store.retrieve("s3").await.is_ok(), "newest (s3) should be present");
306    }
307}