use axum::extract::FromRequestParts;
use http::request::Parts;
use serde::{Deserialize, Serialize};
use crate::session::Session;
const FLASH_SESSION_KEY: &str = "__autumn_flash";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum FlashLevel {
Success,
Info,
Warning,
Error,
}
impl FlashLevel {
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Success => "success",
Self::Info => "info",
Self::Warning => "warning",
Self::Error => "error",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FlashMessage {
pub level: FlashLevel,
pub message: String,
}
#[derive(Debug, Clone)]
pub struct Flash {
session: Session,
}
impl Flash {
#[must_use]
pub const fn new(session: Session) -> Self {
Self { session }
}
pub async fn push(&self, level: FlashLevel, message: impl Into<String>) {
let mut messages = self.peek().await;
messages.push(FlashMessage {
level,
message: message.into(),
});
if let Ok(json) = serde_json::to_string(&messages) {
self.session.insert(FLASH_SESSION_KEY, json).await;
}
}
pub async fn success(&self, message: impl Into<String>) {
self.push(FlashLevel::Success, message).await;
}
pub async fn info(&self, message: impl Into<String>) {
self.push(FlashLevel::Info, message).await;
}
pub async fn warning(&self, message: impl Into<String>) {
self.push(FlashLevel::Warning, message).await;
}
pub async fn error(&self, message: impl Into<String>) {
self.push(FlashLevel::Error, message).await;
}
pub async fn peek(&self) -> Vec<FlashMessage> {
self.session
.get(FLASH_SESSION_KEY)
.await
.map_or_else(Vec::new, |json| {
serde_json::from_str(&json).unwrap_or_default()
})
}
pub async fn consume(&self) -> Vec<FlashMessage> {
let messages = self.peek().await;
if !messages.is_empty() {
self.session.remove(FLASH_SESSION_KEY).await;
}
messages
}
#[cfg(feature = "htmx")]
pub async fn inject_hx_trigger<T: axum::response::IntoResponse>(
&self,
response: T,
) -> axum::response::Response {
let messages = self.consume().await;
let mut res = response.into_response();
if !messages.is_empty() {
let payload = serde_json::json!({
"flash": messages
});
if let Ok(v) = http::header::HeaderValue::from_str(&payload.to_string()) {
res.headers_mut()
.insert(http::header::HeaderName::from_static("hx-trigger"), v);
}
}
res
}
}
#[cfg(feature = "maud")]
impl Flash {
pub async fn render(&self) -> maud::Markup {
self.render_inner(false).await
}
pub async fn render_oob(&self) -> maud::Markup {
self.render_inner(true).await
}
async fn render_inner(&self, oob: bool) -> maud::Markup {
let messages = self.consume().await;
maud::html! {
div id="flash" class="flash-messages" role="status" aria-live="polite"
hx-swap-oob=[oob.then_some("true")] {
(flash_message_divs(&messages))
}
}
}
}
#[cfg(feature = "maud")]
#[must_use]
pub fn flash_message_divs(messages: &[FlashMessage]) -> maud::Markup {
maud::html! {
@for msg in messages {
div class={ "flash flash-" (msg.level.as_str()) } role="alert" {
(msg.message)
}
}
}
}
impl FlashLevel {
#[must_use]
pub const fn live_region(&self) -> (&'static str, &'static str) {
match self {
Self::Error | Self::Warning => ("alert", "assertive"),
Self::Success | Self::Info => ("status", "polite"),
}
}
}
#[cfg(feature = "maud")]
#[derive(Debug, Clone, Copy, Default)]
pub struct FlashMessagesConfig {
dismissible: bool,
}
#[cfg(feature = "maud")]
impl FlashMessagesConfig {
#[must_use]
pub const fn new() -> Self {
Self { dismissible: false }
}
#[must_use]
pub const fn dismissible(mut self, yes: bool) -> Self {
self.dismissible = yes;
self
}
}
#[cfg(feature = "maud")]
#[must_use]
pub fn flash_messages(messages: &[FlashMessage]) -> maud::Markup {
flash_messages_with(messages, &FlashMessagesConfig::new())
}
#[cfg(feature = "maud")]
#[must_use]
pub fn flash_messages_with(
messages: &[FlashMessage],
config: &FlashMessagesConfig,
) -> maud::Markup {
if messages.is_empty() {
return maud::html! {};
}
maud::html! {
div class="autumn-flash-group" {
@for msg in messages {
@let (role, live) = msg.level.live_region();
div class={ "autumn-flash autumn-flash--" (msg.level.as_str()) }
role=(role) aria-live=(live) {
span class="autumn-flash__body" { (msg.message) }
@if config.dismissible {
label class="autumn-flash__dismiss" {
input type="checkbox" class="autumn-flash__dismiss-toggle"
aria-label="Dismiss this message";
span aria-hidden="true" { "×" }
}
}
}
}
}
}
}
pub const FLASH_CSS_PATH: &str = "/static/css/autumn-flash.css";
pub const FLASH_CSS: &str = "\
.flash-messages:empty{display:none}\
.flash{padding:.75rem 1rem;border-radius:.375rem;margin-bottom:.5rem;border:1px solid}\
.flash-success{background:var(--flash-success-bg,#ecfdf5);color:var(--flash-success-fg,#065f46);border-color:var(--flash-success-border,#6ee7b7)}\
.flash-info{background:var(--flash-info-bg,#eff6ff);color:var(--flash-info-fg,#1e3a8a);border-color:var(--flash-info-border,#93c5fd)}\
.flash-warning{background:var(--flash-warning-bg,#fffbeb);color:var(--flash-warning-fg,#92400e);border-color:var(--flash-warning-border,#fcd34d)}\
.flash-error{background:var(--flash-error-bg,#fef2f2);color:var(--flash-error-fg,#991b1b);border-color:var(--flash-error-border,#fca5a5)}\
.autumn-flash-group{display:flex;flex-direction:column;gap:.5rem}\
.autumn-flash{display:flex;align-items:flex-start;justify-content:space-between;gap:.75rem;padding:.75rem 1rem;border-radius:.375rem;border:1px solid}\
.autumn-flash__body{flex:1 1 auto}\
.autumn-flash--success{background:var(--flash-success-bg,#ecfdf5);color:var(--flash-success-fg,#065f46);border-color:var(--flash-success-border,#6ee7b7)}\
.autumn-flash--info{background:var(--flash-info-bg,#eff6ff);color:var(--flash-info-fg,#1e3a8a);border-color:var(--flash-info-border,#93c5fd)}\
.autumn-flash--warning{background:var(--flash-warning-bg,#fffbeb);color:var(--flash-warning-fg,#92400e);border-color:var(--flash-warning-border,#fcd34d)}\
.autumn-flash--error{background:var(--flash-error-bg,#fef2f2);color:var(--flash-error-fg,#991b1b);border-color:var(--flash-error-border,#fca5a5)}\
.autumn-flash__dismiss{flex:0 0 auto;cursor:pointer;line-height:1;font-size:1.25rem;color:inherit;background:none;border:0;padding:0 .25rem}\
.autumn-flash__dismiss-toggle{position:absolute;width:1px;height:1px;padding:0;margin:-1px;overflow:hidden;clip-path:inset(50%);white-space:nowrap;border:0}\
.autumn-flash__dismiss:has(.autumn-flash__dismiss-toggle:focus-visible){outline:2px solid var(--primary,#7c3aed);outline-offset:2px}\
.autumn-flash:has(.autumn-flash__dismiss-toggle:checked){display:none}\
";
impl<S> FromRequestParts<S> for Flash
where
S: Send + Sync,
{
type Rejection = <Session as FromRequestParts<S>>::Rejection;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let session = Session::from_request_parts(parts, state).await?;
Ok(Self::new(session))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[tokio::test]
async fn flash_push_and_consume() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session.clone());
flash.success("Saved!").await;
flash.error("Failed!").await;
let messages = flash.peek().await;
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].level, FlashLevel::Success);
assert_eq!(messages[0].message, "Saved!");
assert_eq!(messages[1].level, FlashLevel::Error);
assert_eq!(messages[1].message, "Failed!");
assert_eq!(flash.peek().await.len(), 2);
let consumed = flash.consume().await;
assert_eq!(consumed.len(), 2);
assert_eq!(flash.peek().await.len(), 0);
}
#[tokio::test]
async fn flash_level_as_str() {
assert_eq!(FlashLevel::Success.as_str(), "success");
assert_eq!(FlashLevel::Info.as_str(), "info");
assert_eq!(FlashLevel::Warning.as_str(), "warning");
assert_eq!(FlashLevel::Error.as_str(), "error");
}
#[tokio::test]
async fn should_not_remove_key_when_consuming_empty_flash() -> Result<(), String> {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
session.insert("dummy", "val").await;
let flash = Flash::new(session.clone());
let messages = flash.consume().await;
assert_eq!(messages.len(), 0);
assert_eq!(
session.get("dummy").await.ok_or("missing key dummy")?,
"val"
);
assert!(!session.contains_key(FLASH_SESSION_KEY).await);
Ok(())
}
#[tokio::test]
async fn should_handle_invalid_json_gracefully() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
session
.insert(FLASH_SESSION_KEY, "{ invalid_json: true")
.await;
let flash = Flash::new(session);
let messages = flash.peek().await;
assert_eq!(messages.len(), 0);
}
#[tokio::test]
async fn should_support_all_convenience_methods() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session);
flash.success("Success msg").await;
flash.info("Info msg").await;
flash.warning("Warning msg").await;
flash.error("Error msg").await;
let messages = flash.peek().await;
assert_eq!(messages.len(), 4);
assert_eq!(messages[0].level, FlashLevel::Success);
assert_eq!(messages[0].message, "Success msg");
assert_eq!(messages[1].level, FlashLevel::Info);
assert_eq!(messages[1].message, "Info msg");
assert_eq!(messages[2].level, FlashLevel::Warning);
assert_eq!(messages[2].message, "Warning msg");
assert_eq!(messages[3].level, FlashLevel::Error);
assert_eq!(messages[3].message, "Error msg");
}
#[tokio::test]
#[cfg(feature = "maud")]
async fn render_emits_messages_and_clears_them() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session.clone());
flash.success("Saved!").await;
flash.error("Oops").await;
let markup = flash.render().await.into_string();
assert!(
markup.contains("id=\"flash\""),
"missing container: {markup}"
);
assert!(markup.contains("aria-live=\"polite\""));
assert!(markup.contains("flash flash-success"));
assert!(markup.contains("Saved!"));
assert!(markup.contains("flash flash-error"));
assert!(markup.contains("Oops"));
assert!(
!markup.contains("style="),
"must not emit inline styles: {markup}"
);
assert!(!markup.contains("hx-swap-oob"));
assert_eq!(flash.peek().await.len(), 0);
}
#[tokio::test]
#[cfg(feature = "maud")]
async fn render_emits_container_even_when_empty() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session);
let markup = flash.render().await.into_string();
assert!(
markup.contains("id=\"flash\""),
"missing container: {markup}"
);
assert!(!markup.contains("flash flash-"));
}
#[tokio::test]
#[cfg(feature = "maud")]
async fn render_oob_marks_container_for_out_of_band_swap() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session.clone());
flash.info("Updated").await;
let markup = flash.render_oob().await.into_string();
assert!(markup.contains("id=\"flash\""));
assert!(
markup.contains("hx-swap-oob=\"true\""),
"missing OOB attr: {markup}"
);
assert!(markup.contains("flash flash-info"));
assert!(markup.contains("Updated"));
assert_eq!(flash.peek().await.len(), 0);
}
#[test]
#[cfg(feature = "maud")]
fn flash_messages_empty_slice_renders_nothing() {
assert_eq!(flash_messages(&[]).into_string(), "");
assert_eq!(
flash_messages_with(&[], &FlashMessagesConfig::new().dismissible(true)).into_string(),
""
);
}
#[test]
#[cfg(feature = "maud")]
fn flash_messages_maps_severity_to_live_region() {
for (level, role, live) in [
(FlashLevel::Success, "status", "polite"),
(FlashLevel::Info, "status", "polite"),
(FlashLevel::Warning, "alert", "assertive"),
(FlashLevel::Error, "alert", "assertive"),
] {
let msg = [FlashMessage {
level,
message: "hi".into(),
}];
let html = flash_messages(&msg).into_string();
assert!(
html.contains(&format!(r#"role="{role}""#)),
"level {level:?} should carry role={role}: {html}"
);
assert!(
html.contains(&format!(r#"aria-live="{live}""#)),
"level {level:?} should carry aria-live={live}: {html}"
);
assert!(
html.contains(&format!("autumn-flash--{}", level.as_str())),
"level {level:?} should carry semantic class: {html}"
);
}
}
#[test]
#[cfg(feature = "maud")]
fn flash_messages_escapes_message_text() {
let msg = [FlashMessage {
level: FlashLevel::Error,
message: "<script>alert(1)</script>".into(),
}];
let html = flash_messages(&msg).into_string();
assert!(
!html.contains("<script>"),
"must escape message text: {html}"
);
assert!(html.contains("<script>"), "{html}");
}
#[test]
#[cfg(feature = "maud")]
fn flash_messages_emits_no_inline_style() {
let msg = [FlashMessage {
level: FlashLevel::Success,
message: "Saved!".into(),
}];
assert!(!flash_messages(&msg).into_string().contains("style="));
}
#[test]
#[cfg(feature = "maud")]
fn flash_messages_dismiss_control_is_opt_in_and_js_free() {
let msg = [FlashMessage {
level: FlashLevel::Info,
message: "Heads up".into(),
}];
assert!(
!flash_messages(&msg)
.into_string()
.contains("autumn-flash__dismiss")
);
let dismissible =
flash_messages_with(&msg, &FlashMessagesConfig::new().dismissible(true)).into_string();
assert!(
dismissible.contains("autumn-flash__dismiss"),
"{dismissible}"
);
assert!(dismissible.contains(r#"type="checkbox""#), "{dismissible}");
assert!(!dismissible.contains("onclick"), "{dismissible}");
assert!(!dismissible.contains("<script"), "{dismissible}");
assert!(
!dismissible.contains(r#"autumn-flash__dismiss-toggle" hidden"#),
"toggle must not be `hidden`: {dismissible}"
);
assert!(
dismissible.contains("autumn-flash__dismiss-toggle"),
"sr-only focusable toggle class present: {dismissible}"
);
assert!(
dismissible.contains(r#"aria-label="Dismiss this message""#),
"{dismissible}"
);
}
#[test]
#[cfg(feature = "maud")]
fn flash_css_backs_the_autumn_flash_classes() {
for selector in [
".autumn-flash",
".autumn-flash--success",
".autumn-flash--info",
".autumn-flash--warning",
".autumn-flash--error",
".autumn-flash__dismiss",
".autumn-flash__dismiss-toggle",
] {
assert!(FLASH_CSS.contains(selector), "FLASH_CSS missing {selector}");
}
assert!(
FLASH_CSS.contains("clip-path:inset(50%)"),
"FLASH_CSS must ship the sr-only (focusable) toggle rule"
);
assert!(
FLASH_CSS.contains(":focus-visible"),
"FLASH_CSS must ship a focus indicator"
);
}
#[tokio::test]
#[cfg(feature = "htmx")]
async fn should_inject_hx_trigger() {
let session = Session::new_for_test("test_id".to_string(), HashMap::new());
let flash = Flash::new(session.clone());
flash.success("Item saved").await;
let response = flash.inject_hx_trigger("OK").await;
let header = response.headers().get("hx-trigger");
assert!(header.is_some());
let json_str = header.unwrap().to_str().unwrap();
let payload: serde_json::Value = serde_json::from_str(json_str).unwrap();
assert_eq!(payload["flash"][0]["level"], "success");
assert_eq!(payload["flash"][0]["message"], "Item saved");
}
}