Skip to main content

docx_rs/xml/
writer.rs

1use std::borrow::Cow;
2use std::fmt;
3use std::io::Write;
4use std::marker::PhantomData;
5
6use quick_xml::events::{BytesDecl, BytesText, Event};
7use quick_xml::Writer as QuickWriter;
8use smallvec::SmallVec;
9
10use super::common::XmlVersion;
11
12pub type Result<T> = std::result::Result<T, Error>;
13
14#[derive(Debug)]
15pub enum Error {
16    Io(std::io::Error),
17    Xml(quick_xml::Error),
18    UnbalancedEndTag,
19}
20
21impl fmt::Display for Error {
22    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
23        match self {
24            Error::Io(err) => write!(f, "io error: {err}"),
25            Error::Xml(err) => write!(f, "xml error: {err}"),
26            Error::UnbalancedEndTag => write!(f, "attempted to close more elements than opened"),
27        }
28    }
29}
30
31impl std::error::Error for Error {
32    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
33        match self {
34            Error::Io(err) => Some(err),
35            Error::Xml(err) => Some(err),
36            Error::UnbalancedEndTag => None,
37        }
38    }
39}
40
41impl From<std::io::Error> for Error {
42    fn from(value: std::io::Error) -> Self {
43        Error::Io(value)
44    }
45}
46
47impl From<quick_xml::Error> for Error {
48    fn from(value: quick_xml::Error) -> Self {
49        Error::Xml(value)
50    }
51}
52
53#[derive(Clone, Debug)]
54pub struct EmitterConfig {
55    pub write_document_declaration: bool,
56    pub perform_escaping: bool,
57    pub perform_indent: bool,
58    pub line_separator: Cow<'static, str>,
59}
60
61impl Default for EmitterConfig {
62    fn default() -> Self {
63        Self {
64            write_document_declaration: true,
65            perform_escaping: true,
66            perform_indent: false,
67            line_separator: Cow::Borrowed("\n"),
68        }
69    }
70}
71
72impl EmitterConfig {
73    pub fn create_writer<W: Write>(&self, writer: W) -> EventWriter<W> {
74        EventWriter {
75            writer: QuickWriter::new(writer),
76            perform_escaping: self.perform_escaping,
77            element_stack: Vec::new(),
78            element_names: Vec::new(),
79            pending_start: SmallVec::new(),
80        }
81    }
82}
83
84#[derive(Debug)]
85struct ElementState {
86    name_start: usize,
87    name_len: usize,
88    pending: bool,
89}
90
91pub struct EventWriter<W: Write> {
92    writer: QuickWriter<W>,
93    perform_escaping: bool,
94    element_stack: Vec<ElementState>,
95    element_names: Vec<u8>,
96    pending_start: SmallVec<[u8; 128]>,
97}
98
99impl<W: Write> EventWriter<W> {
100    pub fn write<'a, E>(&mut self, event: E) -> Result<()>
101    where
102        E: Into<XmlEvent<'a>>,
103    {
104        match event.into() {
105            XmlEvent::StartDocument {
106                version,
107                encoding,
108                standalone,
109            } => {
110                let standalone_text = standalone.map(|flag| if flag { "yes" } else { "no" });
111                let decl = BytesDecl::new(version.as_str(), encoding, standalone_text);
112                self.writer.write_event(Event::Decl(decl))?;
113            }
114            XmlEvent::StartElement(element) => {
115                self.flush_pending()?;
116                let name = &element.encoded[1..1 + element.name_len];
117                let name_start = self.element_names.len();
118                self.element_names.extend_from_slice(name);
119                self.element_stack.push(ElementState {
120                    name_start,
121                    name_len: element.name_len,
122                    pending: true,
123                });
124                self.pending_start = element.encoded;
125            }
126            XmlEvent::EndElement => {
127                let state = self.element_stack.pop().ok_or(Error::UnbalancedEndTag)?;
128                if state.pending {
129                    let writer = self.writer.get_mut();
130                    writer.write_all(&self.pending_start).map_err(Error::from)?;
131                    writer.write_all(b" />").map_err(Error::from)?;
132                    self.pending_start.clear();
133                } else {
134                    let name_end = state.name_start + state.name_len;
135                    let writer = self.writer.get_mut();
136                    writer.write_all(b"</").map_err(Error::from)?;
137                    writer
138                        .write_all(&self.element_names[state.name_start..name_end])
139                        .map_err(Error::from)?;
140                    writer.write_all(b">").map_err(Error::from)?;
141                }
142                self.element_names.truncate(state.name_start);
143            }
144            XmlEvent::Characters(text) => {
145                self.flush_pending()?;
146                let text_event = if self.perform_escaping {
147                    BytesText::new(text.as_ref())
148                } else {
149                    BytesText::from_escaped(text.as_ref())
150                };
151                self.writer.write_event(Event::Text(text_event))?;
152            }
153        }
154        Ok(())
155    }
156
157    pub fn into_inner(self) -> Result<W> {
158        Ok(self.writer.into_inner())
159    }
160
161    pub fn inner_mut(&mut self) -> Result<&mut W> {
162        Ok(self.writer.get_mut())
163    }
164
165    fn flush_pending(&mut self) -> Result<()> {
166        if let Some(state) = self.element_stack.last_mut() {
167            if state.pending {
168                state.pending = false;
169                let writer = self.writer.get_mut();
170                writer.write_all(&self.pending_start).map_err(Error::from)?;
171                writer.write_all(b">").map_err(Error::from)?;
172                self.pending_start.clear();
173            }
174        }
175        Ok(())
176    }
177}
178
179#[derive(Clone, Debug)]
180pub struct StartElement<'a> {
181    encoded: SmallVec<[u8; 128]>,
182    name_len: usize,
183    _marker: PhantomData<&'a ()>,
184}
185
186#[derive(Clone, Debug)]
187pub enum XmlEvent<'a> {
188    StartDocument {
189        version: XmlVersion,
190        encoding: Option<&'a str>,
191        standalone: Option<bool>,
192    },
193    StartElement(StartElement<'a>),
194    EndElement,
195    Characters(Cow<'a, str>),
196}
197
198impl<'a> XmlEvent<'a> {
199    pub fn start_element(name: &'a str) -> StartElement<'a> {
200        let mut encoded = SmallVec::new();
201        encoded.push(b'<');
202        encoded.extend_from_slice(name.as_bytes());
203        StartElement {
204            encoded,
205            name_len: name.len(),
206            _marker: std::marker::PhantomData,
207        }
208    }
209
210    pub fn end_element() -> XmlEvent<'static> {
211        XmlEvent::EndElement
212    }
213}
214
215impl<'a> From<StartElement<'a>> for XmlEvent<'a> {
216    fn from(value: StartElement<'a>) -> Self {
217        XmlEvent::StartElement(value)
218    }
219}
220
221impl<'a> From<&'a str> for XmlEvent<'a> {
222    fn from(value: &'a str) -> Self {
223        XmlEvent::Characters(Cow::Borrowed(value))
224    }
225}
226
227impl From<String> for XmlEvent<'static> {
228    fn from(value: String) -> Self {
229        XmlEvent::Characters(Cow::Owned(value))
230    }
231}
232
233impl<'a> StartElement<'a> {
234    /// Appends an attribute with an already escaped string value.
235    ///
236    /// The value is copied directly into the start tag and is not XML-escaped.
237    pub fn attr(mut self, name: &str, value: &str) -> Self {
238        self.encoded.push(b' ');
239        self.encoded.extend_from_slice(name.as_bytes());
240        self.encoded.extend_from_slice(b"=\"");
241        self.encoded.extend_from_slice(value.as_bytes());
242        self.encoded.push(b'"');
243        self
244    }
245
246    /// Appends an attribute by formatting its value directly into the start tag.
247    ///
248    /// This avoids allocating an intermediate `String`. The formatted value is
249    /// not XML-escaped.
250    pub fn attr_display(mut self, name: &str, value: impl fmt::Display) -> Self {
251        self.encoded.push(b' ');
252        self.encoded.extend_from_slice(name.as_bytes());
253        self.encoded.extend_from_slice(b"=\"");
254        fmt::write(
255            &mut AttributeValueWriter(&mut self.encoded),
256            format_args!("{value}"),
257        )
258        .expect("writing to an in-memory attribute buffer cannot fail");
259        self.encoded.push(b'"');
260        self
261    }
262}
263
264struct AttributeValueWriter<'a>(&'a mut SmallVec<[u8; 128]>);
265
266impl fmt::Write for AttributeValueWriter<'_> {
267    fn write_str(&mut self, value: &str) -> fmt::Result {
268        self.0.extend_from_slice(value.as_bytes());
269        Ok(())
270    }
271}