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, 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#[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 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#[derive(Debug, Clone, PartialEq, Eq, Default)]
71pub struct ParseOptions {
72 pub safety: XmlSafetyPolicy,
73 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
128pub 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
135pub 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}