1use std::convert::Infallible;
2
3use axum::extract::FromRequestParts;
4use axum::http::header::REFERER;
5use axum::http::request::Parts;
6use axum::http::{HeaderMap, HeaderName, HeaderValue, StatusCode};
7use axum::response::{IntoResponse, IntoResponseParts, Redirect, Response, ResponseParts};
8
9#[derive(Debug, Clone, Default)]
19#[non_exhaustive]
20pub struct Htmx {
21 pub request: bool,
23 pub boosted: bool,
25 pub target: Option<String>,
27 pub trigger: Option<String>,
29 pub current_url: Option<String>,
31}
32
33impl Htmx {
34 pub fn from_headers(headers: &HeaderMap) -> Self {
36 let text = |name: &str| {
37 headers
38 .get(name)
39 .and_then(|v| v.to_str().ok())
40 .map(str::to_owned)
41 };
42 Self {
43 request: text("hx-request").as_deref() == Some("true"),
44 boosted: text("hx-boosted").as_deref() == Some("true"),
45 target: text("hx-target"),
46 trigger: text("hx-trigger"),
47 current_url: text("hx-current-url"),
48 }
49 }
50
51 pub fn redirect(&self, to: &str) -> axum::response::Response {
54 use axum::response::IntoResponse;
55 if self.request {
56 HxRedirect(to.to_owned()).into_response()
57 } else {
58 axum::response::Redirect::to(to).into_response()
59 }
60 }
61
62 pub fn wants_fragment(&self) -> bool {
64 self.request && !self.boosted
65 }
66}
67
68impl<S: Send + Sync> FromRequestParts<S> for Htmx {
69 type Rejection = Infallible;
70
71 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
72 Ok(Self::from_headers(&parts.headers))
73 }
74}
75
76fn set(res: &mut ResponseParts, name: &'static str, value: &str) {
77 match HeaderValue::from_str(value) {
78 Ok(value) => {
79 res.headers_mut()
80 .insert(HeaderName::from_static(name), value);
81 }
82 Err(_) => tracing::warn!(header = name, "invalid header value dropped"),
83 }
84}
85
86pub struct HxRedirect(pub String);
88
89impl IntoResponseParts for HxRedirect {
90 type Error = Infallible;
91
92 fn into_response_parts(self, mut res: ResponseParts) -> Result<ResponseParts, Infallible> {
93 set(&mut res, "hx-redirect", &self.0);
94 Ok(res)
95 }
96}
97
98impl IntoResponse for HxRedirect {
99 fn into_response(self) -> Response {
100 (self, StatusCode::OK).into_response()
101 }
102}
103
104pub struct HxRefresh;
106
107impl IntoResponseParts for HxRefresh {
108 type Error = Infallible;
109
110 fn into_response_parts(self, mut res: ResponseParts) -> Result<ResponseParts, Infallible> {
111 set(&mut res, "hx-refresh", "true");
112 Ok(res)
113 }
114}
115
116impl IntoResponse for HxRefresh {
117 fn into_response(self) -> Response {
118 (self, StatusCode::OK).into_response()
119 }
120}
121
122pub struct HxTrigger(pub String);
125
126impl IntoResponseParts for HxTrigger {
127 type Error = Infallible;
128
129 fn into_response_parts(self, mut res: ResponseParts) -> Result<ResponseParts, Infallible> {
130 set(&mut res, "hx-trigger", &self.0);
131 Ok(res)
132 }
133}
134
135pub struct HxRetarget(pub String);
138
139pub struct HxReswap(pub String);
142
143pub struct HxPushUrl(pub String);
146
147macro_rules! hx_header {
148 ($type:ty, $header:literal) => {
149 impl IntoResponseParts for $type {
150 type Error = Infallible;
151
152 fn into_response_parts(
153 self,
154 mut res: ResponseParts,
155 ) -> Result<ResponseParts, Infallible> {
156 set(&mut res, $header, &self.0);
157 Ok(res)
158 }
159 }
160 };
161}
162
163hx_header!(HxRetarget, "hx-retarget");
164hx_header!(HxReswap, "hx-reswap");
165hx_header!(HxPushUrl, "hx-push-url");
166
167pub(crate) fn add_trigger<B>(
170 res: &mut axum::http::Response<B>,
171 name: &str,
172 detail: serde_json::Value,
173) {
174 use serde_json::{Map, Value};
175 let mut triggers = match res
176 .headers()
177 .get("hx-trigger")
178 .and_then(|v| v.to_str().ok())
179 {
180 Some(existing) if existing.trim_start().starts_with('{') => {
181 serde_json::from_str::<Map<String, Value>>(existing).unwrap_or_default()
182 }
183 Some(existing) => existing
184 .split(',')
185 .map(str::trim)
186 .filter(|name| !name.is_empty())
187 .map(|name| (name.to_owned(), Value::Null))
188 .collect(),
189 None => Map::new(),
190 };
191 triggers.insert(name.to_owned(), detail);
192 let json = ascii_json(&Value::Object(triggers).to_string());
195 match axum::http::HeaderValue::from_str(&json) {
196 Ok(value) => {
197 res.headers_mut().insert("hx-trigger", value);
198 }
199 Err(err) => tracing::warn!(error = %err, "could not send an HX-Trigger header"),
200 }
201}
202
203fn ascii_json(json: &str) -> String {
207 let mut out = String::with_capacity(json.len());
208 for c in json.chars() {
209 if c.is_ascii() {
210 out.push(c);
211 } else {
212 let mut units = [0u16; 2];
213 for unit in c.encode_utf16(&mut units) {
214 out.push_str(&format!("\\u{unit:04x}"));
215 }
216 }
217 }
218 out
219}
220
221pub struct Back(Option<String>);
233
234impl<S: Send + Sync> FromRequestParts<S> for Back {
235 type Rejection = Infallible;
236
237 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
238 Ok(Self(same_site_referer(&parts.headers)))
239 }
240}
241
242pub(crate) fn same_site_referer(headers: &axum::http::HeaderMap) -> Option<String> {
246 let referer = headers.get(REFERER)?.to_str().ok()?;
247 if referer.starts_with('/') {
248 return is_local_path(referer).then(|| referer.to_owned());
249 }
250 let host = headers.get(axum::http::header::HOST)?.to_str().ok()?;
251 let rest = referer
252 .strip_prefix("https://")
253 .or_else(|| referer.strip_prefix("http://"))?;
254 let (authority, path) = match rest.find('/') {
255 Some(i) => (&rest[..i], &rest[i..]),
256 None => (rest, "/"),
257 };
258 (authority.eq_ignore_ascii_case(host) && is_local_path(path)).then(|| path.to_owned())
259}
260
261pub(crate) fn is_local_path(path: &str) -> bool {
264 path.starts_with('/')
265 && !path.starts_with("//")
266 && !path.starts_with("/\\")
267 && !path.contains(['\\', '\r', '\n'])
268}
269
270impl IntoResponse for Back {
271 fn into_response(self) -> Response {
272 Redirect::to(self.0.as_deref().unwrap_or("/")).into_response()
273 }
274}
275
276#[cfg(test)]
277mod tests {
278 use super::*;
279
280 #[test]
281 fn triggers_carry_any_text_as_ascii_json() {
282 let mut res = axum::http::Response::new(());
283 res.headers_mut().insert(
284 "hx-trigger",
285 axum::http::HeaderValue::from_static("task-added"),
286 );
287 let text = "“Coffee” added to the cart ✓ 🎉";
288 add_trigger(
289 &mut res,
290 "renox:toast",
291 serde_json::json!({ "message": text }),
292 );
293 let header = res.headers()["hx-trigger"].to_str().unwrap();
294 assert!(header.is_ascii(), "{header}");
295 let back: serde_json::Value = serde_json::from_str(header).unwrap();
296 assert_eq!(back["renox:toast"]["message"], text);
297 assert!(back.get("task-added").is_some(), "earlier events are kept");
298 }
299}