use crate::ParserLimits;
use crate::types::{Content, Entry, FeedMeta, MimeType, ParsedFeed, TextConstruct, TextType};
use ammonia::Builder;
use std::collections::HashSet;
use std::sync::LazyLock;
const SAFE_TAGS: &[&str] = &[
"a",
"abbr",
"acronym",
"b",
"cite",
"code",
"em",
"i",
"kbd",
"mark",
"s",
"samp",
"small",
"strike",
"strong",
"sub",
"sup",
"u",
"var", "br",
"div",
"hr",
"p",
"span", "h1",
"h2",
"h3",
"h4",
"h5",
"h6", "dd",
"dl",
"dt",
"li",
"ol",
"ul", "caption",
"table",
"tbody",
"td",
"tfoot",
"th",
"thead",
"tr", "blockquote",
"q", "pre", "img",
];
static SAFE_HTML_BUILDER: LazyLock<Builder<'static>> = LazyLock::new(|| {
let safe_tags: HashSet<&'static str> = SAFE_TAGS.iter().copied().collect();
let safe_attrs: HashSet<&'static str> = ["alt", "cite", "class", "href", "id", "src", "title"]
.into_iter()
.collect();
let safe_url_schemes: HashSet<&'static str> = ["http", "https", "mailto"].into_iter().collect();
let mut builder = Builder::default();
builder
.tags(safe_tags)
.generic_attributes(safe_attrs)
.link_rel(Some("nofollow noopener noreferrer"))
.url_schemes(safe_url_schemes);
builder
});
pub fn sanitize_html(input: &str) -> String {
SAFE_HTML_BUILDER.clean(input).to_string()
}
pub fn decode_entities(input: &str) -> String {
html_escape::decode_html_entities(input).to_string()
}
pub fn strip_tags(input: &str) -> String {
Builder::default()
.tags(HashSet::new())
.clean(input)
.to_string()
}
pub fn sanitize_feed(feed: &mut ParsedFeed, limits: &ParserLimits) {
sanitize_feed_meta(&mut feed.feed, limits);
for entry in &mut feed.entries {
sanitize_entry(entry, limits);
}
}
fn sanitize_feed_meta(meta: &mut FeedMeta, limits: &ParserLimits) {
sanitize_pair(&mut meta.title, &mut meta.title_detail, limits);
sanitize_pair(&mut meta.subtitle, &mut meta.subtitle_detail, limits);
sanitize_pair(&mut meta.summary, &mut meta.summary_detail, limits);
sanitize_pair(&mut meta.rights, &mut meta.rights_detail, limits);
sanitize_opt(&mut meta.dc_rights, limits);
if let Some(image) = &mut meta.image {
sanitize_opt(&mut image.title, limits);
sanitize_opt(&mut image.description, limits);
}
if let Some(textinput) = &mut meta.textinput {
sanitize_opt(&mut textinput.title, limits);
sanitize_opt(&mut textinput.description, limits);
}
if let Some(itunes) = &mut meta.itunes {
sanitize_opt(&mut itunes.subtitle, limits);
sanitize_opt(&mut itunes.summary, limits);
}
}
fn sanitize_entry(entry: &mut Entry, limits: &ParserLimits) {
sanitize_pair(&mut entry.title, &mut entry.title_detail, limits);
sanitize_pair(&mut entry.subtitle, &mut entry.subtitle_detail, limits);
sanitize_pair(&mut entry.summary, &mut entry.summary_detail, limits);
sanitize_pair(&mut entry.rights, &mut entry.rights_detail, limits);
sanitize_opt(&mut entry.dc_rights, limits);
sanitize_opt(&mut entry.media_title, limits);
sanitize_opt(&mut entry.media_description, limits);
for content in &mut entry.content {
sanitize_content(content, limits);
}
if let Some(source) = &mut entry.source {
sanitize_opt(&mut source.title, limits);
sanitize_opt(&mut source.rights, limits);
}
if let Some(itunes) = &mut entry.itunes {
sanitize_opt(&mut itunes.title, limits);
sanitize_opt(&mut itunes.subtitle, limits);
sanitize_opt(&mut itunes.summary, limits);
}
}
fn sanitize_pair(
value: &mut Option<String>,
detail: &mut Option<TextConstruct>,
limits: &ParserLimits,
) {
if matches!(
detail.as_ref().map(|d| d.content_type),
Some(TextType::Text)
) {
return;
}
sanitize_opt(value, limits);
if let Some(detail) = detail {
detail.value = sanitize_html_bounded(&detail.value, limits.max_html_nesting_depth);
}
}
fn sanitize_content(content: &mut Content, limits: &ParserLimits) {
if content
.content_type
.as_deref()
.is_some_and(|t| t.eq_ignore_ascii_case(MimeType::TEXT_PLAIN))
{
return;
}
content.value = sanitize_html_bounded(&content.value, limits.max_html_nesting_depth);
}
fn sanitize_opt(value: &mut Option<String>, limits: &ParserLimits) {
if let Some(v) = value {
*v = sanitize_html_bounded(v, limits.max_html_nesting_depth);
}
}
fn sanitize_html_bounded(input: &str, max_depth: usize) -> String {
if html_nesting_exceeds(input, max_depth) {
escape_html_plain(input)
} else {
sanitize_html(input)
}
}
const VOID_ELEMENTS: &[&str] = &[
"area", "base", "br", "col", "embed", "hr", "img", "input", "link", "meta", "param", "source",
"track", "wbr",
];
const AUTO_CLOSE_ELEMENTS: &[&str] = &["li", "option", "p", "td", "th", "tr"];
const SCOPE_BARRIERS: &[&str] = &[
"table", "td", "th", "caption", "object", "marquee", "applet",
];
const FORMATTING_ELEMENTS: &[&str] = &[
"a", "b", "big", "code", "em", "font", "i", "nobr", "s", "small", "strike", "strong", "tt", "u",
];
const MAX_TAGS_PER_FIELD: usize = 10_000;
fn tag_name(inner: &[u8]) -> &[u8] {
let end = inner
.iter()
.position(|b| b.is_ascii_whitespace() || *b == b'/')
.unwrap_or(inner.len());
&inner[..end]
}
fn contains_name_ci(names: &[&str], name: &[u8]) -> bool {
names
.iter()
.any(|n| name.eq_ignore_ascii_case(n.as_bytes()))
}
fn find_in_scope(stack: &[&[u8]], name: &[u8]) -> Option<usize> {
for (idx, tag) in stack.iter().enumerate().rev() {
if tag.eq_ignore_ascii_case(name) {
return Some(idx);
}
if contains_name_ci(SCOPE_BARRIERS, tag) {
return None;
}
}
None
}
fn html_nesting_exceeds(html: &str, max_depth: usize) -> bool {
let bytes = html.as_bytes();
let mut stack: Vec<&[u8]> = Vec::new();
let mut tag_count: usize = 0;
let mut i = 0;
while i < bytes.len() {
if bytes[i] != b'<' {
i += 1;
continue;
}
let Some(rel_end) = bytes[i..].iter().position(|&b| b == b'>') else {
break; };
let inner = &bytes[i + 1..i + rel_end];
i += rel_end + 1;
if inner.first() == Some(&b'!') || inner.first() == Some(&b'?') {
continue; }
tag_count += 1;
if tag_count > MAX_TAGS_PER_FIELD {
return true;
}
if inner.first() == Some(&b'/') {
let name = tag_name(&inner[1..]);
if let Some(pos) = find_in_scope(&stack, name) {
stack.truncate(pos);
}
continue;
}
let self_closing = inner.last() == Some(&b'/');
let name = tag_name(if self_closing {
&inner[..inner.len() - 1]
} else {
inner
});
if self_closing
|| contains_name_ci(VOID_ELEMENTS, name)
|| contains_name_ci(FORMATTING_ELEMENTS, name)
{
continue; }
if contains_name_ci(AUTO_CLOSE_ELEMENTS, name)
&& let Some(pos) = find_in_scope(&stack, name)
{
stack.truncate(pos);
}
stack.push(name);
if stack.len() > max_depth {
return true;
}
}
false
}
fn escape_html_plain(input: &str) -> String {
let mut out = String::with_capacity(input.len());
for ch in input.chars() {
match ch {
'&' => out.push_str("&"),
'<' => out.push_str("<"),
'>' => out.push_str(">"),
'"' => out.push_str("""),
'\'' => out.push_str("'"),
_ => out.push(ch),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sanitize_removes_script() {
let html = r"<p>Hello</p><script>alert('XSS')</script>";
let clean = sanitize_html(html);
assert!(!clean.contains("script"));
assert!(clean.contains("Hello"));
}
#[test]
fn test_sanitize_allows_safe_tags() {
let html = r#"<p>Hello <b>world</b> <a href="http://example.com">link</a></p>"#;
let clean = sanitize_html(html);
assert!(clean.contains("<p>"));
assert!(clean.contains("<b>"));
assert!(clean.contains("<a"));
}
#[test]
fn test_sanitize_removes_onclick() {
let html = r#"<a href="/" onclick="alert('XSS')">Click</a>"#;
let clean = sanitize_html(html);
assert!(!clean.contains("onclick"));
assert!(clean.contains("href"));
}
#[test]
fn test_decode_entities() {
assert_eq!(decode_entities("<p>"), "<p>");
assert_eq!(decode_entities("&"), "&");
assert_eq!(decode_entities("""), "\"");
assert_eq!(decode_entities("'"), "'");
}
#[test]
fn test_decode_numeric_entities() {
assert_eq!(decode_entities("<"), "<");
assert_eq!(decode_entities("<"), "<");
}
#[test]
fn test_strip_tags() {
let html = "<p>Hello <b>world</b></p>";
assert_eq!(strip_tags(html), "Hello world");
}
#[test]
fn test_xss_img_onerror() {
let html = r#"<img src="x" onerror="alert('XSS')">"#;
let clean = sanitize_html(html);
assert!(!clean.contains("onerror"));
}
#[test]
fn test_xss_javascript_url() {
let html = r#"<a href="javascript:alert('XSS')">Click</a>"#;
let clean = sanitize_html(html);
assert!(!clean.contains("javascript:"));
}
#[test]
fn test_xss_iframe() {
let html = r#"<iframe src="http://evil.com"></iframe>"#;
let clean = sanitize_html(html);
assert!(!clean.contains("iframe"));
}
#[test]
fn test_xss_data_url() {
let html = r#"<a href="data:text/html,<script>alert('XSS')</script>">Click</a>"#;
let clean = sanitize_html(html);
assert!(!clean.contains("data:"));
}
#[test]
fn test_sanitize_empty_string() {
assert_eq!(sanitize_html(""), "");
}
#[test]
fn test_sanitize_plain_text() {
let text = "Plain text with no tags";
assert_eq!(sanitize_html(text), text);
}
#[test]
fn test_decode_entities_no_entities() {
let text = "No entities here";
assert_eq!(decode_entities(text), text);
}
#[test]
fn test_strip_tags_nested() {
let html = "<div><p>Hello <span><b>world</b></span></p></div>";
assert_eq!(strip_tags(html), "Hello world");
}
#[test]
fn test_sanitize_link_rel_attribute() {
let html = r#"<a href="http://example.com">Link</a>"#;
let clean = sanitize_html(html);
assert!(clean.contains("nofollow"));
assert!(clean.contains("noopener"));
assert!(clean.contains("noreferrer"));
}
}