Skip to main content

aster_forge_xml/
stream.rs

1//! Bounded, namespace-aware streaming XML reader.
2
3use std::borrow::Cow;
4use std::io::{BufRead, Take, Write};
5
6use aster_forge_utils::numbers::usize_to_u64;
7use quick_xml::XmlVersion;
8use quick_xml::encoding::Decoder;
9use quick_xml::escape::unescape;
10use quick_xml::events::attributes::{Attribute, Attributes as QuickAttributes};
11use quick_xml::events::{BytesCData, BytesEnd, BytesPI, BytesStart, BytesText, Event};
12use quick_xml::name::{NamespaceResolver, PrefixDeclaration, ResolveResult};
13use quick_xml::reader::NsReader;
14use quick_xml::writer::Writer;
15
16use crate::syntax::{map_quick_xml_error, utf8};
17use crate::{Error, ValidatedXml, XmlSafetyError, XmlSafetyPolicy};
18
19/// A namespace-resolved XML name borrowed from one streaming event.
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub struct StreamName<'a> {
22    qualified: &'a str,
23    local: &'a str,
24    namespace: Option<&'a str>,
25}
26
27impl<'a> StreamName<'a> {
28    pub fn qualified(self) -> &'a str {
29        self.qualified
30    }
31
32    pub fn local(self) -> &'a str {
33        self.local
34    }
35
36    pub fn namespace(self) -> Option<&'a str> {
37        self.namespace
38    }
39
40    pub fn matches(self, local: &str, namespace: Option<&str>) -> bool {
41        self.local == local && self.namespace == namespace
42    }
43}
44
45/// A start or empty element event.
46pub struct StreamStart<'a> {
47    raw: BytesStart<'a>,
48    namespace: Option<&'a str>,
49    resolver: &'a NamespaceResolver,
50    decoder: Decoder,
51    cached_attribute_values: &'a [CachedAttributeValue],
52}
53
54impl StreamStart<'_> {
55    pub fn name(&self) -> Result<StreamName<'_>, Error> {
56        let qualified = utf8(self.raw.name().into_inner())?;
57        let local = utf8(self.raw.local_name().into_inner())?;
58        Ok(StreamName {
59            qualified,
60            local,
61            namespace: self.namespace,
62        })
63    }
64
65    pub fn attributes(&self) -> StreamAttributes<'_> {
66        StreamAttributes {
67            inner: self.raw.attributes(),
68            resolver: self.resolver,
69            decoder: self.decoder,
70            cached_values: self.cached_attribute_values,
71            cached_index: 0,
72            index: 0,
73        }
74    }
75
76    pub fn attribute(&self, qualified_name: &str) -> Result<Option<Cow<'_, str>>, Error> {
77        for attribute in self.attributes() {
78            let attribute = attribute?;
79            if attribute.name()?.qualified() == qualified_name {
80                return attribute.into_value().map(Some);
81            }
82        }
83        Ok(None)
84    }
85
86    pub fn attribute_ns(
87        &self,
88        local: &str,
89        namespace: Option<&str>,
90    ) -> Result<Option<Cow<'_, str>>, Error> {
91        for attribute in self.attributes() {
92            let attribute = attribute?;
93            if attribute.name()?.matches(local, namespace) {
94                return attribute.into_value().map(Some);
95            }
96        }
97        Ok(None)
98    }
99}
100
101/// Iterator over attributes of a streaming start event.
102pub struct StreamAttributes<'a> {
103    inner: QuickAttributes<'a>,
104    resolver: &'a NamespaceResolver,
105    decoder: Decoder,
106    cached_values: &'a [CachedAttributeValue],
107    cached_index: usize,
108    index: usize,
109}
110
111impl<'a> Iterator for StreamAttributes<'a> {
112    type Item = Result<StreamAttribute<'a>, Error>;
113
114    fn next(&mut self) -> Option<Self::Item> {
115        let index = self.index;
116        self.index = self.index.saturating_add(1);
117        let cached_value = self
118            .cached_values
119            .get(self.cached_index)
120            .and_then(|cached| {
121                if cached.index == index {
122                    self.cached_index += 1;
123                    Some(cached.value.as_str())
124                } else {
125                    None
126                }
127            });
128        self.inner.next().map(|attribute| {
129            attribute
130                .map(|raw| StreamAttribute {
131                    raw,
132                    resolver: self.resolver,
133                    decoder: self.decoder,
134                    cached_value,
135                })
136                .map_err(|error| Error::InvalidXml(error.to_string()))
137        })
138    }
139}
140
141/// A namespace-resolved attribute borrowed from a streaming start event.
142pub struct StreamAttribute<'a> {
143    raw: Attribute<'a>,
144    resolver: &'a NamespaceResolver,
145    decoder: Decoder,
146    cached_value: Option<&'a str>,
147}
148
149impl<'a> StreamAttribute<'a> {
150    pub fn name(&self) -> Result<StreamName<'_>, Error> {
151        let qualified = utf8(self.raw.key.into_inner())?;
152        let local = utf8(self.raw.key.local_name().into_inner())?;
153        let namespace = resolve_namespace(
154            self.resolver.resolve_attribute(self.raw.key).0,
155            "attribute namespace",
156        )?;
157        Ok(StreamName {
158            qualified,
159            local,
160            namespace,
161        })
162    }
163
164    pub fn value(&self) -> Result<Cow<'_, str>, Error> {
165        if let Some(value) = self.cached_value {
166            return Ok(Cow::Borrowed(value));
167        }
168        self.raw
169            .decoded_and_normalized_value(XmlVersion::Explicit1_0, self.decoder)
170            .map_err(|error| Error::InvalidXml(error.to_string()))
171    }
172
173    pub fn into_value(self) -> Result<Cow<'a, str>, Error> {
174        if let Some(value) = self.cached_value {
175            return Ok(Cow::Borrowed(value));
176        }
177        self.raw
178            .decoded_and_normalized_value(XmlVersion::Explicit1_0, self.decoder)
179            .map_err(|error| Error::InvalidXml(error.to_string()))
180    }
181}
182
183/// An end element event.
184#[derive(Debug)]
185pub struct StreamEnd<'a> {
186    raw: BytesEnd<'a>,
187    namespace: Option<&'a str>,
188}
189
190impl StreamEnd<'_> {
191    pub fn name(&self) -> Result<StreamName<'_>, Error> {
192        let qualified = utf8(self.raw.name().into_inner())?;
193        let local = utf8(self.raw.local_name().into_inner())?;
194        Ok(StreamName {
195            qualified,
196            local,
197            namespace: self.namespace,
198        })
199    }
200}
201
202/// Decoded and unescaped character data.
203pub struct StreamText<'a> {
204    value: Cow<'a, str>,
205}
206
207impl StreamText<'_> {
208    pub fn value(&self) -> &str {
209        &self.value
210    }
211}
212
213/// Decoded CDATA content.
214pub struct StreamCData<'a> {
215    value: Cow<'a, str>,
216}
217
218impl StreamCData<'_> {
219    pub fn value(&self) -> &str {
220        &self.value
221    }
222}
223
224/// Decoded XML comment content.
225pub struct StreamComment<'a> {
226    value: Cow<'a, str>,
227}
228
229impl StreamComment<'_> {
230    pub fn value(&self) -> &str {
231        &self.value
232    }
233}
234
235/// A processing instruction.
236pub struct StreamProcessingInstruction<'a> {
237    raw: BytesPI<'a>,
238}
239
240impl StreamProcessingInstruction<'_> {
241    pub fn target(&self) -> Result<&str, Error> {
242        utf8(self.raw.target())
243    }
244
245    pub fn content(&self) -> Result<Option<&str>, Error> {
246        let content = utf8(self.raw.content())?
247            .trim_start_matches(|character: char| character.is_ascii_whitespace());
248        Ok((!content.is_empty()).then_some(content))
249    }
250}
251
252/// One bounded streaming XML event.
253pub enum XmlStreamEvent<'a> {
254    Start(StreamStart<'a>),
255    Empty(StreamStart<'a>),
256    End(StreamEnd<'a>),
257    Text(StreamText<'a>),
258    CData(StreamCData<'a>),
259    Comment(StreamComment<'a>),
260    ProcessingInstruction(StreamProcessingInstruction<'a>),
261    Declaration,
262    DocType,
263    Eof,
264}
265
266struct StreamState {
267    policy: XmlSafetyPolicy,
268    max_input_bytes_u64: u64,
269    depth: usize,
270    elements: usize,
271    text_bytes: usize,
272    events: usize,
273    root_seen: bool,
274    root_complete: bool,
275    current_start_available: bool,
276    current_start_depth: usize,
277    finished: bool,
278}
279
280/// A streaming XML reader that enforces [`XmlSafetyPolicy`] without retaining a full document.
281pub struct XmlStreamReader<R: BufRead> {
282    reader: NsReader<Take<R>>,
283    buffer: Vec<u8>,
284    cached_attribute_values: Vec<CachedAttributeValue>,
285    state: StreamState,
286}
287
288struct CachedAttributeValue {
289    index: usize,
290    value: String,
291}
292
293impl<R: BufRead> XmlStreamReader<R> {
294    pub fn new(reader: R, policy: XmlSafetyPolicy) -> Result<Self, Error> {
295        policy.validate()?;
296        let read_limit = policy.max_input_bytes.saturating_add(1);
297        let read_limit = usize_to_u64(read_limit, "XML stream byte limit").unwrap_or(u64::MAX);
298        let max_input_bytes_u64 =
299            usize_to_u64(policy.max_input_bytes, "XML stream byte limit").unwrap_or(u64::MAX);
300        let mut reader = NsReader::from_reader(reader.take(read_limit));
301        reader.config_mut().trim_text(false);
302        reader
303            .resolver_mut()
304            .set_max_declarations_per_element(policy.max_attributes_per_element);
305        Ok(Self {
306            reader,
307            buffer: Vec::new(),
308            cached_attribute_values: Vec::new(),
309            state: StreamState {
310                policy,
311                max_input_bytes_u64,
312                depth: 0,
313                elements: 0,
314                text_bytes: 0,
315                events: 0,
316                root_seen: false,
317                root_complete: false,
318                current_start_available: false,
319                current_start_depth: 0,
320                finished: false,
321            },
322        })
323    }
324
325    pub fn read_event(&mut self) -> Result<XmlStreamEvent<'_>, Error> {
326        if self.state.finished {
327            return Ok(XmlStreamEvent::Eof);
328        }
329        self.state.current_start_available = false;
330        self.buffer.clear();
331        self.cached_attribute_values.clear();
332        let event = self
333            .reader
334            .read_event_into(&mut self.buffer)
335            .map_err(map_quick_xml_error)?;
336        check_stream_position(&self.reader, &self.state)?;
337        count_event(&mut self.state, &event)?;
338
339        match event {
340            Event::Start(start) => {
341                let namespace = begin_element(
342                    &mut self.state,
343                    &self.reader,
344                    &start,
345                    false,
346                    &mut self.cached_attribute_values,
347                )?;
348                Ok(XmlStreamEvent::Start(StreamStart {
349                    raw: start,
350                    namespace,
351                    resolver: self.reader.resolver(),
352                    decoder: self.reader.decoder(),
353                    cached_attribute_values: &self.cached_attribute_values,
354                }))
355            }
356            Event::Empty(start) => {
357                let namespace = begin_element(
358                    &mut self.state,
359                    &self.reader,
360                    &start,
361                    true,
362                    &mut self.cached_attribute_values,
363                )?;
364                Ok(XmlStreamEvent::Empty(StreamStart {
365                    raw: start,
366                    namespace,
367                    resolver: self.reader.resolver(),
368                    decoder: self.reader.decoder(),
369                    cached_attribute_values: &self.cached_attribute_values,
370                }))
371            }
372            Event::End(end) => {
373                if self.state.depth == 0 {
374                    return Err(XmlSafetyError::Malformed.into());
375                }
376                let namespace = resolve_namespace(
377                    self.reader.resolver().resolve_element(end.name()).0,
378                    "element namespace",
379                )?;
380                utf8(end.name().into_inner())?;
381                self.state.depth -= 1;
382                if self.state.depth == 0 {
383                    self.state.root_complete = true;
384                }
385                Ok(XmlStreamEvent::End(StreamEnd {
386                    raw: end,
387                    namespace,
388                }))
389            }
390            Event::Text(text) => {
391                let value = decode_text(&text)?;
392                count_text(&mut self.state, &value)?;
393                Ok(XmlStreamEvent::Text(StreamText { value }))
394            }
395            Event::CData(cdata) => {
396                let value = cdata
397                    .decode()
398                    .map_err(|_| XmlSafetyError::InvalidEncoding)?;
399                count_text(&mut self.state, &value)?;
400                Ok(XmlStreamEvent::CData(StreamCData { value }))
401            }
402            Event::Comment(comment) => {
403                let value = comment
404                    .decode()
405                    .map_err(|_| XmlSafetyError::InvalidEncoding)?;
406                Ok(XmlStreamEvent::Comment(StreamComment { value }))
407            }
408            Event::PI(pi) => {
409                utf8(pi.target())?;
410                utf8(pi.content())?;
411                Ok(XmlStreamEvent::ProcessingInstruction(
412                    StreamProcessingInstruction { raw: pi },
413                ))
414            }
415            Event::GeneralRef(reference) => {
416                let value = decode_reference(&reference)?;
417                count_text(&mut self.state, &value)?;
418                Ok(XmlStreamEvent::Text(StreamText { value }))
419            }
420            Event::Decl(_) => {
421                if self.state.root_seen || self.state.depth != 0 || self.state.root_complete {
422                    return Err(XmlSafetyError::Malformed.into());
423                }
424                Ok(XmlStreamEvent::Declaration)
425            }
426            Event::DocType(_) => {
427                if self.state.policy.reject_doctype {
428                    return Err(XmlSafetyError::ExternalEntity.into());
429                }
430                if self.state.root_seen || self.state.depth != 0 || self.state.root_complete {
431                    return Err(XmlSafetyError::Malformed.into());
432                }
433                Ok(XmlStreamEvent::DocType)
434            }
435            Event::Eof => {
436                if self.state.depth != 0 || !self.state.root_complete {
437                    return Err(XmlSafetyError::Malformed.into());
438                }
439                self.state.finished = true;
440                Ok(XmlStreamEvent::Eof)
441            }
442        }
443    }
444
445    /// Reads direct text and CDATA until the end of the current start element.
446    pub fn read_text_current(&mut self) -> Result<String, Error> {
447        self.require_current_start()?;
448        self.state.current_start_available = false;
449        let mut output = String::new();
450        loop {
451            match self.read_event()? {
452                XmlStreamEvent::Text(text) => output.push_str(text.value()),
453                XmlStreamEvent::CData(cdata) => output.push_str(cdata.value()),
454                XmlStreamEvent::Comment(_) | XmlStreamEvent::ProcessingInstruction(_) => {}
455                XmlStreamEvent::End(_) => return Ok(output),
456                XmlStreamEvent::Start(_) | XmlStreamEvent::Empty(_) => {
457                    return Err(Error::InvalidXml(
458                        "text helper encountered a nested element".into(),
459                    ));
460                }
461                XmlStreamEvent::Declaration | XmlStreamEvent::DocType | XmlStreamEvent::Eof => {
462                    return Err(XmlSafetyError::Malformed.into());
463                }
464            }
465        }
466    }
467
468    /// Skips the current start element and all descendants with constant retained memory.
469    pub fn skip_current(&mut self) -> Result<(), Error> {
470        self.require_current_start()?;
471        self.state.current_start_available = false;
472        let mut nested = 1usize;
473        while nested > 0 {
474            match self.read_event()? {
475                XmlStreamEvent::Start(_) => {
476                    nested = nested.checked_add(1).ok_or(XmlSafetyError::TooDeep)?;
477                }
478                XmlStreamEvent::End(_) => nested -= 1,
479                XmlStreamEvent::Eof => return Err(XmlSafetyError::Malformed.into()),
480                _ => {}
481            }
482        }
483        Ok(())
484    }
485
486    /// Materializes only the current subtree as a validated owned XML value.
487    pub fn capture_current(&mut self, max_bytes: usize) -> Result<ValidatedXml, Error> {
488        self.require_current_start()?;
489        if max_bytes == 0 {
490            return Err(XmlSafetyError::InvalidPolicy.into());
491        }
492        self.state.current_start_available = false;
493        let mut event_reader = quick_xml::Reader::from_reader(self.buffer.as_slice());
494        let Event::Start(mut captured_start) =
495            event_reader.read_event().map_err(map_quick_xml_error)?
496        else {
497            return Err(Error::InvalidData(
498                "stream start buffer is incomplete".into(),
499            ));
500        };
501        for (prefix, namespace) in self.reader.resolver().bindings() {
502            let already_declared = captured_start.attributes().any(|attribute| {
503                let Ok(attribute) = attribute else {
504                    return false;
505                };
506                match prefix {
507                    PrefixDeclaration::Default => attribute.key.as_ref() == b"xmlns",
508                    PrefixDeclaration::Named(prefix) => {
509                        attribute.key.as_ref().strip_prefix(b"xmlns:") == Some(prefix)
510                    }
511                }
512            });
513            if already_declared {
514                continue;
515            }
516            let namespace = utf8(namespace.into_inner())?;
517            match prefix {
518                PrefixDeclaration::Default => captured_start.push_attribute(("xmlns", namespace)),
519                PrefixDeclaration::Named(prefix) => {
520                    let prefix = utf8(prefix)?;
521                    let name = format!("xmlns:{prefix}");
522                    captured_start.push_attribute((name.as_str(), namespace));
523                }
524            }
525        }
526        let mut writer = Writer::new(LimitedVec::new(Vec::new(), max_bytes));
527        write_capture_event(&mut writer, Event::Start(captured_start))?;
528        let mut nested = 1usize;
529        while nested > 0 {
530            let event = self.read_event()?;
531            match event {
532                XmlStreamEvent::Start(start) => {
533                    nested = nested.checked_add(1).ok_or(XmlSafetyError::TooDeep)?;
534                    write_capture_event(&mut writer, Event::Start(start.raw.borrow()))?;
535                }
536                XmlStreamEvent::Empty(start) => {
537                    write_capture_event(&mut writer, Event::Empty(start.raw.borrow()))?;
538                }
539                XmlStreamEvent::End(end) => {
540                    write_capture_event(&mut writer, Event::End(end.raw.borrow()))?;
541                    nested -= 1;
542                }
543                XmlStreamEvent::Text(text) => {
544                    write_capture_event(&mut writer, Event::Text(BytesText::new(text.value())))?
545                }
546                XmlStreamEvent::CData(cdata) => {
547                    write_capture_event(&mut writer, Event::CData(BytesCData::new(cdata.value())))?;
548                }
549                XmlStreamEvent::Comment(comment) => {
550                    write_capture_event(
551                        &mut writer,
552                        Event::Comment(BytesText::from_escaped(comment.value())),
553                    )?;
554                }
555                XmlStreamEvent::ProcessingInstruction(pi) => {
556                    write_capture_event(&mut writer, Event::PI(pi.raw.borrow()))?;
557                }
558                XmlStreamEvent::Declaration | XmlStreamEvent::DocType | XmlStreamEvent::Eof => {
559                    return Err(XmlSafetyError::Malformed.into());
560                }
561            }
562        }
563        let bytes = writer.into_inner().bytes;
564        let policy = XmlSafetyPolicy {
565            max_input_bytes: max_bytes,
566            ..self.state.policy
567        };
568        ValidatedXml::with_policy(bytes, policy)
569    }
570
571    pub fn into_inner(self) -> R {
572        self.reader.into_inner().into_inner()
573    }
574
575    fn require_current_start(&self) -> Result<(), Error> {
576        if self.state.current_start_available && self.state.current_start_depth == self.state.depth
577        {
578            Ok(())
579        } else {
580            Err(Error::InvalidData(
581                "operation requires the most recently read event to be Start".into(),
582            ))
583        }
584    }
585}
586
587fn check_stream_position<R: BufRead>(
588    reader: &NsReader<Take<R>>,
589    state: &StreamState,
590) -> Result<(), Error> {
591    if reader.buffer_position() > state.max_input_bytes_u64 {
592        Err(XmlSafetyError::InputTooLarge.into())
593    } else {
594        Ok(())
595    }
596}
597
598fn count_event(state: &mut StreamState, event: &Event<'_>) -> Result<(), Error> {
599    if matches!(event, Event::Eof) {
600        return Ok(());
601    }
602    if state.events >= state.policy.max_events {
603        return Err(XmlSafetyError::TooManyEvents.into());
604    }
605    state.events += 1;
606    Ok(())
607}
608
609fn begin_element<'a, R: BufRead>(
610    state: &mut StreamState,
611    reader: &'a NsReader<Take<R>>,
612    start: &BytesStart<'_>,
613    empty: bool,
614    cached_attribute_values: &mut Vec<CachedAttributeValue>,
615) -> Result<Option<&'a str>, Error> {
616    if state.depth == 0 && state.root_complete {
617        return Err(XmlSafetyError::Malformed.into());
618    }
619    if state.depth >= state.policy.max_depth {
620        return Err(XmlSafetyError::TooDeep.into());
621    }
622    let depth = state.depth + 1;
623    if state.elements >= state.policy.max_elements {
624        return Err(XmlSafetyError::TooManyElements.into());
625    }
626    state.elements += 1;
627    let namespace = resolve_namespace(
628        reader.resolver().resolve_element(start.name()).0,
629        "element namespace",
630    )?;
631    utf8(start.name().into_inner())?;
632
633    for (index, attribute) in start.attributes().enumerate() {
634        let attribute = attribute.map_err(|error| Error::InvalidXml(error.to_string()))?;
635        if index >= state.policy.max_attributes_per_element {
636            return Err(XmlSafetyError::TooManyAttributes.into());
637        }
638        utf8(attribute.key.into_inner())?;
639        resolve_namespace(
640            reader.resolver().resolve_attribute(attribute.key).0,
641            "attribute namespace",
642        )?;
643        let value = attribute
644            .decoded_and_normalized_value(XmlVersion::Explicit1_0, reader.decoder())
645            .map_err(|error| Error::InvalidXml(error.to_string()))?;
646        if let Cow::Owned(value) = value {
647            cached_attribute_values.push(CachedAttributeValue { index, value });
648        }
649    }
650
651    state.root_seen = true;
652    if empty {
653        if state.depth == 0 {
654            state.root_complete = true;
655        }
656    } else {
657        state.depth = depth;
658        state.current_start_depth = depth;
659        state.current_start_available = true;
660    }
661    Ok(namespace)
662}
663
664fn count_text(state: &mut StreamState, value: &str) -> Result<(), Error> {
665    let remaining = state.policy.max_text_bytes - state.text_bytes;
666    if value.len() > remaining {
667        return Err(XmlSafetyError::TextTooLarge.into());
668    }
669    state.text_bytes += value.len();
670    if state.depth == 0 && !value.chars().all(char::is_whitespace) {
671        return Err(XmlSafetyError::Malformed.into());
672    }
673    Ok(())
674}
675
676fn resolve_namespace<'a>(result: ResolveResult<'a>, label: &str) -> Result<Option<&'a str>, Error> {
677    match result {
678        ResolveResult::Unbound => Ok(None),
679        ResolveResult::Bound(namespace) => utf8(namespace.into_inner()).map(Some),
680        ResolveResult::Unknown(prefix) => Err(Error::InvalidXml(format!(
681            "unknown {label} prefix `{}`",
682            String::from_utf8_lossy(&prefix)
683        ))),
684    }
685}
686
687fn decode_text<'a>(text: &BytesText<'a>) -> Result<Cow<'a, str>, Error> {
688    match text.decode().map_err(|_| XmlSafetyError::InvalidEncoding)? {
689        Cow::Borrowed(value) => {
690            unescape(value).map_err(|error| Error::InvalidXml(error.to_string()))
691        }
692        Cow::Owned(value) => unescape(&value)
693            .map(Cow::into_owned)
694            .map(Cow::Owned)
695            .map_err(|error| Error::InvalidXml(error.to_string())),
696    }
697}
698
699fn decode_reference<'a>(
700    reference: &quick_xml::events::BytesRef<'a>,
701) -> Result<Cow<'a, str>, Error> {
702    if let Some(character) = reference
703        .resolve_char_ref()
704        .map_err(|error| Error::InvalidXml(error.to_string()))?
705    {
706        return Ok(Cow::Owned(character.to_string()));
707    }
708    Ok(Cow::Borrowed(match utf8(reference.as_ref())? {
709        "amp" => "&",
710        "lt" => "<",
711        "gt" => ">",
712        "apos" => "'",
713        "quot" => "\"",
714        _ => return Err(XmlSafetyError::ExternalEntity.into()),
715    }))
716}
717
718fn write_capture_event(writer: &mut Writer<LimitedVec>, event: Event<'_>) -> Result<(), Error> {
719    if let Err(error) = writer.write_event(event) {
720        if writer.get_ref().exceeded {
721            return Err(XmlSafetyError::InputTooLarge.into());
722        }
723        return Err(Error::Io(error));
724    }
725    Ok(())
726}
727
728struct LimitedVec {
729    bytes: Vec<u8>,
730    max_bytes: usize,
731    exceeded: bool,
732}
733
734impl LimitedVec {
735    fn new(bytes: Vec<u8>, max_bytes: usize) -> Self {
736        Self {
737            bytes,
738            max_bytes,
739            exceeded: false,
740        }
741    }
742}
743
744impl Write for LimitedVec {
745    fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
746        let Some(new_len) = self.bytes.len().checked_add(buffer.len()) else {
747            self.exceeded = true;
748            return Err(std::io::Error::other("XML capture exceeds byte limit"));
749        };
750        if new_len > self.max_bytes {
751            self.exceeded = true;
752            return Err(std::io::Error::other("XML capture exceeds byte limit"));
753        }
754        self.bytes.extend_from_slice(buffer);
755        Ok(buffer.len())
756    }
757
758    fn flush(&mut self) -> std::io::Result<()> {
759        Ok(())
760    }
761}