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 /// Clean up expired states and check capacity
102 ///
103 /// # SECURITY
104 /// Called before inserting new states to:
105 /// 1. Remove expired states (automatic cleanup)
106 /// 2. Check if store is at capacity
107 /// 3. Return eviction needed flag if cleanup doesn't free space
108 fn cleanup_expired(&self) -> bool {
109 let now = std::time::SystemTime::now()
110 .duration_since(std::time::UNIX_EPOCH)
111 .unwrap_or_default()
112 .as_secs();
113
114 // Remove all expired states
115 self.states.retain(|_key, (_provider, expiry)| *expiry > now);
116
117 // Return true if we're still over capacity after cleanup
118 self.states.len() >= self.max_states
119 }
120}
121
122impl Default for InMemoryStateStore {
123 fn default() -> Self {
124 Self::new()
125 }
126}
127
128// Reason: StateStore is defined with #[async_trait]; all implementations must match
129// its transformed method signatures to satisfy the trait contract
130// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
131#[async_trait]
132impl StateStore for InMemoryStateStore {
133 async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
134 // SECURITY: Clean up expired states before inserting new one
135 if self.cleanup_expired() {
136 // Still over capacity after cleanup - reject to prevent memory exhaustion
137 return Err(crate::error::AuthError::ConfigError {
138 message: "State store at capacity, cannot store new state".to_string(),
139 });
140 }
141
142 self.states.insert(state, (provider, expiry_secs));
143 Ok(())
144 }
145
146 async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
147 let (_key, value) =
148 self.states.remove(state).ok_or_else(|| crate::error::AuthError::InvalidState)?;
149 Ok(value)
150 }
151}
152
153/// Redis-backed state store for distributed deployments
154///
155/// Uses Redis to store OAuth state parameters, allowing state validation
156/// across multiple server instances. Automatically expires states after TTL.
157#[cfg(feature = "redis-rate-limiting")]
158#[derive(Clone)]
159pub struct RedisStateStore {
160 client: redis::aio::ConnectionManager,
161}
162
163#[cfg(feature = "redis-rate-limiting")]
164impl RedisStateStore {
165 /// Create a new Redis state store
166 ///
167 /// # Arguments
168 /// * `redis_url` - Connection string (e.g., "redis://localhost:6379")
169 ///
170 /// # Example
171 /// ```no_run
172 /// // Requires: live Redis server.
173 /// # async fn example() -> fraiseql_auth::error::Result<()> {
174 /// use fraiseql_auth::state_store::RedisStateStore;
175 /// let store = RedisStateStore::new("redis://localhost:6379").await?;
176 /// # Ok(())
177 /// # }
178 /// ```
179 /// # Errors
180 ///
181 /// Returns [`AuthError::ConfigError`](crate::error::AuthError::ConfigError) if the Redis URL is
182 /// invalid or if the connection manager cannot be established.
183 pub async fn new(redis_url: &str) -> Result<Self> {
184 let client =
185 redis::Client::open(redis_url).map_err(|e| crate::error::AuthError::ConfigError {
186 message: e.to_string(),
187 })?;
188
189 let connection_manager = client.get_connection_manager().await.map_err(|e| {
190 crate::error::AuthError::ConfigError {
191 message: e.to_string(),
192 }
193 })?;
194
195 Ok(Self {
196 client: connection_manager,
197 })
198 }
199
200 /// Get Redis key for state
201 fn state_key(state: &str) -> String {
202 format!("oauth:state:{}", state)
203 }
204}
205
206#[cfg(feature = "redis-rate-limiting")]
207// Reason: StateStore is defined with #[async_trait]; all implementations must match
208// its transformed method signatures to satisfy the trait contract
209// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
210#[async_trait]
211impl StateStore for RedisStateStore {
212 async fn store(&self, state: String, provider: String, expiry_secs: u64) -> Result<()> {
213 use redis::AsyncCommands;
214
215 let key = Self::state_key(&state);
216 let ttl = expiry_secs
217 .saturating_sub(
218 std::time::SystemTime::now()
219 .duration_since(std::time::UNIX_EPOCH)
220 .unwrap_or_default()
221 .as_secs(),
222 )
223 .max(1); // Minimum 1 second TTL
224
225 let mut conn = self.client.clone();
226 let _: () = conn.set_ex(&key, &provider, ttl).await.map_err(|e| {
227 crate::error::AuthError::ConfigError {
228 message: e.to_string(),
229 }
230 })?;
231
232 Ok(())
233 }
234
235 async fn retrieve(&self, state: &str) -> Result<(String, u64)> {
236 use redis::AsyncCommands;
237
238 let key = Self::state_key(state);
239 let mut conn = self.client.clone();
240
241 // SECURITY: Use GETDEL (atomic get-and-delete, Redis ≥6.2) to prevent the
242 // GET+DEL race condition where two concurrent requests could both read the
243 // same state token before either deletes it, enabling replay attacks.
244 let provider: Option<String> =
245 conn.get_del(&key).await.map_err(|e| crate::error::AuthError::ConfigError {
246 message: e.to_string(),
247 })?;
248
249 let provider = provider.ok_or(crate::error::AuthError::InvalidState)?;
250
251 // Return current time as expiry (it was already validated by Redis TTL)
252 let expiry_secs = std::time::SystemTime::now()
253 .duration_since(std::time::UNIX_EPOCH)
254 .unwrap_or_default()
255 .as_secs();
256
257 Ok((provider, expiry_secs))
258 }
259}