aster_forge_xml/
parser.rs

1//! Event-based XML validation and shared parsing limits.
2
3use std::borrow::Cow;
4
5use quick_xml::Reader;
6use quick_xml::XmlVersion;
7use quick_xml::escape::unescape;
8use quick_xml::events::{BytesStart, Event};
9
10use crate::syntax::{
11    XML_NAMESPACE_URI, map_quick_xml_error_at, split_qualified_name, validate_namespace_binding,
12    validate_qualified_name,
13};
14use crate::{DEFAULT_XML_MAX_DEPTH, Error, XmlSafetyError};
15
16const DEFAULT_MAX_INPUT_BYTES: usize = 10 * 1024 * 1024;
17const DEFAULT_MAX_ELEMENTS: usize = 100_000;
18const DEFAULT_MAX_ATTRIBUTES_PER_ELEMENT: usize = 1_024;
19const DEFAULT_MAX_TEXT_BYTES: usize = 10 * 1024 * 1024;
20const DEFAULT_MAX_EVENTS: usize = 1_000_000;
21
22/// Finite resource and declaration limits applied to untrusted XML.
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub struct XmlSafetyPolicy {
25    pub max_input_bytes: usize,
26    pub max_depth: usize,
27    pub max_elements: usize,
28    pub max_attributes_per_element: usize,
29    pub max_text_bytes: usize,
30    pub max_events: usize,
31    pub reject_doctype: bool,
32}
33
34impl XmlSafetyPolicy {
35    /// A conservative policy suitable for network and storage protocol input.
36    #[must_use]
37    pub const fn untrusted() -> Self {
38        Self {
39            max_input_bytes: DEFAULT_MAX_INPUT_BYTES,
40            max_depth: DEFAULT_XML_MAX_DEPTH,
41            max_elements: DEFAULT_MAX_ELEMENTS,
42            max_attributes_per_element: DEFAULT_MAX_ATTRIBUTES_PER_ELEMENT,
43            max_text_bytes: DEFAULT_MAX_TEXT_BYTES,
44            max_events: DEFAULT_MAX_EVENTS,
45            reject_doctype: true,
46        }
47    }
48
49    pub(crate) fn validate(self) -> Result<(), XmlSafetyError> {
50        if self.max_input_bytes == 0
51            || self.max_depth == 0
52            || self.max_elements == 0
53            || self.max_attributes_per_element == 0
54            || self.max_text_bytes == 0
55            || self.max_events == 0
56        {
57            Err(XmlSafetyError::InvalidPolicy)
58        } else {
59            Ok(())
60        }
61    }
62}
63
64impl Default for XmlSafetyPolicy {
65    fn default() -> Self {
66        Self::untrusted()
67    }
68}
69
70/// Tree parsing behavior. Safety limits remain finite by default.
71#[derive(Debug, Clone, PartialEq, Eq, Default)]
72pub struct ParseOptions {
73    pub safety: XmlSafetyPolicy,
74    /// Drops whitespace-only text nodes and trims retained text nodes.
75    pub trim_whitespace: bool,
76}
77
78impl ParseOptions {
79    #[must_use]
80    pub fn new() -> Self {
81        Self::default()
82    }
83
84    #[must_use]
85    pub fn safety_policy(mut self, policy: XmlSafetyPolicy) -> Self {
86        self.safety = policy;
87        self
88    }
89
90    #[must_use]
91    pub fn max_depth(mut self, value: usize) -> Self {
92        self.safety.max_depth = value;
93        self
94    }
95
96    #[must_use]
97    pub fn max_elements(mut self, value: usize) -> Self {
98        self.safety.max_elements = value;
99        self
100    }
101
102    #[must_use]
103    pub fn max_size(mut self, value: usize) -> Self {
104        self.safety.max_input_bytes = value;
105        self
106    }
107
108    #[must_use]
109    pub fn max_attributes_per_element(mut self, value: usize) -> Self {
110        self.safety.max_attributes_per_element = value;
111        self
112    }
113
114    #[must_use]
115    pub fn max_text_bytes(mut self, value: usize) -> Self {
116        self.safety.max_text_bytes = value;
117        self
118    }
119
120    #[must_use]
121    pub fn max_events(mut self, value: usize) -> Self {
122        self.safety.max_events = value;
123        self
124    }
125
126    #[must_use]
127    pub fn allow_dtd(mut self, allow: bool) -> Self {
128        self.safety.reject_doctype = !allow;
129        self
130    }
131
132    #[must_use]
133    pub fn trim_whitespace(mut self, trim: bool) -> Self {
134        self.trim_whitespace = trim;
135        self
136    }
137}
138
139/// Validates one complete XML document without constructing a DOM.
140///
141/// # Errors
142///
143/// Returns a safety error when `policy` is invalid or the document is malformed, exceeds a limit,
144/// uses forbidden DTD/entity features, or contains invalid encoding.
145pub fn validate_xml_input(bytes: &[u8], policy: XmlSafetyPolicy) -> Result<(), XmlSafetyError> {
146    scan_xml(bytes, &ParseOptions::new().safety_policy(policy))
147        .map(|_| ())
148        .map_err(|error| safety_error(&error))
149}
150
151/// Returns the local name of a validated document root.
152///
153/// # Errors
154///
155/// Returns a safety error when `policy` is invalid, the document has no single valid root, exceeds
156/// a limit, uses forbidden DTD/entity features, or contains invalid encoding.
157pub fn xml_root_local_name(
158    bytes: &[u8],
159    policy: XmlSafetyPolicy,
160) -> Result<String, XmlSafetyError> {
161    scan_xml(bytes, &ParseOptions::new().safety_policy(policy))
162        .map_err(|error| safety_error(&error))?
163        .ok_or(XmlSafetyError::Malformed)
164}
165
166fn safety_error(error: &Error) -> XmlSafetyError {
167    match error {
168        Error::Safety(error) => *error,
169        Error::InvalidXml(_) | Error::InvalidData(_) | Error::Io(_) => XmlSafetyError::Malformed,
170    }
171}
172
173#[derive(Debug)]
174struct Frame {
175    qualified_name: String,
176    binding_start: usize,
177}
178
179#[derive(Debug)]
180struct NamespaceBinding {
181    prefix: String,
182    uri: Option<String>,
183}
184
185#[derive(Debug, Default)]
186struct ScanState {
187    frames: Vec<Frame>,
188    bindings: Vec<NamespaceBinding>,
189    root_name: Option<String>,
190    root_complete: bool,
191    elements: usize,
192    text_bytes: usize,
193    events: usize,
194}
195
196impl ScanState {
197    fn count_event(&mut self, policy: XmlSafetyPolicy) -> Result<(), Error> {
198        self.events = self
199            .events
200            .checked_add(1)
201            .ok_or(XmlSafetyError::TooManyEvents)?;
202        if self.events > policy.max_events {
203            return Err(XmlSafetyError::TooManyEvents.into());
204        }
205        Ok(())
206    }
207
208    fn count_element(&mut self, policy: XmlSafetyPolicy) -> Result<(), Error> {
209        let depth = self
210            .frames
211            .len()
212            .checked_add(1)
213            .ok_or(XmlSafetyError::TooDeep)?;
214        if depth > policy.max_depth {
215            return Err(XmlSafetyError::TooDeep.into());
216        }
217        self.elements = self
218            .elements
219            .checked_add(1)
220            .ok_or(XmlSafetyError::TooManyElements)?;
221        if self.elements > policy.max_elements {
222            return Err(XmlSafetyError::TooManyElements.into());
223        }
224        Ok(())
225    }
226
227    fn count_text(&mut self, text: &str, policy: XmlSafetyPolicy) -> Result<(), Error> {
228        self.text_bytes = self
229            .text_bytes
230            .checked_add(text.len())
231            .ok_or(XmlSafetyError::TextTooLarge)?;
232        if self.text_bytes > policy.max_text_bytes {
233            return Err(XmlSafetyError::TextTooLarge.into());
234        }
235        if self.frames.is_empty() && !text.chars().all(char::is_whitespace) {
236            return Err(XmlSafetyError::Malformed.into());
237        }
238        Ok(())
239    }
240
241    fn namespace(&self, prefix: &str) -> Option<&str> {
242        if prefix == "xml" {
243            return Some(XML_NAMESPACE_URI);
244        }
245        self.bindings
246            .iter()
247            .rev()
248            .find(|binding| binding.prefix == prefix)
249            .and_then(|binding| binding.uri.as_deref())
250    }
251}
252
253fn scan_xml(bytes: &[u8], options: &ParseOptions) -> Result<Option<String>, Error> {
254    options.safety.validate()?;
255    if bytes.len() > options.safety.max_input_bytes {
256        return Err(XmlSafetyError::InputTooLarge.into());
257    }
258
259    let mut reader = Reader::from_reader(bytes);
260    reader.config_mut().trim_text(false);
261    let mut state = ScanState::default();
262
263    loop {
264        let event = reader.read_event().map_err(|error| {
265            let error_position = usize::try_from(reader.error_position()).unwrap_or(bytes.len());
266            map_quick_xml_error_at(error, error_position, bytes, options.safety.reject_doctype)
267        })?;
268        if !matches!(event, Event::Eof) {
269            state.count_event(options.safety)?;
270        }
271        match event {
272            Event::Start(start) => {
273                state.count_element(options.safety)?;
274                let frame = scan_element(&reader, &mut state, &start, options.safety)?;
275                state.frames.push(frame);
276            }
277            Event::Empty(start) => {
278                state.count_element(options.safety)?;
279                let frame = scan_element(&reader, &mut state, &start, options.safety)?;
280                state.bindings.truncate(frame.binding_start);
281                if state.frames.is_empty() {
282                    state.root_complete = true;
283                }
284            }
285            Event::End(end) => {
286                let end_name = end.name();
287                let qualified_name = end_name.as_ref();
288                let frame = state.frames.pop().ok_or(XmlSafetyError::Malformed)?;
289                if frame.qualified_name != qualified_name {
290                    return Err(XmlSafetyError::Malformed.into());
291                }
292                state.bindings.truncate(frame.binding_start);
293                if state.frames.is_empty() {
294                    state.root_complete = true;
295                }
296            }
297            Event::Text(text) => {
298                let raw = text.as_ref();
299                let value = unescape(raw).map_err(|error| Error::InvalidXml(error.to_string()))?;
300                state.count_text(value.as_ref(), options.safety)?;
301            }
302            Event::CData(text) => state.count_text(text.as_ref(), options.safety)?,
303            Event::GeneralRef(reference) => {
304                let value = decode_reference(&reference)?;
305                state.count_text(value.as_ref(), options.safety)?;
306            }
307            Event::Decl(_) => {
308                if state.root_name.is_some() || !state.frames.is_empty() || state.root_complete {
309                    return Err(XmlSafetyError::Malformed.into());
310                }
311            }
312            Event::DocType(_) => {
313                if options.safety.reject_doctype {
314                    return Err(XmlSafetyError::ExternalEntity.into());
315                }
316                if state.root_name.is_some() || !state.frames.is_empty() || state.root_complete {
317                    return Err(XmlSafetyError::Malformed.into());
318                }
319            }
320            Event::Comment(_) | Event::PI(_) => {}
321            Event::Eof => {
322                if !state.frames.is_empty() || !state.root_complete {
323                    return Err(XmlSafetyError::Malformed.into());
324                }
325                return Ok(state.root_name);
326            }
327        }
328    }
329}
330
331fn scan_element(
332    _reader: &Reader<&[u8]>,
333    state: &mut ScanState,
334    start: &BytesStart<'_>,
335    policy: XmlSafetyPolicy,
336) -> Result<Frame, Error> {
337    if state.frames.is_empty() && state.root_complete {
338        return Err(XmlSafetyError::Malformed.into());
339    }
340    let start_name = start.name();
341    let qualified_name = start_name.as_ref();
342    let (prefix, local_name) = validate_qualified_name(qualified_name)?;
343    let binding_start = state.bindings.len();
344    let mut attribute_count = 0usize;
345
346    for attribute in start.attributes() {
347        attribute_count = attribute_count
348            .checked_add(1)
349            .ok_or(XmlSafetyError::TooManyAttributes)?;
350        if attribute_count > policy.max_attributes_per_element {
351            return Err(XmlSafetyError::TooManyAttributes.into());
352        }
353        let attribute = attribute.map_err(|error| Error::InvalidXml(error.to_string()))?;
354        let name = attribute.key.as_ref();
355        validate_qualified_name(name)?;
356        if name == "xmlns" || name.starts_with("xmlns:") {
357            let namespace_prefix = name.strip_prefix("xmlns:").unwrap_or("");
358            let uri = attribute
359                .normalized_value(XmlVersion::Explicit1_0)
360                .map_err(|error| Error::InvalidXml(error.to_string()))?;
361            validate_namespace_binding(namespace_prefix, &uri)?;
362            state.bindings.push(NamespaceBinding {
363                prefix: namespace_prefix.to_owned(),
364                uri: (!uri.is_empty()).then(|| uri.into_owned()),
365            });
366        }
367    }
368
369    if let Some(prefix) = prefix
370        && state.namespace(prefix).is_none()
371    {
372        return Err(XmlSafetyError::Malformed.into());
373    }
374    for attribute in start.attributes() {
375        let attribute = attribute.map_err(|error| Error::InvalidXml(error.to_string()))?;
376        let name = attribute.key.as_ref();
377        if name == "xmlns" || name.starts_with("xmlns:") {
378            continue;
379        }
380        let (prefix, _) = split_qualified_name(name);
381        if let Some(prefix) = prefix
382            && prefix != "xml"
383            && state.namespace(prefix).is_none()
384        {
385            return Err(XmlSafetyError::Malformed.into());
386        }
387    }
388
389    if state.root_name.is_none() {
390        state.root_name = Some(local_name.to_owned());
391    }
392    Ok(Frame {
393        qualified_name: qualified_name.to_owned(),
394        binding_start,
395    })
396}
397
398fn decode_reference<'a>(
399    reference: &quick_xml::events::BytesRef<'a>,
400) -> Result<Cow<'a, str>, Error> {
401    if let Some(character) = reference
402        .resolve_char_ref()
403        .map_err(|error| Error::InvalidXml(error.to_string()))?
404    {
405        return Ok(Cow::Owned(character.to_string()));
406    }
407    Ok(Cow::Borrowed(match reference.as_ref() {
408        "amp" => "&",
409        "lt" => "<",
410        "gt" => ">",
411        "apos" => "'",
412        "quot" => "\"",
413        _ => return Err(XmlSafetyError::ExternalEntity.into()),
414    }))
415}