1use 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
27pub const MESSAGE_HEADER: &str = "X-Messages";
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
32pub enum MessageLevel {
33 Debug,
35 Info,
37 Success,
39 Warning,
41 Error,
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct Message {
48 pub level: MessageLevel,
50 pub text: String,
52}
53
54impl Message {
55 pub fn new(level: MessageLevel, text: String) -> Self {
66 Self { level, text }
67 }
68
69 pub fn debug(text: String) -> Self {
71 Self::new(MessageLevel::Debug, text)
72 }
73
74 pub fn info(text: String) -> Self {
76 Self::new(MessageLevel::Info, text)
77 }
78
79 pub fn success(text: String) -> Self {
81 Self::new(MessageLevel::Success, text)
82 }
83
84 pub fn warning(text: String) -> Self {
86 Self::new(MessageLevel::Warning, text)
87 }
88
89 pub fn error(text: String) -> Self {
91 Self::new(MessageLevel::Error, text)
92 }
93}
94
95pub trait MessageStorage: Send + Sync {
97 fn add_message(&self, session_id: &str, message: Message);
99 fn get_and_clear_messages(&self, session_id: &str) -> Vec<Message>;
101 fn get_messages(&self, session_id: &str) -> Vec<Message>;
103}
104
105pub struct SessionStorage {
110 messages: Arc<RwLock<HashMap<String, Vec<Message>>>>,
111}
112
113impl SessionStorage {
114 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
156pub struct CookieStorage {
161 messages: Arc<RwLock<HashMap<String, Vec<Message>>>>,
162}
163
164impl CookieStorage {
165 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#[allow(dead_code)]
258pub struct MessageMiddleware {
259 storage: Arc<dyn MessageStorage>,
260}
261
262impl MessageMiddleware {
263 pub fn new(storage: Arc<dyn MessageStorage>) -> Self {
275 Self { storage }
276 }
277
278 fn get_session_id(request: &Request) -> String {
283 if let Some(session_data) = request.extensions.get::<SessionData>() {
285 return session_data.id.clone();
286 }
287
288 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 let response = match handler.handle(request).await {
316 Ok(resp) => resp,
317 Err(e) => Response::from(e),
318 };
319
320 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 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 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 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 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 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 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}