1use 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#[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 #[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#[derive(Debug, Clone, PartialEq, Eq, Default)]
72pub struct ParseOptions {
73 pub safety: XmlSafetyPolicy,
74 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
139pub 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
151pub 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}