Skip to main content

rama_http/protocols/html/rewrite/
rewriter.rs

1//! The selector-driven HTML rewriter.
2
3use rama_core::error::BoxError;
4
5use super::super::selector::Selector;
6use super::super::tokenizer::{
7    Cdata, Comment, Doctype, EndTag, StartTag, Text, TokenSink, Tokenizer,
8};
9use super::SelectorMatcher;
10use super::element::{Element, ElementContentHandler, EndActions, HandlerResult};
11
12/// The [`TokenSink`] that drives matching + mutation + serialization.
13struct RewriteSink<H> {
14    matcher: SelectorMatcher,
15    handler: H,
16    output: Vec<u8>,
17    /// Reused scratch for the selectors matching the current element.
18    matched: Vec<usize>,
19    /// Deferred end-tag actions, one entry per open scope — kept in lockstep
20    /// with the matcher's open-element stack (see [`SelectorMatcher`]).
21    pending: Vec<EndActions>,
22    /// Number of open ancestors currently suppressing their content (from a
23    /// `remove` / `replace` / `set_inner_content`). While non-zero, token
24    /// output is swallowed.
25    suppress_depth: usize,
26    /// First handler error, if any (aborts the rewrite).
27    error: Option<BoxError>,
28}
29
30impl<H: ElementContentHandler> TokenSink for RewriteSink<H> {
31    fn start_tag(&mut self, tag: &StartTag<'_>) {
32        if self.error.is_some() {
33            self.output.extend_from_slice(tag.raw());
34            return;
35        }
36        let Self {
37            matcher,
38            handler,
39            output,
40            matched,
41            pending,
42            suppress_depth,
43            error,
44        } = self;
45
46        let implied = matcher.pop_implied_for_start(tag.name_hash());
47        close_pending(output, pending, suppress_depth, implied, None);
48
49        // Visible iff no enclosing element is suppressing its content after
50        // any optional-end-tag frames closed by this start tag.
51        let visible = *suppress_depth == 0;
52        matched.clear();
53        let opened = matcher.push_element(tag, |index| matched.push(index));
54
55        if matched.is_empty() {
56            if visible {
57                output.extend_from_slice(tag.raw());
58            }
59            if opened {
60                pending.push(EndActions::passthrough());
61            }
62            return;
63        }
64
65        let mut element = Element::new(tag);
66        for &index in matched.iter() {
67            if let Err(err) = handler.handle_element(index, &mut element) {
68                *error = Some(err);
69                break;
70            }
71        }
72        let actions = element.serialize(output, visible);
73
74        if opened {
75            if actions.suppress_content {
76                *suppress_depth += 1;
77            }
78            pending.push(actions);
79        } else if visible {
80            // Void / self-closing: no children and no end tag, so the
81            // end-anchored content (if any) lands right here.
82            output.extend_from_slice(actions.append.as_bytes());
83            output.extend_from_slice(actions.after.as_bytes());
84        }
85    }
86
87    fn end_tag(&mut self, tag: &EndTag<'_>) {
88        if self.error.is_some() {
89            self.output.extend_from_slice(tag.raw());
90            return;
91        }
92        let popped = self.matcher.pop_element(tag.name_hash());
93        if popped == 0 {
94            // Stray end tag: emit verbatim unless inside suppressed content.
95            if self.suppress_depth == 0 {
96                self.output.extend_from_slice(tag.raw());
97            }
98            return;
99        }
100        // The end tag closes `popped` frames: the named element plus any
101        // still-open descendants it implicitly closes. Every closed frame gets
102        // its deferred end actions at this point; only the named frame owns
103        // the source end tag bytes.
104        close_pending(
105            &mut self.output,
106            &mut self.pending,
107            &mut self.suppress_depth,
108            popped,
109            Some(tag.raw()),
110        );
111    }
112
113    fn text(&mut self, text: &Text<'_>) {
114        self.passthrough(text.raw());
115    }
116
117    fn comment(&mut self, comment: &Comment<'_>) {
118        self.passthrough(comment.raw());
119    }
120
121    fn cdata(&mut self, cdata: &Cdata<'_>) {
122        self.passthrough(cdata.raw());
123    }
124
125    fn doctype(&mut self, doctype: &Doctype<'_>) {
126        self.passthrough(doctype.raw());
127    }
128}
129
130/// A streaming, selector-driven HTML rewriter.
131///
132/// Feed input with [`write`](Self::write) and finish with
133/// [`end`](Self::end); the rewritten bytes accumulate in an internal buffer,
134/// drained with [`take_output`](Self::take_output). Unmatched content passes
135/// through byte-for-byte. For one-shot use, prefer [`rewrite_str`].
136pub struct HtmlRewriter<H> {
137    tokenizer: Tokenizer,
138    sink: RewriteSink<H>,
139}
140
141impl<H: ElementContentHandler> HtmlRewriter<H> {
142    /// Creates a rewriter that runs `handler` for elements matching
143    /// `selectors` (the `selector` argument to the handler is the index into
144    /// this slice).
145    #[must_use]
146    pub fn new(selectors: &[Selector], handler: H) -> Self {
147        Self {
148            tokenizer: Tokenizer::new(),
149            sink: RewriteSink {
150                matcher: SelectorMatcher::new(selectors),
151                handler,
152                output: Vec::new(),
153                matched: Vec::new(),
154                pending: Vec::new(),
155                suppress_depth: 0,
156                error: None,
157            },
158        }
159    }
160
161    /// Feeds a chunk of input, appending rewritten bytes to the output.
162    ///
163    /// # Errors
164    ///
165    /// Surfaces a handler error, or a [`ParsingAmbiguityError`] if the input
166    /// is ambiguous for streaming parsing.
167    ///
168    /// [`ParsingAmbiguityError`]: crate::protocols::html::tokenizer::ParsingAmbiguityError
169    pub fn write(&mut self, chunk: &[u8]) -> Result<(), BoxError> {
170        self.tokenizer.write(chunk, &mut self.sink)?;
171        self.sink.take_error()
172    }
173
174    /// Finalizes the stream, flushing any remaining input.
175    ///
176    /// # Errors
177    ///
178    /// See [`write`](Self::write).
179    pub fn end(&mut self) -> Result<(), BoxError> {
180        self.tokenizer.end(&mut self.sink)?;
181        self.sink.finish();
182        self.sink.take_error()
183    }
184
185    /// Removes and returns the rewritten output accumulated so far.
186    #[must_use]
187    pub fn take_output(&mut self) -> Vec<u8> {
188        std::mem::take(&mut self.sink.output)
189    }
190
191    /// Consumes the rewriter, returning the handler (e.g. to read state
192    /// accumulated during the rewrite).
193    #[must_use]
194    pub fn into_handler(self) -> H {
195        self.sink.handler
196    }
197}
198
199impl<H> RewriteSink<H> {
200    /// Emits a leaf token's raw bytes, unless it is inside suppressed content.
201    /// (On error the rewrite is doomed and its output discarded, so bytes are
202    /// passed through to keep the buffer well-formed for inspection.)
203    fn passthrough(&mut self, raw: &[u8]) {
204        if self.error.is_some() || self.suppress_depth == 0 {
205            self.output.extend_from_slice(raw);
206        }
207    }
208
209    fn take_error(&mut self) -> Result<(), BoxError> {
210        match self.error.take() {
211            Some(err) => Err(err),
212            None => Ok(()),
213        }
214    }
215
216    fn finish(&mut self) {
217        let popped = self.matcher.finish();
218        close_pending(
219            &mut self.output,
220            &mut self.pending,
221            &mut self.suppress_depth,
222            popped,
223            None,
224        );
225    }
226}
227
228fn close_pending(
229    output: &mut Vec<u8>,
230    pending: &mut Vec<EndActions>,
231    suppress_depth: &mut usize,
232    popped: usize,
233    named_end_tag: Option<&[u8]>,
234) {
235    for i in 0..popped {
236        let Some(actions) = pending.pop() else {
237            break;
238        };
239        if actions.suppress_content {
240            *suppress_depth = (*suppress_depth).saturating_sub(1);
241        }
242        if *suppress_depth == 0 {
243            output.extend_from_slice(actions.append.as_bytes());
244            if i + 1 == popped
245                && let Some(raw) = named_end_tag
246                && !actions.suppress_end_tag
247            {
248                output.extend_from_slice(raw);
249            }
250            output.extend_from_slice(actions.after.as_bytes());
251        }
252    }
253}
254
255type BoxedHandler<'h> = Box<dyn FnMut(&mut Element<'_>) -> HandlerResult + 'h>;
256
257/// A builder bundling `(selector, closure)` pairs — the closure-based escape
258/// hatch over the [`ElementContentHandler`] trait, for one-off rewrites that
259/// don't need a dedicated state struct.
260#[derive(Default)]
261pub struct ElementContentHandlers<'h> {
262    selectors: Vec<Selector>,
263    handlers: Vec<BoxedHandler<'h>>,
264}
265
266impl<'h> ElementContentHandlers<'h> {
267    /// Creates an empty set of handlers.
268    #[must_use]
269    pub fn new() -> Self {
270        Self::default()
271    }
272
273    /// Registers `handler` for elements matching `selector`.
274    #[must_use]
275    pub fn on(
276        mut self,
277        selector: Selector,
278        handler: impl FnMut(&mut Element<'_>) -> HandlerResult + 'h,
279    ) -> Self {
280        self.selectors.push(selector);
281        self.handlers.push(Box::new(handler));
282        self
283    }
284}
285
286impl ElementContentHandler for ElementContentHandlers<'_> {
287    fn handle_element(&mut self, selector: usize, element: &mut Element<'_>) -> HandlerResult {
288        match self.handlers.get_mut(selector) {
289            Some(handler) => handler(element),
290            None => Ok(()),
291        }
292    }
293}
294
295impl<'h> HtmlRewriter<ElementContentHandlers<'h>> {
296    /// Creates a rewriter from a closure-based [`ElementContentHandlers`].
297    #[must_use]
298    pub fn from_handlers(handlers: ElementContentHandlers<'h>) -> Self {
299        let selectors = handlers.selectors.clone();
300        Self::new(&selectors, handlers)
301    }
302}
303
304/// One-shot rewrite of a complete HTML string.
305///
306/// # Errors
307///
308/// Surfaces a handler error, a parsing-ambiguity error, or invalid UTF-8 in
309/// the rewritten output.
310pub fn rewrite_str(html: &str, handlers: ElementContentHandlers<'_>) -> Result<String, BoxError> {
311    let mut rewriter = HtmlRewriter::from_handlers(handlers);
312    rewriter.write(html.as_bytes())?;
313    rewriter.end()?;
314    String::from_utf8(rewriter.take_output()).map_err(Into::into)
315}
316
317#[cfg(test)]
318mod tests {
319    use super::{ElementContentHandlers, HtmlRewriter, rewrite_str};
320    use crate::protocols::html::PreEscaped;
321    use crate::protocols::html::rewrite::{
322        AttributeName, Element, ElementContentHandler, HandlerResult,
323    };
324    use crate::protocols::html::selector::Selector;
325
326    fn sel(s: &str) -> Selector {
327        s.parse()
328            .unwrap_or_else(|e| panic!("`{s}` should parse: {e}"))
329    }
330
331    fn rewrite(html: &str, handlers: ElementContentHandlers<'_>) -> String {
332        rewrite_str(html, handlers).expect("rewrite succeeds")
333    }
334
335    #[test]
336    fn unmatched_passes_through_verbatim() {
337        let out = rewrite("<p>hi <b>x</b></p>", ElementContentHandlers::new());
338        assert_eq!(out, "<p>hi <b>x</b></p>");
339    }
340
341    #[test]
342    fn set_and_remove_attributes() {
343        let out = rewrite(
344            r#"<a href="/old" data-x="1">link</a>"#,
345            ElementContentHandlers::new().on(sel("a"), |el| {
346                el.set_attribute(AttributeName::from_static("href"), "/new");
347                el.remove_attribute("data-x");
348                el.set_attribute(AttributeName::from_static("rel"), "nofollow");
349                Ok(())
350            }),
351        );
352        assert_eq!(out, r#"<a href="/new" rel="nofollow">link</a>"#);
353    }
354
355    #[test]
356    fn attribute_values_are_escaped() {
357        let out = rewrite(
358            "<a>x</a>",
359            ElementContentHandlers::new().on(sel("a"), |el| {
360                el.set_attribute(AttributeName::from_static("title"), r#"a "b" & c"#);
361                Ok(())
362            }),
363        );
364        assert_eq!(out, r#"<a title="a &quot;b&quot; &amp; c">x</a>"#);
365    }
366
367    #[test]
368    fn before_and_prepend() {
369        let out = rewrite(
370            "<body>content</body>",
371            ElementContentHandlers::new().on(sel("body"), |el| {
372                el.before("X");
373                el.prepend("Y<&");
374                Ok(())
375            }),
376        );
377        // `before` precedes the start tag; `prepend` follows it; text escaped.
378        assert_eq!(out, "X<body>Y&lt;&amp;content</body>");
379    }
380
381    #[test]
382    fn reading_attributes() {
383        let out = rewrite(
384            r#"<a href="/x" disabled>k</a>"#,
385            ElementContentHandlers::new().on(sel("a"), |el| {
386                assert_eq!(el.attribute("href"), Some(&b"/x"[..]));
387                assert!(el.has_attribute("disabled"));
388                assert_eq!(el.attribute("disabled"), Some(&b""[..]));
389                assert_eq!(el.attribute("missing"), None);
390                Ok(())
391            }),
392        );
393        assert_eq!(out, r#"<a href="/x" disabled>k</a>"#);
394    }
395
396    #[test]
397    fn only_matching_elements_are_touched() {
398        let out = rewrite(
399            "<div><span>a</span><span>b</span></div>",
400            ElementContentHandlers::new().on(sel("div > span"), |el| {
401                el.set_attribute(AttributeName::from_static("data-hit"), "1");
402                Ok(())
403            }),
404        );
405        assert_eq!(
406            out,
407            r#"<div><span data-hit="1">a</span><span data-hit="1">b</span></div>"#
408        );
409    }
410
411    #[test]
412    fn handler_error_aborts() {
413        rewrite_str(
414            "<a></a>",
415            ElementContentHandlers::new().on(sel("a"), |_el| Err("boom".into())),
416        )
417        .expect_err("handler error should abort the rewrite");
418    }
419
420    /// A handler struct carrying its own accumulated state.
421    #[derive(Default)]
422    struct LinkCounter {
423        count: usize,
424    }
425
426    impl ElementContentHandler for LinkCounter {
427        fn handle_element(&mut self, _selector: usize, element: &mut Element<'_>) -> HandlerResult {
428            self.count += 1;
429            element.set_attribute(
430                AttributeName::from_static("data-n"),
431                &self.count.to_string(),
432            );
433            Ok(())
434        }
435    }
436
437    #[test]
438    fn visitor_trait_shares_state() {
439        let selectors = [sel("a")];
440        let mut rewriter = HtmlRewriter::new(&selectors, LinkCounter::default());
441        rewriter.write(b"<a>1</a><a>2</a>").expect("write succeeds");
442        rewriter.end().expect("end succeeds");
443        let out = String::from_utf8(rewriter.take_output()).expect("utf8");
444        assert_eq!(out, r#"<a data-n="1">1</a><a data-n="2">2</a>"#);
445        assert_eq!(rewriter.into_handler().count, 2);
446    }
447
448    // --- slice B: end-anchored edits ------------------------------------
449
450    #[test]
451    fn append_inserts_before_end_tag() {
452        let out = rewrite(
453            "<div>x</div>",
454            ElementContentHandlers::new().on(sel("div"), |el| {
455                el.append("!");
456                Ok(())
457            }),
458        );
459        assert_eq!(out, "<div>x!</div>");
460    }
461
462    #[test]
463    fn after_inserts_after_end_tag() {
464        let out = rewrite(
465            "<div>x</div>",
466            ElementContentHandlers::new().on(sel("div"), |el| {
467                el.after("Y");
468                Ok(())
469            }),
470        );
471        assert_eq!(out, "<div>x</div>Y");
472    }
473
474    #[test]
475    fn set_inner_content_replaces_children() {
476        let out = rewrite(
477            "<div>old<b>stuff</b></div>",
478            ElementContentHandlers::new().on(sel("div"), |el| {
479                el.set_inner_content("new");
480                Ok(())
481            }),
482        );
483        assert_eq!(out, "<div>new</div>");
484    }
485
486    #[test]
487    fn set_inner_content_keeps_attribute_edits() {
488        let out = rewrite(
489            r#"<div class="a">old</div>"#,
490            ElementContentHandlers::new().on(sel("div"), |el| {
491                el.set_attribute(AttributeName::from_static("data-x"), "1");
492                el.set_inner_content("new");
493                Ok(())
494            }),
495        );
496        assert_eq!(out, r#"<div class="a" data-x="1">new</div>"#);
497    }
498
499    #[test]
500    fn replace_swaps_whole_element() {
501        let out = rewrite(
502            "a<p>hi</p>b",
503            ElementContentHandlers::new().on(sel("p"), |el| {
504                el.replace("X");
505                Ok(())
506            }),
507        );
508        assert_eq!(out, "aXb");
509    }
510
511    #[test]
512    fn replace_then_after() {
513        let out = rewrite(
514            "<p>hi</p>",
515            ElementContentHandlers::new().on(sel("p"), |el| {
516                el.replace("R");
517                el.after("A");
518                Ok(())
519            }),
520        );
521        assert_eq!(out, "RA");
522    }
523
524    #[test]
525    fn remove_drops_element_and_children() {
526        let out = rewrite(
527            "a<p>h<b>i</b></p>b",
528            ElementContentHandlers::new().on(sel("p"), |el| {
529                el.remove();
530                Ok(())
531            }),
532        );
533        assert_eq!(out, "ab");
534    }
535
536    #[test]
537    fn remove_keeps_before_and_after() {
538        let out = rewrite(
539            "x<p>hi</p>y",
540            ElementContentHandlers::new().on(sel("p"), |el| {
541                el.before("B");
542                el.remove();
543                el.after("A");
544                Ok(())
545            }),
546        );
547        // `before` then (element gone) then `after`.
548        assert_eq!(out, "xBAy");
549    }
550
551    #[test]
552    fn remove_and_keep_content_drops_only_tags() {
553        let out = rewrite(
554            "a<p>hi</p>b",
555            ElementContentHandlers::new().on(sel("p"), |el| {
556                el.remove_and_keep_content();
557                Ok(())
558            }),
559        );
560        assert_eq!(out, "ahib");
561    }
562
563    #[test]
564    fn match_inside_removed_ancestor_is_swallowed() {
565        // The inner handler still runs, but its output is suppressed by the
566        // removed ancestor.
567        let out = rewrite(
568            "<div><a>x</a></div>",
569            ElementContentHandlers::new()
570                .on(sel("div"), |el| {
571                    el.remove();
572                    Ok(())
573                })
574                .on(sel("a"), |el| {
575                    el.set_attribute(AttributeName::from_static("data-hit"), "1");
576                    Ok(())
577                }),
578        );
579        assert_eq!(out, "");
580    }
581
582    #[test]
583    fn after_on_void_element() {
584        let out = rewrite(
585            "<img src=x>tail",
586            ElementContentHandlers::new().on(sel("img"), |el| {
587                el.after("Y");
588                Ok(())
589            }),
590        );
591        // No end tag for a void element: `after` lands right after it.
592        assert_eq!(out, "<img src=x>Ytail");
593    }
594
595    #[test]
596    fn append_accepts_into_html() {
597        let out = rewrite(
598            "<div>x</div>",
599            ElementContentHandlers::new().on(sel("div"), |el| {
600                el.append(PreEscaped("<i>!</i>"));
601                Ok(())
602            }),
603        );
604        // `PreEscaped` content is written verbatim (the `IntoHtml` path).
605        assert_eq!(out, "<div>x<i>!</i></div>");
606    }
607
608    #[test]
609    fn remove_survives_chunk_boundaries() {
610        // Suppression state must persist across `write` calls.
611        let mut rewriter =
612            HtmlRewriter::from_handlers(ElementContentHandlers::new().on(sel("div"), |el| {
613                el.remove();
614                Ok(())
615            }));
616        for chunk in [&b"a<div>con"[..], b"tent</di", b"v>b"] {
617            rewriter.write(chunk).expect("write succeeds");
618        }
619        rewriter.end().expect("end succeeds");
620        let out = String::from_utf8(rewriter.take_output()).expect("utf8");
621        assert_eq!(out, "ab");
622    }
623
624    // --- robustness on malformed / crossed nesting --------------------------
625
626    #[test]
627    fn remove_survives_crossed_nesting() {
628        // `</d>` implicitly closes the still-open `<e>`; suppression must clear
629        // so the trailing text is not swallowed.
630        let out = rewrite(
631            "<d><e></d>VISIBLE",
632            ElementContentHandlers::new().on(sel("d"), |el| {
633                el.remove();
634                Ok(())
635            }),
636        );
637        assert_eq!(out, "VISIBLE");
638    }
639
640    #[test]
641    fn remove_with_unclosed_child_keeps_the_rest() {
642        // Only the `<a>…</a>` span is removed; the misnested remainder passes
643        // through byte-for-byte (no runaway suppression).
644        let out = rewrite(
645            "keep<a>1<b>2</a>3</b>4",
646            ElementContentHandlers::new().on(sel("a"), |el| {
647                el.remove();
648                Ok(())
649            }),
650        );
651        assert_eq!(out, "keep3</b>4");
652    }
653
654    #[test]
655    fn set_inner_content_survives_crossed_nesting() {
656        let out = rewrite(
657            "<a>1<b>2</a>3",
658            ElementContentHandlers::new().on(sel("a"), |el| {
659                el.set_inner_content("X");
660                Ok(())
661            }),
662        );
663        assert_eq!(out, "<a>X</a>3");
664    }
665
666    #[test]
667    fn nested_match_inside_removed_ancestor_is_swallowed() {
668        // The inner `<a>`'s handler still runs, but its end-anchored output is
669        // suppressed along with the rest of the removed subtree.
670        let out = rewrite(
671            "<div><a>x</a></div>z",
672            ElementContentHandlers::new()
673                .on(sel("div"), |el| {
674                    el.remove();
675                    Ok(())
676                })
677                .on(sel("a"), |el| {
678                    el.after("!");
679                    el.append("?");
680                    Ok(())
681                }),
682        );
683        assert_eq!(out, "z");
684    }
685
686    #[test]
687    fn replace_on_self_closing_element() {
688        // A self-closing (non-void) element has no end tag: the replacement is
689        // emitted inline and suppression is never engaged.
690        let out = rewrite(
691            "<x/>tail",
692            ElementContentHandlers::new().on(sel("x"), |el| {
693                el.replace("R");
694                Ok(())
695            }),
696        );
697        assert_eq!(out, "Rtail");
698    }
699
700    #[test]
701    fn append_applies_to_optional_li_end_tags() {
702        let out = rewrite(
703            "<ul><li>one<li>two</ul>",
704            ElementContentHandlers::new().on(sel("li"), |el| {
705                el.append("!");
706                Ok(())
707            }),
708        );
709        assert_eq!(out, "<ul><li>one!<li>two!</ul>");
710    }
711
712    #[test]
713    fn start_implied_p_close_restores_suppression() {
714        #[derive(Default)]
715        struct RemoveFirst {
716            seen: usize,
717        }
718
719        impl ElementContentHandler for RemoveFirst {
720            fn handle_element(
721                &mut self,
722                _selector: usize,
723                element: &mut Element<'_>,
724            ) -> HandlerResult {
725                if self.seen == 0 {
726                    element.remove();
727                } else {
728                    element.append("!");
729                }
730                self.seen += 1;
731                Ok(())
732            }
733        }
734
735        let selectors = [sel("p")];
736        let mut rewriter = HtmlRewriter::new(&selectors, RemoveFirst::default());
737        rewriter
738            .write(b"<p>drop<p>keep</p>")
739            .expect("write succeeds");
740        rewriter.end().expect("end succeeds");
741        let out = String::from_utf8(rewriter.take_output()).expect("utf8");
742        assert_eq!(out, "<p>keep!</p>");
743    }
744
745    #[test]
746    fn eof_flushes_end_anchored_actions() {
747        let out = rewrite(
748            "<div>tail",
749            ElementContentHandlers::new().on(sel("div"), |el| {
750                el.append("!");
751                el.after("A");
752                Ok(())
753            }),
754        );
755        assert_eq!(out, "<div>tail!A");
756    }
757}