Skip to main content

reinhardt_middleware/
messages.rs

1//! Messages middleware
2//!
3//! Provides Django-style flash messages for one-time notifications.
4//! Messages can be stored in sessions or cookies and are displayed once.
5//!
6//! # Middleware Ordering
7//!
8//! For best results, ensure `SessionMiddleware` runs before `MessageMiddleware`:
9//!
10//! ```ignore
11//! app.middleware(SessionMiddleware::new(config))
12//!    .middleware(MessageMiddleware::new(storage));
13//! ```
14//!
15//! This allows `MessageMiddleware` to use the session ID from
16//! `SessionMiddleware` via request extensions.
17
18use async_trait::async_trait;
19use hyper::header::COOKIE;
20use reinhardt_http::{Handler, Middleware, Request, Response, Result};
21use serde::{Deserialize, Serialize};
22use std::collections::HashMap;
23use std::sync::{Arc, RwLock};
24
25use crate::session::SessionData;
26
27/// Message header for passing messages between middleware and handlers
28pub const MESSAGE_HEADER: &str = "X-Messages";
29
30/// Message severity levels
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
32pub enum MessageLevel {
33	/// Diagnostic information for developers, not shown to end users.
34	Debug,
35	/// Informational message for the user.
36	Info,
37	/// Indicates an operation completed successfully.
38	Success,
39	/// Indicates a potential issue that the user should be aware of.
40	Warning,
41	/// Indicates an operation failed or an error occurred.
42	Error,
43}
44
45/// A single flash message
46#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct Message {
48	/// The severity level of this message.
49	pub level: MessageLevel,
50	/// The message text to display to the user.
51	pub text: String,
52}
53
54impl Message {
55	/// Create a new message
56	///
57	/// # Examples
58	///
59	/// ```
60	/// use reinhardt_middleware::messages::{Message, MessageLevel};
61	///
62	/// let msg = Message::new(MessageLevel::Success, "Saved successfully!".to_string());
63	/// assert_eq!(msg.level, MessageLevel::Success);
64	/// ```
65	pub fn new(level: MessageLevel, text: String) -> Self {
66		Self { level, text }
67	}
68
69	/// Create a debug message
70	pub fn debug(text: String) -> Self {
71		Self::new(MessageLevel::Debug, text)
72	}
73
74	/// Create an info message
75	pub fn info(text: String) -> Self {
76		Self::new(MessageLevel::Info, text)
77	}
78
79	/// Create a success message
80	pub fn success(text: String) -> Self {
81		Self::new(MessageLevel::Success, text)
82	}
83
84	/// Create a warning message
85	pub fn warning(text: String) -> Self {
86		Self::new(MessageLevel::Warning, text)
87	}
88
89	/// Create an error message
90	pub fn error(text: String) -> Self {
91		Self::new(MessageLevel::Error, text)
92	}
93}
94
95/// Message storage trait
96pub trait MessageStorage: Send + Sync {
97	/// Add a message to storage
98	fn add_message(&self, session_id: &str, message: Message);
99	/// Get all messages for a session and clear them
100	fn get_and_clear_messages(&self, session_id: &str) -> Vec<Message>;
101	/// Get messages without clearing
102	fn get_messages(&self, session_id: &str) -> Vec<Message>;
103}
104
105/// Session-based message storage
106///
107/// Stores messages in memory keyed by session ID.
108/// In production, this should be backed by a persistent session store.
109pub struct SessionStorage {
110	messages: Arc<RwLock<HashMap<String, Vec<Message>>>>,
111}
112
113impl SessionStorage {
114	/// Create a new SessionStorage
115	///
116	/// # Examples
117	///
118	/// ```
119	/// use reinhardt_middleware::messages::SessionStorage;
120	///
121	/// let storage = SessionStorage::new();
122	/// ```
123	pub fn new() -> Self {
124		Self {
125			messages: Arc::new(RwLock::new(HashMap::new())),
126		}
127	}
128}
129
130impl Default for SessionStorage {
131	fn default() -> Self {
132		Self::new()
133	}
134}
135
136impl MessageStorage for SessionStorage {
137	fn add_message(&self, session_id: &str, message: Message) {
138		let mut messages = self.messages.write().unwrap_or_else(|e| e.into_inner());
139		messages
140			.entry(session_id.to_string())
141			.or_default()
142			.push(message);
143	}
144
145	fn get_and_clear_messages(&self, session_id: &str) -> Vec<Message> {
146		let mut messages = self.messages.write().unwrap_or_else(|e| e.into_inner());
147		messages.remove(session_id).unwrap_or_default()
148	}
149
150	fn get_messages(&self, session_id: &str) -> Vec<Message> {
151		let messages = self.messages.read().unwrap_or_else(|e| e.into_inner());
152		messages.get(session_id).cloned().unwrap_or_default()
153	}
154}
155
156/// Cookie-based message storage
157///
158/// Stores messages in memory similar to SessionStorage but designed
159/// to be serialized to cookies in production.
160pub struct CookieStorage {
161	messages: Arc<RwLock<HashMap<String, Vec<Message>>>>,
162}
163
164impl CookieStorage {
165	/// Create a new CookieStorage
166	///
167	/// # Examples
168	///
169	/// ```
170	/// use reinhardt_middleware::messages::CookieStorage;
171	///
172	/// let storage = CookieStorage::new();
173	/// ```
174	pub fn new() -> Self {
175		Self {
176			messages: Arc::new(RwLock::new(HashMap::new())),
177		}
178	}
179}
180
181impl Default for CookieStorage {
182	fn default() -> Self {
183		Self::new()
184	}
185}
186
187impl MessageStorage for CookieStorage {
188	fn add_message(&self, session_id: &str, message: Message) {
189		let mut messages = self.messages.write().unwrap_or_else(|e| e.into_inner());
190		messages
191			.entry(session_id.to_string())
192			.or_default()
193			.push(message);
194	}
195
196	fn get_and_clear_messages(&self, session_id: &str) -> Vec<Message> {
197		let mut messages = self.messages.write().unwrap_or_else(|e| e.into_inner());
198		messages.remove(session_id).unwrap_or_default()
199	}
200
201	fn get_messages(&self, session_id: &str) -> Vec<Message> {
202		let messages = self.messages.read().unwrap_or_else(|e| e.into_inner());
203		messages.get(session_id).cloned().unwrap_or_default()
204	}
205}
206
207/// Message framework middleware
208///
209/// Provides flash message functionality similar to Django's messages framework.
210///
211/// # Examples
212///
213/// ```
214/// use std::sync::Arc;
215/// use reinhardt_middleware::messages::{MessageMiddleware, SessionStorage, Message, MessageLevel};
216/// use reinhardt_http::{Handler, Middleware, Request, Response};
217/// use hyper::{StatusCode, Method, Version, HeaderMap};
218/// use bytes::Bytes;
219///
220/// struct TestHandler {
221///     storage: Arc<dyn reinhardt_middleware::messages::MessageStorage>,
222/// }
223///
224/// #[async_trait::async_trait]
225/// impl Handler for TestHandler {
226///     async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
227///         // Add a message
228///         self.storage.add_message("test-session", Message::success("Operation successful!".to_string()));
229///         Ok(Response::new(StatusCode::OK).with_body(Bytes::from("OK")))
230///     }
231/// }
232///
233/// # tokio_test::block_on(async {
234/// let storage: Arc<dyn reinhardt_middleware::messages::MessageStorage> = Arc::new(SessionStorage::new());
235/// let middleware = MessageMiddleware::new(storage.clone());
236/// let handler = Arc::new(TestHandler { storage: storage.clone() });
237///
238/// let mut headers = HeaderMap::new();
239/// headers.insert(hyper::header::COOKIE, "sessionid=test-session".parse().unwrap());
240///
241/// let request = Request::builder()
242///     .method(Method::GET)
243///     .uri("/page")
244///     .version(Version::HTTP_11)
245///     .headers(headers)
246///     .body(Bytes::new())
247///     .build()
248///     .unwrap();
249///
250/// let _response = middleware.process(request, handler).await.unwrap();
251/// let messages = storage.get_and_clear_messages("test-session");
252/// assert_eq!(messages.len(), 1);
253/// assert_eq!(messages[0].level, MessageLevel::Success);
254/// # });
255/// ```
256// Allow dead_code: public API for middleware pipeline integration, not yet wired into default stack
257#[allow(dead_code)]
258pub struct MessageMiddleware {
259	storage: Arc<dyn MessageStorage>,
260}
261
262impl MessageMiddleware {
263	/// Create a new MessageMiddleware with the given storage backend
264	///
265	/// # Examples
266	///
267	/// ```
268	/// use std::sync::Arc;
269	/// use reinhardt_middleware::messages::{MessageMiddleware, SessionStorage};
270	///
271	/// let storage = Arc::new(SessionStorage::new());
272	/// let middleware = MessageMiddleware::new(storage);
273	/// ```
274	pub fn new(storage: Arc<dyn MessageStorage>) -> Self {
275		Self { storage }
276	}
277
278	/// Extract session ID from request
279	///
280	/// This method first checks for `SessionData` in request extensions
281	/// (set by `SessionMiddleware`), then falls back to cookie extraction.
282	fn get_session_id(request: &Request) -> String {
283		// Check for SessionData set by SessionMiddleware
284		if let Some(session_data) = request.extensions.get::<SessionData>() {
285			return session_data.id.clone();
286		}
287
288		// Fallback: extract from cookie
289		request
290			.headers
291			.get(COOKIE)
292			.and_then(|c| c.to_str().ok())
293			.and_then(|cookies| {
294				for cookie in cookies.split(';') {
295					let cookie = cookie.trim();
296					if let Some((name, value)) = cookie.split_once('=')
297						&& name == "sessionid"
298					{
299						return Some(value.to_string());
300					}
301				}
302				None
303			})
304			.unwrap_or_else(|| "default".to_string())
305	}
306}
307
308#[async_trait]
309impl Middleware for MessageMiddleware {
310	async fn process(&self, request: Request, handler: Arc<dyn Handler>) -> Result<Response> {
311		let _session_id = Self::get_session_id(&request);
312
313		// Convert errors to responses so post-processing always runs,
314		// even when invoked outside MiddlewareChain. (#3244)
315		let response = match handler.handle(request).await {
316			Ok(resp) => resp,
317			Err(e) => Response::from(e),
318		};
319
320		// Messages are stored in the storage and can be retrieved by handlers
321		// In a complete implementation, we'd add messages to the response or template context
322		Ok(response)
323	}
324}
325
326#[cfg(test)]
327mod tests {
328	use super::*;
329	use bytes::Bytes;
330	use hyper::{HeaderMap, Method, StatusCode, Version};
331
332	#[test]
333	fn test_message_creation() {
334		let msg = Message::debug("Debug message".to_string());
335		assert_eq!(msg.level, MessageLevel::Debug);
336
337		let msg = Message::info("Info message".to_string());
338		assert_eq!(msg.level, MessageLevel::Info);
339
340		let msg = Message::success("Success message".to_string());
341		assert_eq!(msg.level, MessageLevel::Success);
342
343		let msg = Message::warning("Warning message".to_string());
344		assert_eq!(msg.level, MessageLevel::Warning);
345
346		let msg = Message::error("Error message".to_string());
347		assert_eq!(msg.level, MessageLevel::Error);
348	}
349
350	#[test]
351	fn test_session_storage_add_and_get() {
352		let storage = SessionStorage::new();
353		let session_id = "test-session";
354
355		storage.add_message(session_id, Message::info("Message 1".to_string()));
356		storage.add_message(session_id, Message::success("Message 2".to_string()));
357
358		let messages = storage.get_messages(session_id);
359		assert_eq!(messages.len(), 2);
360		assert_eq!(messages[0].level, MessageLevel::Info);
361		assert_eq!(messages[1].level, MessageLevel::Success);
362	}
363
364	#[test]
365	fn test_session_storage_clear() {
366		let storage = SessionStorage::new();
367		let session_id = "test-session";
368
369		storage.add_message(session_id, Message::info("Message 1".to_string()));
370		storage.add_message(session_id, Message::info("Message 2".to_string()));
371
372		let messages = storage.get_and_clear_messages(session_id);
373		assert_eq!(messages.len(), 2);
374
375		// Messages should be cleared
376		let messages = storage.get_messages(session_id);
377		assert_eq!(messages.len(), 0);
378	}
379
380	#[test]
381	fn test_cookie_storage_add_and_get() {
382		let storage = CookieStorage::new();
383		let session_id = "test-session";
384
385		storage.add_message(session_id, Message::warning("Warning 1".to_string()));
386		storage.add_message(session_id, Message::error("Error 1".to_string()));
387
388		let messages = storage.get_messages(session_id);
389		assert_eq!(messages.len(), 2);
390		assert_eq!(messages[0].level, MessageLevel::Warning);
391		assert_eq!(messages[1].level, MessageLevel::Error);
392	}
393
394	#[test]
395	fn test_cookie_storage_clear() {
396		let storage = CookieStorage::new();
397		let session_id = "test-session";
398
399		storage.add_message(session_id, Message::info("Info 1".to_string()));
400
401		let messages = storage.get_and_clear_messages(session_id);
402		assert_eq!(messages.len(), 1);
403
404		// Messages should be cleared
405		let messages = storage.get_messages(session_id);
406		assert_eq!(messages.len(), 0);
407	}
408
409	#[test]
410	fn test_separate_sessions() {
411		let storage = SessionStorage::new();
412
413		storage.add_message("session1", Message::info("Session 1 message".to_string()));
414		storage.add_message(
415			"session2",
416			Message::success("Session 2 message".to_string()),
417		);
418
419		let messages1 = storage.get_messages("session1");
420		let messages2 = storage.get_messages("session2");
421
422		assert_eq!(messages1.len(), 1);
423		assert_eq!(messages2.len(), 1);
424		assert_eq!(messages1[0].level, MessageLevel::Info);
425		assert_eq!(messages2[0].level, MessageLevel::Success);
426	}
427
428	struct TestHandler {
429		storage: Arc<dyn MessageStorage>,
430	}
431
432	#[async_trait]
433	impl Handler for TestHandler {
434		async fn handle(&self, request: Request) -> Result<Response> {
435			let session_id = MessageMiddleware::get_session_id(&request);
436			self.storage
437				.add_message(&session_id, Message::success("Test message".to_string()));
438			Ok(Response::new(StatusCode::OK).with_body(Bytes::from("OK")))
439		}
440	}
441
442	#[tokio::test]
443	async fn test_middleware_with_session_storage() {
444		let storage: Arc<dyn MessageStorage> = Arc::new(SessionStorage::new());
445		let middleware = MessageMiddleware::new(storage.clone());
446		let handler = Arc::new(TestHandler {
447			storage: storage.clone(),
448		});
449
450		let mut headers = HeaderMap::new();
451		headers.insert(COOKIE, "sessionid=test-session".parse().unwrap());
452
453		let request = Request::builder()
454			.method(Method::GET)
455			.uri("/page")
456			.version(Version::HTTP_11)
457			.headers(headers)
458			.body(Bytes::new())
459			.build()
460			.unwrap();
461
462		let response = middleware.process(request, handler).await.unwrap();
463		assert_eq!(response.status, StatusCode::OK);
464
465		// Verify message was stored
466		let messages = storage.get_and_clear_messages("test-session");
467		assert_eq!(messages.len(), 1);
468		assert_eq!(messages[0].level, MessageLevel::Success);
469	}
470
471	#[tokio::test]
472	async fn test_middleware_default_session() {
473		let storage: Arc<dyn MessageStorage> = Arc::new(SessionStorage::new());
474		let middleware = MessageMiddleware::new(storage.clone());
475		let handler = Arc::new(TestHandler {
476			storage: storage.clone(),
477		});
478
479		// No session cookie
480		let request = Request::builder()
481			.method(Method::GET)
482			.uri("/page")
483			.version(Version::HTTP_11)
484			.headers(HeaderMap::new())
485			.body(Bytes::new())
486			.build()
487			.unwrap();
488
489		let response = middleware.process(request, handler).await.unwrap();
490		assert_eq!(response.status, StatusCode::OK);
491
492		// Should use default session
493		let messages = storage.get_messages("default");
494		assert_eq!(messages.len(), 1);
495	}
496
497	#[tokio::test]
498	async fn test_middleware_with_cookie_storage() {
499		let storage: Arc<dyn MessageStorage> = Arc::new(CookieStorage::new());
500		let middleware = MessageMiddleware::new(storage.clone());
501		let handler = Arc::new(TestHandler {
502			storage: storage.clone(),
503		});
504
505		let mut headers = HeaderMap::new();
506		headers.insert(COOKIE, "sessionid=cookie-session".parse().unwrap());
507
508		let request = Request::builder()
509			.method(Method::GET)
510			.uri("/page")
511			.version(Version::HTTP_11)
512			.headers(headers)
513			.body(Bytes::new())
514			.build()
515			.unwrap();
516
517		let response = middleware.process(request, handler).await.unwrap();
518		assert_eq!(response.status, StatusCode::OK);
519
520		// Verify message was stored
521		let messages = storage.get_and_clear_messages("cookie-session");
522		assert_eq!(messages.len(), 1);
523		assert_eq!(messages[0].level, MessageLevel::Success);
524	}
525}