reinhardt_auth/sessions/csrf.rs
1//! CSRF protection integration with sessions
2//!
3//! This module provides integration between session management and CSRF protection
4//! from reinhardt-forms. CSRF tokens are stored in sessions for validation.
5//!
6//! ## Example
7//!
8//! ```rust
9//! use reinhardt_auth::sessions::csrf::CsrfSessionManager;
10//! use reinhardt_auth::sessions::Session;
11//! use reinhardt_auth::sessions::backends::InMemorySessionBackend;
12//!
13//! # async fn example() -> Result<(), Box<dyn std::error::Error>> {
14//! let backend = InMemorySessionBackend::new();
15//! let mut session = Session::new(backend);
16//!
17//! // Create CSRF manager
18//! let csrf_manager = CsrfSessionManager::new();
19//!
20//! // Generate and store CSRF token in session
21//! let token = csrf_manager.generate_token(&mut session)?;
22//! println!("CSRF token: {}", token);
23//!
24//! // Validate token from session
25//! let is_valid = csrf_manager.validate_token(&mut session, &token)?;
26//! assert!(is_valid);
27//! # Ok(())
28//! # }
29//! ```
30
31use super::backends::SessionBackend;
32use super::session::Session;
33use serde::{Deserialize, Serialize};
34use std::time::SystemTime;
35use subtle::ConstantTimeEq;
36use uuid::Uuid;
37
38/// CSRF token data stored in session
39#[derive(Debug, Clone, Serialize, Deserialize)]
40pub struct CsrfTokenData {
41 /// The token value
42 pub token: String,
43 /// When the token was created
44 pub created_at: SystemTime,
45}
46
47/// CSRF session manager
48///
49/// Manages CSRF tokens in sessions, integrating with reinhardt-forms.
50///
51/// # Example
52///
53/// ```rust
54/// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
55/// use reinhardt_auth::sessions::Session;
56/// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
57///
58/// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
59/// let backend = InMemorySessionBackend::new();
60/// let mut session = Session::new(backend);
61///
62/// let csrf = CsrfSessionManager::new();
63///
64/// // Generate token
65/// let token = csrf.generate_token(&mut session)?;
66///
67/// // Validate token
68/// assert!(csrf.validate_token(&mut session, &token)?);
69/// # Ok(())
70/// # }
71/// ```
72pub struct CsrfSessionManager {
73 /// Session key for storing CSRF token
74 session_key: String,
75}
76
77impl CsrfSessionManager {
78 /// Create a new CSRF session manager
79 ///
80 /// # Example
81 ///
82 /// ```rust
83 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
84 ///
85 /// let csrf = CsrfSessionManager::new();
86 /// ```
87 pub fn new() -> Self {
88 Self {
89 session_key: "_csrf_token".to_string(),
90 }
91 }
92
93 /// Create a new CSRF session manager with custom session key
94 ///
95 /// # Example
96 ///
97 /// ```rust
98 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
99 ///
100 /// let csrf = CsrfSessionManager::with_key("my_csrf_token".to_string());
101 /// ```
102 pub fn with_key(session_key: String) -> Self {
103 Self { session_key }
104 }
105
106 /// Generate a new CSRF token and store it in the session
107 ///
108 /// # Example
109 ///
110 /// ```rust
111 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
112 /// use reinhardt_auth::sessions::Session;
113 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
114 ///
115 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
116 /// let backend = InMemorySessionBackend::new();
117 /// let mut session = Session::new(backend);
118 ///
119 /// let csrf = CsrfSessionManager::new();
120 /// let token = csrf.generate_token(&mut session)?;
121 ///
122 /// assert!(!token.is_empty());
123 /// # Ok(())
124 /// # }
125 /// ```
126 pub fn generate_token<B: SessionBackend>(
127 &self,
128 session: &mut Session<B>,
129 ) -> Result<String, serde_json::Error> {
130 let token = Uuid::new_v4().to_string();
131 let token_data = CsrfTokenData {
132 token: token.clone(),
133 created_at: SystemTime::now(),
134 };
135
136 session.set(&self.session_key, token_data)?;
137 Ok(token)
138 }
139
140 /// Get the current CSRF token from the session
141 ///
142 /// # Example
143 ///
144 /// ```rust
145 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
146 /// use reinhardt_auth::sessions::Session;
147 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
148 ///
149 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
150 /// let backend = InMemorySessionBackend::new();
151 /// let mut session = Session::new(backend);
152 ///
153 /// let csrf = CsrfSessionManager::new();
154 ///
155 /// // Generate token first
156 /// let generated = csrf.generate_token(&mut session)?;
157 ///
158 /// // Get the stored token
159 /// let stored = csrf.get_token(&mut session)?;
160 /// assert_eq!(stored, Some(generated));
161 /// # Ok(())
162 /// # }
163 /// ```
164 pub fn get_token<B: SessionBackend>(
165 &self,
166 session: &mut Session<B>,
167 ) -> Result<Option<String>, serde_json::Error> {
168 let token_data: Option<CsrfTokenData> = session.get(&self.session_key)?;
169 Ok(token_data.map(|data| data.token))
170 }
171
172 /// Validate a CSRF token against the session
173 ///
174 /// # Example
175 ///
176 /// ```rust
177 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
178 /// use reinhardt_auth::sessions::Session;
179 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
180 ///
181 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
182 /// let backend = InMemorySessionBackend::new();
183 /// let mut session = Session::new(backend);
184 ///
185 /// let csrf = CsrfSessionManager::new();
186 /// let token = csrf.generate_token(&mut session)?;
187 ///
188 /// // Valid token
189 /// assert!(csrf.validate_token(&mut session, &token)?);
190 ///
191 /// // Invalid token
192 /// assert!(!csrf.validate_token(&mut session, "wrong_token")?);
193 /// # Ok(())
194 /// # }
195 /// ```
196 pub fn validate_token<B: SessionBackend>(
197 &self,
198 session: &mut Session<B>,
199 submitted_token: &str,
200 ) -> Result<bool, serde_json::Error> {
201 let stored_token = self.get_token(session)?;
202
203 match stored_token {
204 Some(token) => {
205 // Use constant-time comparison to prevent timing attacks
206 Ok(token.as_bytes().ct_eq(submitted_token.as_bytes()).into())
207 }
208 None => Ok(false),
209 }
210 }
211
212 /// Rotate the CSRF token (generate a new one)
213 ///
214 /// This is useful after login or privilege escalation.
215 ///
216 /// # Example
217 ///
218 /// ```rust
219 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
220 /// use reinhardt_auth::sessions::Session;
221 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
222 ///
223 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
224 /// let backend = InMemorySessionBackend::new();
225 /// let mut session = Session::new(backend);
226 ///
227 /// let csrf = CsrfSessionManager::new();
228 ///
229 /// let old_token = csrf.generate_token(&mut session)?;
230 /// let new_token = csrf.rotate_token(&mut session)?;
231 ///
232 /// assert_ne!(old_token, new_token);
233 ///
234 /// // Old token should no longer be valid
235 /// assert!(!csrf.validate_token(&mut session, &old_token)?);
236 /// // New token should be valid
237 /// assert!(csrf.validate_token(&mut session, &new_token)?);
238 /// # Ok(())
239 /// # }
240 /// ```
241 pub fn rotate_token<B: SessionBackend>(
242 &self,
243 session: &mut Session<B>,
244 ) -> Result<String, serde_json::Error> {
245 // Simply generate a new token, which overwrites the old one
246 self.generate_token(session)
247 }
248
249 /// Clear the CSRF token from the session
250 ///
251 /// # Example
252 ///
253 /// ```rust
254 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
255 /// use reinhardt_auth::sessions::Session;
256 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
257 ///
258 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
259 /// let backend = InMemorySessionBackend::new();
260 /// let mut session = Session::new(backend);
261 ///
262 /// let csrf = CsrfSessionManager::new();
263 /// csrf.generate_token(&mut session)?;
264 ///
265 /// csrf.clear_token(&mut session);
266 ///
267 /// assert!(csrf.get_token(&mut session)?.is_none());
268 /// # Ok(())
269 /// # }
270 /// ```
271 pub fn clear_token<B: SessionBackend>(&self, session: &mut Session<B>) {
272 session.delete(&self.session_key);
273 }
274
275 /// Get or create a CSRF token
276 ///
277 /// Returns the existing token if available, otherwise generates a new one.
278 ///
279 /// # Example
280 ///
281 /// ```rust
282 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
283 /// use reinhardt_auth::sessions::Session;
284 /// use reinhardt_auth::sessions::backends::InMemorySessionBackend;
285 ///
286 /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
287 /// let backend = InMemorySessionBackend::new();
288 /// let mut session = Session::new(backend);
289 ///
290 /// let csrf = CsrfSessionManager::new();
291 ///
292 /// let token1 = csrf.get_or_create_token(&mut session)?;
293 /// let token2 = csrf.get_or_create_token(&mut session)?;
294 ///
295 /// // Should be the same token
296 /// assert_eq!(token1, token2);
297 /// # Ok(())
298 /// # }
299 /// ```
300 pub fn get_or_create_token<B: SessionBackend>(
301 &self,
302 session: &mut Session<B>,
303 ) -> Result<String, serde_json::Error> {
304 if let Some(token) = self.get_token(session)? {
305 Ok(token)
306 } else {
307 self.generate_token(session)
308 }
309 }
310}
311
312impl Default for CsrfSessionManager {
313 /// Create default CSRF session manager
314 ///
315 /// # Example
316 ///
317 /// ```rust
318 /// use reinhardt_auth::sessions::csrf::CsrfSessionManager;
319 ///
320 /// let csrf = CsrfSessionManager::default();
321 /// ```
322 fn default() -> Self {
323 Self::new()
324 }
325}
326
327#[cfg(test)]
328mod tests {
329 use super::*;
330 use crate::sessions::InMemorySessionBackend;
331
332 #[tokio::test]
333 async fn test_csrf_manager_new() {
334 let _csrf = CsrfSessionManager::new();
335 }
336
337 #[tokio::test]
338 async fn test_csrf_manager_with_key() {
339 let csrf = CsrfSessionManager::with_key("custom_key".to_string());
340 assert_eq!(csrf.session_key, "custom_key");
341 }
342
343 #[tokio::test]
344 async fn test_generate_token() {
345 let backend = InMemorySessionBackend::new();
346 let mut session = Session::new(backend);
347
348 let csrf = CsrfSessionManager::new();
349 let token = csrf.generate_token(&mut session).unwrap();
350
351 assert!(!token.is_empty());
352 }
353
354 #[tokio::test]
355 async fn test_get_token() {
356 let backend = InMemorySessionBackend::new();
357 let mut session = Session::new(backend);
358
359 let csrf = CsrfSessionManager::new();
360
361 // No token initially
362 assert!(csrf.get_token(&mut session).unwrap().is_none());
363
364 // Generate token
365 let generated = csrf.generate_token(&mut session).unwrap();
366
367 // Get the stored token
368 let stored = csrf.get_token(&mut session).unwrap();
369 assert_eq!(stored, Some(generated));
370 }
371
372 #[tokio::test]
373 async fn test_validate_token() {
374 let backend = InMemorySessionBackend::new();
375 let mut session = Session::new(backend);
376
377 let csrf = CsrfSessionManager::new();
378 let token = csrf.generate_token(&mut session).unwrap();
379
380 // Valid token
381 assert!(csrf.validate_token(&mut session, &token).unwrap());
382
383 // Invalid token
384 assert!(!csrf.validate_token(&mut session, "wrong_token").unwrap());
385 }
386
387 #[tokio::test]
388 async fn test_validate_token_no_token_in_session() {
389 let backend = InMemorySessionBackend::new();
390 let mut session = Session::new(backend);
391
392 let csrf = CsrfSessionManager::new();
393
394 // No token in session
395 assert!(!csrf.validate_token(&mut session, "any_token").unwrap());
396 }
397
398 #[tokio::test]
399 async fn test_rotate_token() {
400 let backend = InMemorySessionBackend::new();
401 let mut session = Session::new(backend);
402
403 let csrf = CsrfSessionManager::new();
404
405 let old_token = csrf.generate_token(&mut session).unwrap();
406 let new_token = csrf.rotate_token(&mut session).unwrap();
407
408 assert_ne!(old_token, new_token);
409
410 // Old token should no longer be valid
411 assert!(!csrf.validate_token(&mut session, &old_token).unwrap());
412
413 // New token should be valid
414 assert!(csrf.validate_token(&mut session, &new_token).unwrap());
415 }
416
417 #[tokio::test]
418 async fn test_clear_token() {
419 let backend = InMemorySessionBackend::new();
420 let mut session = Session::new(backend);
421
422 let csrf = CsrfSessionManager::new();
423 csrf.generate_token(&mut session).unwrap();
424
425 csrf.clear_token(&mut session);
426
427 assert!(csrf.get_token(&mut session).unwrap().is_none());
428 }
429
430 #[tokio::test]
431 async fn test_get_or_create_token() {
432 let backend = InMemorySessionBackend::new();
433 let mut session = Session::new(backend);
434
435 let csrf = CsrfSessionManager::new();
436
437 let token1 = csrf.get_or_create_token(&mut session).unwrap();
438 let token2 = csrf.get_or_create_token(&mut session).unwrap();
439
440 // Should be the same token
441 assert_eq!(token1, token2);
442 }
443
444 #[tokio::test]
445 async fn test_get_or_create_token_creates_if_missing() {
446 let backend = InMemorySessionBackend::new();
447 let mut session = Session::new(backend);
448
449 let csrf = CsrfSessionManager::new();
450
451 // No token initially
452 assert!(csrf.get_token(&mut session).unwrap().is_none());
453
454 // get_or_create should create one
455 let token = csrf.get_or_create_token(&mut session).unwrap();
456 assert!(!token.is_empty());
457
458 // Should now exist
459 assert!(csrf.get_token(&mut session).unwrap().is_some());
460 }
461}