use axum::http::Request;
use axum::{
body::Body,
extract::FromRequestParts,
http::{HeaderMap, HeaderValue, StatusCode, header, request::Parts},
middleware::Next,
response::{IntoResponse, Redirect, Response},
};
use crate::components::nav_origin::{
arrived_from_dashboard, scope_from_dashboard, with_nav_origin,
};
use crate::components::swap::{AppLayoutKey, MainContentKey, SwapKey, oob_delete};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HtmxRequestType {
Full,
Partial,
}
impl HtmxRequestType {
fn parse(raw: &str) -> Option<Self> {
match raw.trim().to_ascii_lowercase().as_str() {
"full" => Some(Self::Full),
"partial" => Some(Self::Partial),
_ => None,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct Htmx {
pub request: bool,
pub boosted: bool,
pub history_restore: bool,
pub request_type: Option<HtmxRequestType>,
pub target_id: Option<String>,
pub source_id: Option<String>,
pub current_url: Option<String>,
}
impl Htmx {
pub fn from_headers(headers: &HeaderMap) -> Self {
let request = header_true(headers, "HX-Request");
let boosted = header_true(headers, "HX-Boosted");
let history_restore = header_true(headers, "HX-History-Restore-Request");
let request_type = headers
.get("HX-Request-Type")
.and_then(|v| v.to_str().ok())
.and_then(HtmxRequestType::parse);
let target_id = headers
.get("HX-Target")
.and_then(|v| v.to_str().ok())
.and_then(parse_element_id);
let source_id = headers
.get("HX-Source")
.and_then(|v| v.to_str().ok())
.and_then(parse_element_id)
.or_else(|| {
headers
.get("HX-Trigger")
.and_then(|v| v.to_str().ok())
.and_then(parse_element_id)
});
let current_url = headers
.get("HX-Current-URL")
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
Self {
request,
boosted,
history_restore,
request_type,
target_id,
source_id,
current_url,
}
}
pub fn wants_partial(&self) -> bool {
self.request && !self.history_restore
}
pub fn targets<K: SwapKey>(&self) -> bool {
self.target_id.as_deref() == Some(K::ID)
}
pub fn wants_app_layout(&self) -> bool {
self.wants_partial() && self.targets::<AppLayoutKey>()
}
pub fn wants_main_content(&self) -> bool {
if !self.wants_partial() || self.targets::<AppLayoutKey>() {
return false;
}
if matches!(self.target_id.as_deref(), Some("body")) {
return false;
}
self.targets::<MainContentKey>() || self.target_id.is_none()
}
pub fn redirect(&self, path: &str) -> Response {
let path = with_nav_origin(path);
if self.request {
let mut response = StatusCode::OK.into_response();
if let Ok(value) = HeaderValue::from_str(&path) {
response.headers_mut().insert("HX-Redirect", value);
}
response
} else {
Redirect::to(&path).into_response()
}
}
}
pub const TABLE_REFRESH_EVENT: &str = "lariv-table-refresh";
pub fn table_refresh_event(table_id: &str) -> String {
format!("{TABLE_REFRESH_EVENT}-{table_id}")
}
pub const FK_CREATED_EVENT: &str = "lariv-fk-created";
pub fn respond_create_modal_done<M: SwapKey>(
htmx: &Htmx,
refresh_table_id: &str,
detail_url: &str,
) -> Response {
respond_create_modal_done_fk::<M>(htmx, refresh_table_id, detail_url, "", "", "")
}
pub fn respond_create_modal_done_fk<M: SwapKey>(
htmx: &Htmx,
refresh_table_id: &str,
detail_url: &str,
fk_value: impl ToString,
fk_display: &str,
target_input: &str,
) -> Response {
respond_create_modal_done_fk_extra::<M>(
htmx,
refresh_table_id,
detail_url,
fk_value,
fk_display,
target_input,
&[],
)
}
pub fn respond_create_modal_done_fk_extra<M: SwapKey>(
htmx: &Htmx,
refresh_table_id: &str,
detail_url: &str,
fk_value: impl ToString,
fk_display: &str,
target_input: &str,
extra: &[(&str, &str)],
) -> Response {
let refresh = refresh_table_id.trim();
let target = target_input.trim();
if !htmx.request || (refresh.is_empty() && target.is_empty()) {
return htmx.redirect(detail_url);
}
let body = oob_delete::<M>().into_string();
let fk_value = fk_value.to_string();
let mut trigger_map = serde_json::Map::new();
if !refresh.is_empty() && target.is_empty() {
trigger_map.insert(
table_refresh_event(refresh),
serde_json::json!({ "target": "document" }),
);
}
if !fk_value.is_empty() {
let mut fk = serde_json::Map::new();
fk.insert("value".into(), serde_json::Value::String(fk_value));
fk.insert(
"display".into(),
serde_json::Value::String(fk_display.to_string()),
);
if !target.is_empty() {
fk.insert("name".into(), serde_json::Value::String(target.to_string()));
}
for (k, v) in extra {
if !k.is_empty() {
fk.insert(
(*k).to_string(),
serde_json::Value::String((*v).to_string()),
);
}
}
trigger_map.insert(FK_CREATED_EVENT.to_string(), serde_json::Value::Object(fk));
}
let trigger = serde_json::Value::Object(trigger_map).to_string();
let mut builder = Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "text/html; charset=utf-8")
.header("HX-Reswap", "none");
if let Some(value) = hx_trigger_header_value(&trigger) {
builder = builder.header("HX-Trigger", value);
}
builder
.body(body.into())
.unwrap_or_else(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())
}
pub fn respond_edit_modal_done<M: SwapKey>(htmx: &Htmx, detail_url: &str) -> Response {
if !htmx.request {
return htmx.redirect(detail_url);
}
let body = oob_delete::<M>().into_string();
let mut builder = Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "text/html; charset=utf-8")
.header("HX-Reswap", "none");
if let Ok(value) = HeaderValue::from_str(detail_url) {
builder = builder.header("HX-Redirect", value);
}
builder
.body(body.into())
.unwrap_or_else(|_| StatusCode::INTERNAL_SERVER_ERROR.into_response())
}
impl<S> FromRequestParts<S> for Htmx
where
S: Send + Sync,
{
type Rejection = std::convert::Infallible;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
if let Some(htmx) = parts.extensions.get::<Htmx>() {
return Ok(htmx.clone());
}
Ok(Htmx::from_headers(&parts.headers))
}
}
pub fn parse_element_id(raw: &str) -> Option<String> {
let raw = raw.trim();
if raw.is_empty() {
return None;
}
if let Some((_, id)) = raw.rsplit_once('#') {
let id = id.trim();
if id.is_empty() {
None
} else {
Some(id.to_owned())
}
} else if let Some(id) = raw.strip_prefix('#') {
let id = id.trim();
if id.is_empty() {
None
} else {
Some(id.to_owned())
}
} else {
Some(raw.to_owned())
}
}
fn hx_trigger_header_value(trigger: &str) -> Option<HeaderValue> {
HeaderValue::from_str(&escape_non_ascii_json(trigger)).ok()
}
fn escape_non_ascii_json(s: &str) -> String {
use std::fmt::Write;
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if c.is_ascii() && !c.is_control() {
out.push(c);
} else {
if let Err(e) = write!(out, "\\u{:04x}", c as u32) {
tracing::error!(error = %e, "failed escaping non-ascii for hx-trigger");
}
}
}
out
}
fn header_true(headers: &HeaderMap, name: &str) -> bool {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.eq_ignore_ascii_case("true"))
}
fn rewrite_redirect_origin(headers: &mut HeaderMap) {
let Some(location) = headers
.get(header::LOCATION)
.and_then(|v| v.to_str().ok())
.map(str::to_owned)
else {
return;
};
let rewritten = with_nav_origin(&location);
if rewritten == location {
return;
}
if let Ok(value) = HeaderValue::from_str(&rewritten) {
headers.insert(header::LOCATION, value);
}
}
const VARY_HTMX: &str = "HX-Request, HX-Target, HX-Request-Type, HX-History-Restore-Request";
pub async fn htmx_middleware(mut req: Request<Body>, next: Next) -> Response {
let htmx = Htmx::from_headers(req.headers());
let is_htmx = htmx.request;
let from_dashboard = arrived_from_dashboard(req.uri(), req.headers());
req.extensions_mut().insert(htmx);
let mut response = scope_from_dashboard(from_dashboard, async move {
let mut response = next.run(req).await;
if from_dashboard {
rewrite_redirect_origin(response.headers_mut());
}
response
})
.await;
if !is_htmx {
return response;
}
response
.headers_mut()
.insert(header::VARY, HeaderValue::from_static(VARY_HTMX));
let status = response.status();
if status.is_redirection()
&& let Some(location) = response.headers().get(header::LOCATION).cloned()
{
let mut headers = response.headers().clone();
headers.remove(header::LOCATION);
headers.insert("HX-Redirect", location);
headers.insert(header::VARY, HeaderValue::from_static(VARY_HTMX));
let body = response.into_body();
let mut rebuilt = Response::new(body);
*rebuilt.status_mut() = StatusCode::OK;
*rebuilt.headers_mut() = headers;
return rebuilt;
}
response
}
#[cfg(test)]
mod tests {
use super::*;
use crate::swap_key;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use axum::middleware::from_fn;
use axum::routing::get;
use axum::{Router, response::Redirect};
use tower::ServiceExt;
swap_key!(TestPaneKey, "app-layout");
swap_key!(TestTableKey, "user-table");
#[test]
fn parse_element_id_formats() {
assert_eq!(
parse_element_id("div#app-layout").as_deref(),
Some("app-layout")
);
assert_eq!(
parse_element_id("#app-layout").as_deref(),
Some("app-layout")
);
assert_eq!(
parse_element_id("app-layout").as_deref(),
Some("app-layout")
);
assert_eq!(parse_element_id("div#").as_deref(), None);
assert_eq!(parse_element_id("").as_deref(), None);
}
#[test]
fn history_restore_requests_full_page_not_pane() {
let mut headers = HeaderMap::new();
headers.insert("HX-Request", HeaderValue::from_static("true"));
headers.insert(
"HX-History-Restore-Request",
HeaderValue::from_static("true"),
);
headers.insert("HX-Target", HeaderValue::from_static("div#app-layout"));
let htmx = Htmx::from_headers(&headers);
assert!(htmx.history_restore);
assert!(!htmx.wants_partial());
assert!(!htmx.wants_app_layout());
assert!(!htmx.wants_main_content());
}
#[test]
fn targets_and_wants_app_layout() {
let mut headers = HeaderMap::new();
headers.insert("HX-Request", HeaderValue::from_static("true"));
headers.insert("HX-Target", HeaderValue::from_static("div#app-layout"));
let htmx = Htmx::from_headers(&headers);
assert!(htmx.request);
assert!(htmx.targets::<TestPaneKey>());
assert!(htmx.wants_app_layout());
assert!(!htmx.wants_main_content());
assert!(!htmx.targets::<TestTableKey>());
}
#[test]
fn wants_main_content_htmx4_target_and_missing_target() {
let mut headers = HeaderMap::new();
headers.insert("HX-Request", HeaderValue::from_static("true"));
headers.insert("HX-Target", HeaderValue::from_static("main#main-content"));
let htmx = Htmx::from_headers(&headers);
assert!(htmx.wants_main_content());
assert!(!htmx.wants_app_layout());
let mut headers = HeaderMap::new();
headers.insert("HX-Request", HeaderValue::from_static("true"));
let htmx = Htmx::from_headers(&headers);
assert!(htmx.wants_main_content());
assert!(!htmx.wants_app_layout());
let mut headers = HeaderMap::new();
headers.insert("HX-Request", HeaderValue::from_static("true"));
headers.insert("HX-Target", HeaderValue::from_static("div#user-table"));
let htmx = Htmx::from_headers(&headers);
assert!(!htmx.wants_main_content());
}
#[test]
fn redirect_htmx_uses_200_hx_redirect() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = htmx.redirect("/users/");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("HX-Redirect")
.and_then(|v| v.to_str().ok()),
Some("/users/")
);
assert!(response.headers().get(header::LOCATION).is_none());
}
#[test]
fn redirect_non_htmx_uses_303() {
let htmx = Htmx::default();
let response = htmx.redirect("/users/");
assert_eq!(response.status(), StatusCode::SEE_OTHER);
assert_eq!(
response
.headers()
.get(header::LOCATION)
.and_then(|v| v.to_str().ok()),
Some("/users/")
);
}
swap_key!(TestCreateModalKey, "role-create-modal");
#[test]
fn create_modal_done_refreshes_table_when_refresh_set() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response =
respond_create_modal_done::<TestCreateModalKey>(&htmx, "role-table", "/users/roles/1/");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("HX-Reswap")
.and_then(|v| v.to_str().ok()),
Some("none")
);
let trigger = response
.headers()
.get("HX-Trigger")
.and_then(|v| v.to_str().ok())
.unwrap_or_default();
assert!(
trigger.contains(&table_refresh_event("role-table")),
"{trigger}"
);
assert!(trigger.contains("\"target\":\"document\""), "{trigger}");
assert!(!trigger.contains(FK_CREATED_EVENT), "{trigger}");
assert!(response.headers().get("HX-Redirect").is_none());
}
#[test]
fn create_modal_done_fk_triggers_picker_fill() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = respond_create_modal_done_fk::<TestCreateModalKey>(
&htmx,
"role-table",
"/users/roles/1/",
1,
"Admin",
"",
);
let trigger = response
.headers()
.get("HX-Trigger")
.and_then(|v| v.to_str().ok())
.unwrap_or_default();
assert!(
trigger.contains(&table_refresh_event("role-table")),
"{trigger}"
);
assert!(trigger.contains(FK_CREATED_EVENT), "{trigger}");
assert!(trigger.contains("\"value\":\"1\""), "{trigger}");
assert!(trigger.contains("\"display\":\"Admin\""), "{trigger}");
}
#[test]
fn create_modal_done_fk_keeps_table_refresh_with_unicode_display() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = respond_create_modal_done_fk::<TestCreateModalKey>(
&htmx,
"role-table",
"/users/roles/1/",
1,
"Café",
"",
);
let trigger = response
.headers()
.get("HX-Trigger")
.map(|v| String::from_utf8_lossy(v.as_bytes()).into_owned())
.unwrap_or_default();
assert!(
trigger.contains(&table_refresh_event("role-table")),
"{trigger}"
);
assert!(trigger.contains("\"target\":\"document\""), "{trigger}");
assert!(
trigger.contains("Caf\\u00e9") || trigger.contains("Café"),
"{trigger}"
);
}
#[test]
fn create_modal_done_fk_fills_field_without_table_refresh() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = respond_create_modal_done_fk::<TestCreateModalKey>(
&htmx,
"role-selection-table",
"/users/roles/1/",
1,
"Admin",
"RoleID",
);
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().get("HX-Redirect").is_none());
let trigger = response
.headers()
.get("HX-Trigger")
.and_then(|v| v.to_str().ok())
.unwrap_or_default();
assert!(
!trigger.contains(&table_refresh_event("role-selection-table")),
"{trigger}"
);
assert!(trigger.contains(FK_CREATED_EVENT), "{trigger}");
assert!(trigger.contains("\"name\":\"RoleID\""), "{trigger}");
assert!(trigger.contains("\"value\":\"1\""), "{trigger}");
assert!(trigger.contains("\"display\":\"Admin\""), "{trigger}");
}
#[test]
fn create_modal_done_fk_stays_on_page_with_target_input_only() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = respond_create_modal_done_fk::<TestCreateModalKey>(
&htmx,
"",
"/users/roles/1/",
1,
"Admin",
"RoleID",
);
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().get("HX-Redirect").is_none());
let trigger = response
.headers()
.get("HX-Trigger")
.and_then(|v| v.to_str().ok())
.unwrap_or_default();
assert!(trigger.contains(FK_CREATED_EVENT), "{trigger}");
assert!(trigger.contains("\"name\":\"RoleID\""), "{trigger}");
}
#[test]
fn create_modal_done_redirects_without_refresh() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response =
respond_create_modal_done::<TestCreateModalKey>(&htmx, "", "/users/roles/1/");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("HX-Redirect")
.and_then(|v| v.to_str().ok()),
Some("/users/roles/1/")
);
}
#[test]
fn edit_modal_done_closes_modal_and_redirects() {
let htmx = Htmx {
request: true,
..Default::default()
};
let response = respond_edit_modal_done::<TestCreateModalKey>(&htmx, "/crm/leads/1/");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("HX-Reswap")
.and_then(|v| v.to_str().ok()),
Some("none")
);
assert_eq!(
response
.headers()
.get("HX-Redirect")
.and_then(|v| v.to_str().ok()),
Some("/crm/leads/1/")
);
}
#[tokio::test]
async fn middleware_rewrites_3xx_to_hx_redirect() {
let app = Router::new()
.route("/go", get(|| async { Redirect::to("/users/login") }))
.layer(from_fn(htmx_middleware));
let response = app
.oneshot(
Request::builder()
.uri("/go")
.header("HX-Request", "true")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("HX-Redirect")
.and_then(|v| v.to_str().ok()),
Some("/users/login")
);
assert!(response.headers().get(header::LOCATION).is_none());
assert_eq!(
response
.headers()
.get(header::VARY)
.and_then(|v| v.to_str().ok()),
Some(VARY_HTMX)
);
}
#[tokio::test]
async fn middleware_leaves_non_htmx_redirect() {
let app = Router::new()
.route("/go", get(|| async { Redirect::to("/users/login") }))
.layer(from_fn(htmx_middleware));
let response = app
.oneshot(Request::builder().uri("/go").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::SEE_OTHER);
assert_eq!(
response
.headers()
.get(header::LOCATION)
.and_then(|v| v.to_str().ok()),
Some("/users/login")
);
}
#[tokio::test]
async fn middleware_scopes_from_dashboard_query() {
let app = Router::new()
.route(
"/users/",
get(|| async { crate::components::nav_origin::from_dashboard().to_string() }),
)
.layer(from_fn(htmx_middleware));
let with_origin = app
.clone()
.oneshot(
Request::builder()
.uri("/users/?from=dashboard")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let body = axum::body::to_bytes(with_origin.into_body(), 64)
.await
.unwrap();
#[cfg(feature = "plugin-dashboard")]
assert_eq!(&body[..], b"true");
#[cfg(not(feature = "plugin-dashboard"))]
assert_eq!(&body[..], b"false");
let without_origin = app
.oneshot(
Request::builder()
.uri("/users/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let body = axum::body::to_bytes(without_origin.into_body(), 64)
.await
.unwrap();
assert_eq!(&body[..], b"false");
}
#[tokio::test]
async fn middleware_rewrites_redirects_to_keep_dashboard_origin() {
let app = Router::new()
.route("/go", get(|| async { Redirect::to("/crm/contacts") }))
.layer(from_fn(htmx_middleware));
let response = app
.oneshot(
Request::builder()
.uri("/go?from=dashboard")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
#[cfg(feature = "plugin-dashboard")]
assert_eq!(
response
.headers()
.get(header::LOCATION)
.and_then(|v| v.to_str().ok()),
Some("/crm/contacts?from=dashboard")
);
#[cfg(not(feature = "plugin-dashboard"))]
assert_eq!(
response
.headers()
.get(header::LOCATION)
.and_then(|v| v.to_str().ok()),
Some("/crm/contacts")
);
}
}