Skip to main content

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}