use std::convert::Infallible;
use axum::extract::FromRequestParts;
use serde::{Deserialize, Serialize};
use tower_sessions::Session as TowerSession;
pub struct Flash {
session: TowerSession,
messages: Vec<FlashMessage>,
data: std::collections::BTreeMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FlashMessage {
pub level: FlashLevel,
pub message: String,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum FlashLevel {
Success,
Error,
Warning,
Info,
}
const FLASH_KEY: &str = "_flash";
pub(crate) const FLASH_DATA_KEY: &str = "_flash_data";
impl Flash {
#[must_use]
pub fn messages(&self) -> &[FlashMessage] {
&self.messages
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.messages.is_empty() && self.data.is_empty()
}
#[must_use]
pub fn get(&self, key: &str) -> Option<&str> {
self.data.get(key).map(String::as_str)
}
#[must_use]
pub fn data(&self) -> &std::collections::BTreeMap<String, String> {
&self.data
}
pub async fn success(&self, message: &str) -> Result<(), FlashError> {
self.add(FlashLevel::Success, message).await
}
pub async fn error(&self, message: &str) -> Result<(), FlashError> {
self.add(FlashLevel::Error, message).await
}
pub async fn warning(&self, message: &str) -> Result<(), FlashError> {
self.add(FlashLevel::Warning, message).await
}
pub async fn info(&self, message: &str) -> Result<(), FlashError> {
self.add(FlashLevel::Info, message).await
}
async fn add(&self, level: FlashLevel, message: &str) -> Result<(), FlashError> {
let mut messages: Vec<FlashMessage> = self
.session
.get(FLASH_KEY)
.await
.map_err(|e| FlashError::Session(e.to_string()))?
.unwrap_or_default();
messages.push(FlashMessage {
level,
message: message.to_string(),
});
self.session
.insert(FLASH_KEY, &messages)
.await
.map_err(|e| FlashError::Session(e.to_string()))?;
Ok(())
}
}
impl<S> FromRequestParts<S> for Flash
where
S: Send + Sync,
{
type Rejection = Infallible;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection> {
let session = TowerSession::from_request_parts(parts, state)
.await
.map_err(|_| unreachable!("Session extraction is infallible"))?;
let messages: Vec<FlashMessage> = session
.get(FLASH_KEY)
.await
.map_err(|e| FlashError::Session(e.to_string()))
.unwrap_or(None)
.unwrap_or_default();
let _ = session.remove::<Vec<FlashMessage>>(FLASH_KEY).await;
let data: std::collections::BTreeMap<String, String> = session
.get(FLASH_DATA_KEY)
.await
.unwrap_or(None)
.unwrap_or_default();
if !data.is_empty() {
let _ = session
.remove::<std::collections::BTreeMap<String, String>>(FLASH_DATA_KEY)
.await;
}
Ok(Flash {
session,
messages,
data,
})
}
}
#[derive(Debug)]
pub enum FlashError {
Session(String),
}
impl std::fmt::Display for FlashError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Session(msg) => write!(f, "session error: {msg}"),
}
}
}
impl std::error::Error for FlashError {}