Skip to main content

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