use crate::error::SamlError;
use quick_xml::escape::resolve_predefined_entity;
use quick_xml::events::{BytesRef, Event};
use quick_xml::name::QName;
use quick_xml::{Reader, XmlVersion};
pub const DEFAULT_XML_MAX_BYTES: usize = 1024 * 1024;
pub const DEFAULT_XML_MAX_DEPTH: usize = 1024;
pub const DEFAULT_XML_MAX_NODES: usize = 50_000;
pub const DEFAULT_XML_MAX_ATTRIBUTES_PER_ELEMENT: usize = 64;
pub const DEFAULT_XML_MAX_ATTRIBUTE_VALUE_BYTES: usize = 16 * 1024;
pub const DEFAULT_XML_MAX_TEXT_BYTES: usize = DEFAULT_XML_MAX_BYTES;
const XML_LIMIT_EXCEEDED: &str = "ERR_XML_LIMIT_EXCEEDED";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct XmlLimits {
pub max_bytes: usize,
pub max_depth: usize,
pub max_nodes: usize,
pub max_attributes_per_element: usize,
pub max_attribute_value_bytes: usize,
pub max_text_bytes: usize,
}
impl XmlLimits {
pub const fn unbounded() -> Self {
Self {
max_bytes: usize::MAX,
max_depth: usize::MAX,
max_nodes: usize::MAX,
max_attributes_per_element: usize::MAX,
max_attribute_value_bytes: usize::MAX,
max_text_bytes: usize::MAX,
}
}
pub(crate) fn check_input_bytes(self, len: usize) -> Result<(), SamlError> {
if len > self.max_bytes {
return Err(limit_exceeded("max XML bytes", self.max_bytes));
}
Ok(())
}
}
impl Default for XmlLimits {
fn default() -> Self {
Self {
max_bytes: DEFAULT_XML_MAX_BYTES,
max_depth: DEFAULT_XML_MAX_DEPTH,
max_nodes: DEFAULT_XML_MAX_NODES,
max_attributes_per_element: DEFAULT_XML_MAX_ATTRIBUTES_PER_ELEMENT,
max_attribute_value_bytes: DEFAULT_XML_MAX_ATTRIBUTE_VALUE_BYTES,
max_text_bytes: DEFAULT_XML_MAX_TEXT_BYTES,
}
}
}
#[derive(Debug, Clone)]
pub struct Node {
pub local_name: String,
pub attrs: Vec<(String, String)>,
pub children: Vec<Node>,
pub text: String,
pub start: usize,
pub end: usize,
}
impl Node {
pub fn attr(&self, name: &str) -> Option<&str> {
self.attrs
.iter()
.find(|(k, _)| k == name)
.map(|(_, v)| v.as_str())
}
}
#[derive(Debug, Clone)]
pub struct Document {
pub root: Node,
}
fn local_name_str(name: QName) -> String {
String::from_utf8_lossy(name.local_name().as_ref()).into_owned()
}
fn limit_exceeded(limit: &str, max: usize) -> SamlError {
SamlError::Invalid(format!("{XML_LIMIT_EXCEEDED}: {limit} exceeded ({max})"))
}
fn checked_node_count(count: &mut usize, limits: XmlLimits) -> Result<(), SamlError> {
*count = count
.checked_add(1)
.ok_or_else(|| limit_exceeded("max XML nodes", limits.max_nodes))?;
if *count > limits.max_nodes {
return Err(limit_exceeded("max XML nodes", limits.max_nodes));
}
Ok(())
}
fn check_depth(depth: usize, limits: XmlLimits) -> Result<(), SamlError> {
if depth > limits.max_depth {
return Err(limit_exceeded("max XML depth", limits.max_depth));
}
Ok(())
}
fn checked_append_text(
target: &mut String,
text: &str,
limits: XmlLimits,
) -> Result<(), SamlError> {
let next_len = target
.len()
.checked_add(text.len())
.ok_or_else(|| limit_exceeded("max XML text bytes", limits.max_text_bytes))?;
if next_len > limits.max_text_bytes {
return Err(limit_exceeded("max XML text bytes", limits.max_text_bytes));
}
target.push_str(text);
Ok(())
}
fn read_attrs(
e: &quick_xml::events::BytesStart,
limits: XmlLimits,
) -> Result<Vec<(String, String)>, SamlError> {
let mut out = Vec::new();
for attr in e.attributes() {
if out.len() >= limits.max_attributes_per_element {
return Err(limit_exceeded(
"max XML attributes per element",
limits.max_attributes_per_element,
));
}
let attr = attr.map_err(|err| SamlError::Xml(err.to_string()))?;
let key = local_name_str(attr.key);
let value = attr
.decoded_and_normalized_value(XmlVersion::Implicit1_0, e.decoder())
.map_err(|err| SamlError::Xml(err.to_string()))?
.into_owned();
if value.len() > limits.max_attribute_value_bytes {
return Err(limit_exceeded(
"max XML attribute value bytes",
limits.max_attribute_value_bytes,
));
}
out.push((key, value));
}
Ok(out)
}
fn find_lt(bytes: &[u8], before: usize, after: usize) -> usize {
bytes[before..after]
.iter()
.position(|&b| b == b'<')
.map(|p| before + p)
.unwrap_or(before)
}
fn push_child(stack: &mut [Node], roots: &mut Vec<Node>, node: Node) {
match stack.last_mut() {
Some(parent) => parent.children.push(node),
None => roots.push(node),
}
}
#[derive(Clone, Copy)]
enum ParseMode {
Document,
Roots,
}
fn is_xml_whitespace(text: &str) -> bool {
text.bytes()
.all(|byte| matches!(byte, b' ' | b'\t' | b'\r' | b'\n'))
}
fn reject_non_document_content(mode: ParseMode) -> Result<(), SamlError> {
match mode {
ParseMode::Document => Err(SamlError::Xml(
"content is not allowed outside the document element".into(),
)),
ParseMode::Roots => Ok(()),
}
}
fn push_general_ref(node: &mut Node, e: BytesRef, limits: XmlLimits) -> Result<(), SamlError> {
if let Some(ch) = e
.resolve_char_ref()
.map_err(|err| SamlError::Xml(err.to_string()))?
{
let mut buf = [0; 4];
checked_append_text(&mut node.text, ch.encode_utf8(&mut buf), limits)?;
return Ok(());
}
let entity = e.decode().map_err(|err| SamlError::Xml(err.to_string()))?;
let resolved = resolve_predefined_entity(&entity)
.ok_or_else(|| SamlError::Xml(format!("unrecognized entity `{entity}`")))?;
checked_append_text(&mut node.text, resolved, limits)?;
Ok(())
}
pub fn parse(xml: &str) -> Result<Document, SamlError> {
parse_with_limits(xml, XmlLimits::default())
}
pub fn parse_with_limits(xml: &str, limits: XmlLimits) -> Result<Document, SamlError> {
let mut roots = parse_roots_inner(xml, limits, ParseMode::Document)?;
let root = match roots.len() {
0 => return Err(SamlError::Xml("no document element".into())),
1 => roots
.pop()
.ok_or_else(|| SamlError::Xml("no document element".into()))?,
_ => return Err(SamlError::Xml("multiple document elements".into())),
};
Ok(Document { root })
}
pub fn parse_roots(xml: &str) -> Result<Vec<Node>, SamlError> {
parse_roots_with_limits(xml, XmlLimits::default())
}
pub fn parse_roots_with_limits(xml: &str, limits: XmlLimits) -> Result<Vec<Node>, SamlError> {
parse_roots_inner(xml, limits, ParseMode::Roots)
}
fn parse_roots_inner(
xml: &str,
limits: XmlLimits,
mode: ParseMode,
) -> Result<Vec<Node>, SamlError> {
limits.check_input_bytes(xml.len())?;
let mut reader = Reader::from_str(xml);
let bytes = xml.as_bytes();
let mut stack: Vec<Node> = Vec::new();
let mut roots: Vec<Node> = Vec::new();
let mut node_count = 0usize;
loop {
let before = reader.buffer_position() as usize;
let event = reader
.read_event()
.map_err(|err| SamlError::Xml(err.to_string()))?;
let after = reader.buffer_position() as usize;
match event {
Event::Start(e) => {
checked_node_count(&mut node_count, limits)?;
check_depth(stack.len().saturating_add(1), limits)?;
let start = find_lt(bytes, before, after);
stack.push(Node {
local_name: local_name_str(e.name()),
attrs: read_attrs(&e, limits)?,
children: Vec::new(),
text: String::new(),
start,
end: 0,
});
}
Event::Empty(e) => {
checked_node_count(&mut node_count, limits)?;
check_depth(stack.len().saturating_add(1), limits)?;
let start = find_lt(bytes, before, after);
let node = Node {
local_name: local_name_str(e.name()),
attrs: read_attrs(&e, limits)?,
children: Vec::new(),
text: String::new(),
start,
end: after,
};
push_child(&mut stack, &mut roots, node);
}
Event::End(_) => {
if let Some(mut node) = stack.pop() {
node.end = after;
push_child(&mut stack, &mut roots, node);
}
}
Event::Text(e) => {
let text = e.decode().map_err(|err| SamlError::Xml(err.to_string()))?;
if let Some(top) = stack.last_mut() {
checked_append_text(&mut top.text, &text, limits)?;
} else if !is_xml_whitespace(&text) {
reject_non_document_content(mode)?;
}
}
Event::CData(e) => {
if let Some(top) = stack.last_mut() {
let inner = e.into_inner();
let text = String::from_utf8_lossy(&inner);
checked_append_text(&mut top.text, &text, limits)?;
} else {
reject_non_document_content(mode)?;
}
}
Event::GeneralRef(e) => {
if let Some(top) = stack.last_mut() {
push_general_ref(top, e, limits)?;
} else {
reject_non_document_content(mode)?;
}
}
Event::DocType(_) => {
return Err(SamlError::Xml("DOCTYPE is not allowed".into()));
}
Event::Eof => break,
_ => {}
}
}
Ok(roots)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_doctype() {
let xml = "<!DOCTYPE foo [<!ENTITY x \"y\">]><foo/>";
assert!(parse(xml).is_err());
}
#[test]
fn parses_root_and_attrs() -> Result<(), SamlError> {
let doc = parse("<a:Root xmlns:a=\"urn:x\" id=\"1\"><b>hi</b></a:Root>")?;
assert_eq!(doc.root.local_name, "Root");
assert_eq!(doc.root.attr("id"), Some("1"));
assert_eq!(doc.root.children.len(), 1);
assert_eq!(doc.root.children[0].text, "hi");
Ok(())
}
#[test]
fn parses_escaped_attribute_and_text_values() -> Result<(), SamlError> {
let doc = parse("<Root value=\"one & two\"><Child>three < four</Child></Root>")?;
assert_eq!(doc.root.attr("value"), Some("one & two"));
assert_eq!(doc.root.children[0].text, "three < four");
Ok(())
}
#[test]
fn parse_rejects_multiple_document_elements() {
let result = parse("<Root/><Second/>");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn parse_rejects_text_before_document_element() {
let result = parse("outside<Root/>");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn parse_rejects_text_after_document_element() {
let result = parse("<Root/>outside");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn parse_rejects_cdata_outside_document_element() {
let result = parse("<Root/><![CDATA[outside]]>");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn parse_rejects_reference_outside_document_element() {
let result = parse("<Root/>&");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn parse_accepts_xml_misc_around_document_element() -> Result<(), SamlError> {
let xml = concat!(
"<?xml version=\"1.0\"?>\n",
"<!-- before --><?before allowed?>\n",
"<Root/>\n",
"<?after allowed?><!-- after -->",
);
let document = parse(xml)?;
assert_eq!(document.root.local_name, "Root");
Ok(())
}
#[test]
fn parse_roots_preserves_multiple_root_collection() -> Result<(), SamlError> {
let roots = parse_roots("<Root/><Second/>")?;
assert_eq!(roots.len(), 2);
Ok(())
}
fn limit_hit<T>(result: Result<T, SamlError>) -> bool {
matches!(result, Err(SamlError::Invalid(message)) if message.contains(XML_LIMIT_EXCEEDED))
}
#[test]
fn rejects_xml_over_byte_limit() {
let limits = XmlLimits {
max_bytes: 4,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<Root/>", limits)));
}
#[test]
fn rejects_xml_over_depth_limit() {
let limits = XmlLimits {
max_depth: 2,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<a><b><c/></b></a>", limits)));
}
#[test]
fn rejects_xml_over_node_limit() {
let limits = XmlLimits {
max_nodes: 2,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<a><b/><c/></a>", limits)));
}
#[test]
fn rejects_xml_over_attribute_count_limit() {
let limits = XmlLimits {
max_attributes_per_element: 1,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<a x=\"1\" y=\"2\"/>", limits)));
}
#[test]
fn rejects_xml_over_default_attribute_count_limit() {
let mut xml = String::from("<samlp:Response");
for index in 0..=DEFAULT_XML_MAX_ATTRIBUTES_PER_ELEMENT {
xml.push_str(&format!(" attr{index}=\"value{index}\""));
}
xml.push_str("/>");
assert!(limit_hit(parse(&xml)));
}
#[test]
fn rejects_xml_over_default_namespace_declaration_count_limit() {
let mut xml = String::from("<samlp:Response");
for index in 0..=DEFAULT_XML_MAX_ATTRIBUTES_PER_ELEMENT {
xml.push_str(&format!(" xmlns:p{index}=\"urn:test:{index}\""));
}
xml.push_str("/>");
assert!(limit_hit(parse(&xml)));
}
#[test]
fn duplicate_attribute_returns_xml_error() {
let result = parse("<samlp:Response ID=\"first\" ID=\"second\"/>");
assert!(matches!(result, Err(SamlError::Xml(_))));
}
#[test]
fn rejects_xml_over_attribute_value_limit() {
let limits = XmlLimits {
max_attribute_value_bytes: 3,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<a x=\"1234\"/>", limits)));
}
#[test]
fn rejects_xml_over_text_limit() {
let limits = XmlLimits {
max_text_bytes: 3,
..Default::default()
};
assert!(limit_hit(parse_with_limits("<a>1234</a>", limits)));
}
}