Skip to main content

datastar/
warp.rs

1//! Warp integration for Datastar.
2
3use {
4    crate::{
5        consts::{self, DATASTAR_REQ_HEADER_STR},
6        prelude::{DatastarEvent, ExecuteScript, PatchElements, PatchSignals},
7    },
8    bytes::Bytes,
9    serde::{Deserialize, de::DeserializeOwned},
10    std::{convert::Infallible, fmt::Write},
11    warp::{
12        Filter, Rejection, Reply,
13        filters::sse::Event,
14        http::{Method, StatusCode},
15    },
16};
17
18impl PatchElements {
19    /// Write this [`PatchElements`] into a Warp SSE [`Event`].
20    pub fn write_as_warp_sse_event(&self) -> Event {
21        self.as_datastar_event().write_as_warp_sse_event()
22    }
23}
24
25impl From<PatchElements> for Event {
26    fn from(value: PatchElements) -> Self {
27        value.write_as_warp_sse_event()
28    }
29}
30
31impl From<&PatchElements> for Event {
32    fn from(value: &PatchElements) -> Self {
33        value.write_as_warp_sse_event()
34    }
35}
36
37impl PatchSignals {
38    /// Write this [`PatchSignals`] into a Warp SSE [`Event`].
39    pub fn write_as_warp_sse_event(&self) -> Event {
40        self.as_datastar_event().write_as_warp_sse_event()
41    }
42}
43
44impl From<PatchSignals> for Event {
45    fn from(value: PatchSignals) -> Self {
46        value.write_as_warp_sse_event()
47    }
48}
49
50impl From<&PatchSignals> for Event {
51    fn from(value: &PatchSignals) -> Self {
52        value.write_as_warp_sse_event()
53    }
54}
55
56impl ExecuteScript {
57    /// Write this [`ExecuteScript`] into a Warp SSE [`Event`].
58    pub fn write_as_warp_sse_event(&self) -> Event {
59        self.as_datastar_event().write_as_warp_sse_event()
60    }
61}
62
63impl From<ExecuteScript> for Event {
64    fn from(value: ExecuteScript) -> Self {
65        value.write_as_warp_sse_event()
66    }
67}
68
69impl From<&ExecuteScript> for Event {
70    fn from(value: &ExecuteScript) -> Self {
71        value.write_as_warp_sse_event()
72    }
73}
74
75impl DatastarEvent {
76    /// Turn this [`DatastarEvent`] into a Warp SSE [`Event`].
77    pub fn write_as_warp_sse_event(&self) -> Event {
78        let mut event = Event::default().event(self.event.as_str());
79
80        if self.retry.as_millis() != (consts::DEFAULT_SSE_RETRY_DURATION as u128) {
81            event = event.retry(self.retry);
82        }
83
84        event = match self.id.as_deref() {
85            Some(id) => event.id(id),
86            None => event,
87        };
88
89        let mut data = String::with_capacity(
90            (self.data.iter().map(|s| s.len()).sum::<usize>() + self.data.len()).saturating_sub(1),
91        );
92
93        let mut sep = "";
94        for line in self.data.iter() {
95            // Assumption: std::fmt::write does not fail ever for [`String`].
96            let _ = write!(&mut data, "{sep}{line}");
97            sep = "\n";
98        }
99
100        event.data(data)
101    }
102}
103
104impl From<DatastarEvent> for Event {
105    fn from(value: DatastarEvent) -> Self {
106        value.write_as_warp_sse_event()
107    }
108}
109
110impl From<&DatastarEvent> for Event {
111    fn from(value: &DatastarEvent) -> Self {
112        value.write_as_warp_sse_event()
113    }
114}
115
116#[derive(Deserialize)]
117struct DatastarParam {
118    datastar: Option<serde_json::Value>,
119}
120
121/// Error type for [`ReadSignals`] extraction failures.
122#[derive(Debug)]
123pub struct ReadSignalsError {
124    message: String,
125    status: StatusCode,
126}
127
128impl warp::reject::Reject for ReadSignalsError {}
129
130/// [`ReadSignals`] is a wrapper type for extracted Datastar signals.
131///
132/// # Examples
133///
134/// ```
135/// use datastar::warp::{read_signals, ReadSignals};
136/// use serde::Deserialize;
137/// use warp::Filter;
138///
139/// #[derive(Deserialize)]
140/// struct Signals {
141///     foo: String,
142///     bar: i32,
143/// }
144///
145/// let route = warp::path("hello")
146///     .and(read_signals::<Signals>())
147///     .map(|signals: ReadSignals<Signals>| {
148///         format!("foo: {}, bar: {}", signals.0.foo, signals.0.bar)
149///     });
150/// ```
151#[derive(Debug)]
152pub struct ReadSignals<T>(pub T);
153
154/// Creates a Warp Filter that extracts Datastar signals from the request.
155///
156/// For GET and DELETE requests, signals are extracted from the `datastar` query
157/// parameter. A missing parameter is treated as JSON `null`, allowing
158/// `ReadSignals<Option<T>>` to produce `None`. For POST, PUT, and PATCH
159/// requests, signals are extracted from the JSON body.
160///
161/// # Examples
162///
163/// ```
164/// use datastar::warp::{read_signals, ReadSignals};
165/// use serde::Deserialize;
166/// use warp::Filter;
167///
168/// #[derive(Deserialize)]
169/// struct Signals {
170///     delay: u64,
171/// }
172///
173/// let route = warp::path("hello")
174///     .and(warp::get())
175///     .and(read_signals::<Signals>())
176///     .map(|ReadSignals(signals): ReadSignals<Signals>| {
177///         format!("delay: {}", signals.delay)
178///     });
179/// ```
180pub fn read_signals<T>() -> impl Filter<Extract = (ReadSignals<T>,), Error = Rejection> + Clone
181where
182    T: DeserializeOwned + Send,
183{
184    warp::method()
185        .and(warp::query::raw().or(warp::any().map(String::new)).unify())
186        .and(warp::body::bytes().or(warp::any().map(Bytes::new)).unify())
187        .and_then(extract_signals::<T>)
188}
189
190async fn extract_signals<T>(
191    method: Method,
192    query: String,
193    body: Bytes,
194) -> Result<ReadSignals<T>, Rejection>
195where
196    T: DeserializeOwned,
197{
198    match method {
199        Method::GET | Method::DELETE => {
200            // Parse ?datastar={json} from query string
201            let params: DatastarParam = serde_urlencoded::from_str(&query).map_err(|err| {
202                #[cfg(feature = "tracing")]
203                tracing::debug!(%err, "failed to parse query string");
204
205                warp::reject::custom(ReadSignalsError {
206                    message: format!("Failed to parse query: {err}"),
207                    status: StatusCode::BAD_REQUEST,
208                })
209            })?;
210
211            let signals_str = match params.datastar.as_ref() {
212                Some(value) => value.as_str().ok_or_else(|| {
213                    warp::reject::custom(ReadSignalsError {
214                        message: "datastar parameter must be a JSON string".into(),
215                        status: StatusCode::BAD_REQUEST,
216                    })
217                })?,
218                None => "null",
219            };
220
221            let signals: T = serde_json::from_str(signals_str).map_err(|err| {
222                #[cfg(feature = "tracing")]
223                tracing::debug!(%err, "failed to parse JSON value from query");
224
225                let _ = &err; // silence unused warning when tracing is disabled
226
227                warp::reject::custom(ReadSignalsError {
228                    message: format!("Failed to parse JSON: {err}"),
229                    status: StatusCode::BAD_REQUEST,
230                })
231            })?;
232
233            Ok(ReadSignals(signals))
234        }
235        _ => {
236            // POST/PUT/PATCH: parse body as JSON
237            let signals: T = serde_json::from_slice(&body).map_err(|err| {
238                #[cfg(feature = "tracing")]
239                tracing::debug!(%err, "failed to parse JSON value from body");
240
241                let _ = &err; // silence unused warning when tracing is disabled
242
243                warp::reject::custom(ReadSignalsError {
244                    message: format!("Failed to parse JSON body: {err}"),
245                    status: StatusCode::BAD_REQUEST,
246                })
247            })?;
248
249            Ok(ReadSignals(signals))
250        }
251    }
252}
253
254/// Creates a Filter that checks for the datastar-request header.
255/// Returns `true` if the header is present, `false` otherwise.
256pub fn is_datastar_request() -> impl Filter<Extract = (bool,), Error = Rejection> + Clone {
257    warp::header::optional::<String>(DATASTAR_REQ_HEADER_STR)
258        .map(|header: Option<String>| header.is_some())
259}
260
261/// Creates a Filter that optionally extracts Datastar signals from the request.
262///
263/// Returns `Some(ReadSignals<T>)` if signals are present and parseable,
264/// `None` if the `datastar-request` header is not present.
265///
266/// # Examples
267///
268/// ```
269/// use datastar::warp::{read_signals_optional, ReadSignals};
270/// use serde::Deserialize;
271/// use warp::Filter;
272///
273/// #[derive(Deserialize)]
274/// struct Signals {
275///     delay: u64,
276/// }
277///
278/// let route = warp::path("hello")
279///     .and(read_signals_optional::<Signals>())
280///     .map(|signals: Option<ReadSignals<Signals>>| {
281///         match signals {
282///             Some(ReadSignals(s)) => format!("delay: {}", s.delay),
283///             None => "no signals".to_string(),
284///         }
285///     });
286/// ```
287pub fn read_signals_optional<T>()
288-> impl Filter<Extract = (Option<ReadSignals<T>>,), Error = Rejection> + Clone
289where
290    T: DeserializeOwned + Send,
291{
292    warp::header::optional::<String>(DATASTAR_REQ_HEADER_STR)
293        .and(
294            read_signals::<T>()
295                .map(Some)
296                .or(warp::any().map(|| None::<ReadSignals<T>>))
297                .unify(),
298        )
299        .map(
300            |is_datastar: Option<String>, signals: Option<ReadSignals<T>>| {
301                if is_datastar.is_some() { signals } else { None }
302            },
303        )
304}
305
306/// Rejection handler for [`ReadSignals`] errors.
307///
308/// Use this with `warp::Filter::recover` to convert rejections into proper HTTP responses.
309///
310/// # Examples
311///
312/// ```
313/// use datastar::warp::{read_signals, handle_rejection, ReadSignals};
314/// use serde::Deserialize;
315/// use warp::Filter;
316///
317/// #[derive(Deserialize)]
318/// struct Signals {
319///     delay: u64,
320/// }
321///
322/// let route = warp::path("hello")
323///     .and(read_signals::<Signals>())
324///     .map(|ReadSignals(signals): ReadSignals<Signals>| {
325///         format!("delay: {}", signals.delay)
326///     })
327///     .recover(handle_rejection);
328/// ```
329pub async fn handle_rejection(err: Rejection) -> Result<impl Reply, Infallible> {
330    if let Some(e) = err.find::<ReadSignalsError>() {
331        Ok(warp::reply::with_status(e.message.clone(), e.status))
332    } else {
333        Ok(warp::reply::with_status(
334            "Internal Server Error".to_owned(),
335            StatusCode::INTERNAL_SERVER_ERROR,
336        ))
337    }
338}
339
340#[cfg(test)]
341mod tests {
342    use {super::*, crate::consts::ElementPatchMode, core::time::Duration, serde::Deserialize};
343
344    fn assert_event(event: Event, expected: &str) {
345        assert_eq!(event.to_string(), expected);
346    }
347
348    #[test]
349    fn writes_patch_elements_and_conversions() {
350        let patch = PatchElements::new("<div>one</div>\n<div>two</div>")
351            .id("elements-1")
352            .retry(Duration::from_millis(2_500))
353            .selector("#main")
354            .mode(ElementPatchMode::Append);
355        let expected = concat!(
356            "event:datastar-patch-elements\n",
357            "data:selector #main\n",
358            "data:mode append\n",
359            "data:elements <div>one</div>\n",
360            "data:elements <div>two</div>\n",
361            "id:elements-1\n",
362            "retry:2500\n\n",
363        );
364
365        assert_event(patch.write_as_warp_sse_event(), expected);
366        assert_event(Event::from(&patch), expected);
367        assert_event(Event::from(patch), expected);
368    }
369
370    #[test]
371    fn writes_patch_signals_and_conversions() {
372        let patch = PatchSignals::new("{count: 1}").only_if_missing(true);
373        let expected = concat!(
374            "event:datastar-patch-signals\n",
375            "data:onlyIfMissing true\n",
376            "data:signals {count: 1}\n\n",
377        );
378
379        assert_event(patch.write_as_warp_sse_event(), expected);
380        assert_event(Event::from(&patch), expected);
381        assert_event(Event::from(patch), expected);
382    }
383
384    #[test]
385    fn writes_execute_script_and_conversions() {
386        let script = ExecuteScript::new("console.log('hello')");
387        let expected = concat!(
388            "event:datastar-patch-elements\n",
389            "data:selector body\n",
390            "data:mode append\n",
391            "data:elements <script data-effect=\"el.remove()\">",
392            "console.log('hello')</script>\n\n",
393        );
394
395        assert_event(script.write_as_warp_sse_event(), expected);
396        assert_event(Event::from(&script), expected);
397        assert_event(Event::from(script), expected);
398    }
399
400    #[test]
401    fn writes_generic_events_and_conversions() {
402        let event = PatchSignals::new("{count: 1}")
403            .id("signals-1")
404            .retry(Duration::from_millis(2_500))
405            .into_datastar_event();
406        let expected = concat!(
407            "event:datastar-patch-signals\n",
408            "data:signals {count: 1}\n",
409            "id:signals-1\n",
410            "retry:2500\n\n",
411        );
412
413        assert_event(event.write_as_warp_sse_event(), expected);
414        assert_event(Event::from(&event), expected);
415        assert_event(Event::from(event), expected);
416    }
417
418    #[derive(Debug, Deserialize, PartialEq)]
419    struct TestSignals {
420        count: u64,
421    }
422
423    #[tokio::test]
424    async fn extracts_get_and_body_signals() {
425        let get = warp::test::request()
426            .method("GET")
427            .path("/?datastar=%7B%22count%22%3A7%7D")
428            .filter(&read_signals::<TestSignals>())
429            .await
430            .unwrap();
431        assert_eq!(get.0, TestSignals { count: 7 });
432
433        let delete = warp::test::request()
434            .method("DELETE")
435            .path("/?datastar=%7B%22count%22%3A8%7D")
436            .filter(&read_signals::<TestSignals>())
437            .await
438            .unwrap();
439        assert_eq!(delete.0, TestSignals { count: 8 });
440
441        let post = warp::test::request()
442            .method("POST")
443            .body(r#"{"count":9}"#)
444            .filter(&read_signals::<TestSignals>())
445            .await
446            .unwrap();
447        assert_eq!(post.0, TestSignals { count: 9 });
448    }
449
450    #[tokio::test]
451    async fn handles_optional_signals_and_request_header() {
452        let present = warp::test::request()
453            .method("GET")
454            .path("/?datastar=%7B%22count%22%3A7%7D")
455            .header(DATASTAR_REQ_HEADER_STR, "true")
456            .filter(&read_signals_optional::<TestSignals>())
457            .await
458            .unwrap();
459        assert_eq!(present.unwrap().0, TestSignals { count: 7 });
460
461        let missing = warp::test::request()
462            .filter(&read_signals_optional::<TestSignals>())
463            .await
464            .unwrap();
465        assert!(missing.is_none());
466
467        let present = warp::test::request()
468            .header(DATASTAR_REQ_HEADER_STR, "true")
469            .filter(&is_datastar_request())
470            .await
471            .unwrap();
472        assert!(present);
473
474        let missing = warp::test::request()
475            .filter(&is_datastar_request())
476            .await
477            .unwrap();
478        assert!(!missing);
479    }
480
481    #[tokio::test]
482    async fn handles_missing_get_signals() {
483        let optional = warp::test::request()
484            .method("GET")
485            .filter(&read_signals::<Option<TestSignals>>())
486            .await
487            .unwrap();
488        assert_eq!(optional.0, None);
489
490        let rejection = warp::test::request()
491            .method("GET")
492            .filter(&read_signals::<TestSignals>())
493            .await
494            .unwrap_err();
495        let response = handle_rejection(rejection).await.unwrap().into_response();
496        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
497    }
498
499    #[tokio::test]
500    async fn maps_signal_rejections_to_responses() {
501        let rejection = warp::test::request()
502            .method("GET")
503            .path("/?datastar=not-json")
504            .filter(&read_signals::<TestSignals>())
505            .await
506            .unwrap_err();
507        let response = handle_rejection(rejection).await.unwrap().into_response();
508        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
509
510        let response = handle_rejection(warp::reject::not_found())
511            .await
512            .unwrap()
513            .into_response();
514        assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
515    }
516}