1use 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
12struct RewriteSink<H> {
14 matcher: SelectorMatcher,
15 handler: H,
16 output: Vec<u8>,
17 matched: Vec<usize>,
19 pending: Vec<EndActions>,
22 suppress_depth: usize,
26 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 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 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 if self.suppress_depth == 0 {
96 self.output.extend_from_slice(tag.raw());
97 }
98 return;
99 }
100 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
130pub struct HtmlRewriter<H> {
137 tokenizer: Tokenizer,
138 sink: RewriteSink<H>,
139}
140
141impl<H: ElementContentHandler> HtmlRewriter<H> {
142 #[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 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 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 #[must_use]
187 pub fn take_output(&mut self) -> Vec<u8> {
188 std::mem::take(&mut self.sink.output)
189 }
190
191 #[must_use]
194 pub fn into_handler(self) -> H {
195 self.sink.handler
196 }
197}
198
199impl<H> RewriteSink<H> {
200 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#[derive(Default)]
261pub struct ElementContentHandlers<'h> {
262 selectors: Vec<Selector>,
263 handlers: Vec<BoxedHandler<'h>>,
264}
265
266impl<'h> ElementContentHandlers<'h> {
267 #[must_use]
269 pub fn new() -> Self {
270 Self::default()
271 }
272
273 #[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 #[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
304pub 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 "b" & 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 assert_eq!(out, "X<body>Y<&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 #[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 #[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 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 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 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 assert_eq!(out, "<div>x<i>!</i></div>");
606 }
607
608 #[test]
609 fn remove_survives_chunk_boundaries() {
610 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 #[test]
627 fn remove_survives_crossed_nesting() {
628 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 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 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 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}