Skip to main content

rusty_xml_parser/
dtd.rs

1//! DTD subset parser (internal + caller-supplied external). No network.
2
3use rusty_xml_tree::{AttrDecl, AttrDefault, ElementDecl, XmlDtd};
4use crate::error::XmlError;
5
6/// `xmlParseDTD` — parse a DTD from memory (caller already loaded the bytes).
7#[doc(alias = "xmlParseDTD")]
8pub fn xml_parse_dtd(
9    buffer: &[u8],
10    public_id: Option<&str>,
11    system_id: Option<&str>,
12) -> Result<XmlDtd, XmlError> {
13    let text = String::from_utf8_lossy(buffer);
14    let mut dtd = parse_dtd_subset(&text)?;
15    dtd.public_id = public_id.map(str::to_string);
16    dtd.system_id = system_id.map(str::to_string);
17    Ok(dtd)
18}
19
20/// Parse a DTD internal/external subset into declarations.
21pub fn parse_dtd_subset(src: &str) -> Result<XmlDtd, XmlError> {
22    let expanded = expand_pe(src);
23    let mut dtd = XmlDtd::default();
24    dtd.int_subset = Some(src.to_string());
25    let mut p = DtdParser {
26        src: expanded.as_str(),
27        pos: 0,
28        dtd: &mut dtd,
29    };
30    p.parse_markup()?;
31    Ok(dtd)
32}
33
34fn expand_pe(src: &str) -> String {
35    // Multi-pass PE expansion so `%percent;` can invent new PE names.
36    let mut cur = src.to_string();
37    for _ in 0..16 {
38        let mut pes: std::collections::HashMap<String, String> = std::collections::HashMap::new();
39        harvest_pe(&cur, &mut pes);
40        let next = subst_pe(&cur, &pes);
41        if next == cur {
42            return cur;
43        }
44        cur = next;
45    }
46    cur
47}
48
49fn harvest_pe(src: &str, pes: &mut std::collections::HashMap<String, String>) {
50    let bytes = src.as_bytes();
51    let mut i = 0;
52    while i + 8 < bytes.len() {
53        if bytes[i] == b'<' && bytes.get(i..i + 9) == Some(b"<!ENTITY ") {
54            i += 9;
55            while i < bytes.len() && bytes[i].is_ascii_whitespace() {
56                i += 1;
57            }
58            if i < bytes.len() && bytes[i] == b'%' {
59                i += 1;
60                while i < bytes.len() && bytes[i].is_ascii_whitespace() {
61                    i += 1;
62                }
63                let start = i;
64                while i < bytes.len() && !bytes[i].is_ascii_whitespace() && bytes[i] != b'"' && bytes[i] != b'\'' {
65                    i += 1;
66                }
67                let name = src[start..i].to_string();
68                while i < bytes.len() && bytes[i].is_ascii_whitespace() {
69                    i += 1;
70                }
71                if i < bytes.len() && (bytes[i] == b'"' || bytes[i] == b'\'') {
72                    let q = bytes[i];
73                    i += 1;
74                    let vs = i;
75                    while i < bytes.len() && bytes[i] != q {
76                        i += 1;
77                    }
78                    let val = decode_charrefs(&src[vs..i]);
79                    pes.insert(name, val);
80                }
81            }
82        } else {
83            i += 1;
84        }
85    }
86}
87
88fn subst_pe(src: &str, pes: &std::collections::HashMap<String, String>) -> String {
89    let mut out = String::new();
90    let mut chars = src.chars().peekable();
91    let mut in_comment = false;
92    while let Some(c) = chars.next() {
93        if in_comment {
94            out.push(c);
95            if c == '-' && chars.peek() == Some(&'-') {
96                out.push(chars.next().unwrap());
97                if chars.peek() == Some(&'>') {
98                    out.push(chars.next().unwrap());
99                    in_comment = false;
100                }
101            }
102            continue;
103        }
104        if c == '<' && chars.peek() == Some(&'!') {
105            out.push(c);
106            out.push(chars.next().unwrap());
107            if chars.peek() == Some(&'-') {
108                out.push(chars.next().unwrap());
109                if chars.peek() == Some(&'-') {
110                    out.push(chars.next().unwrap());
111                    in_comment = true;
112                }
113            }
114            continue;
115        }
116        if c == '%' {
117            let mut name = String::new();
118            while let Some(&n) = chars.peek() {
119                if n == ';' {
120                    chars.next();
121                    break;
122                }
123                if n.is_ascii_whitespace() || n == '"' || n == '\'' {
124                    break;
125                }
126                name.push(n);
127                chars.next();
128            }
129            if let Some(v) = pes.get(&name) {
130                out.push_str(v);
131            } else {
132                out.push('%');
133                out.push_str(&name);
134                if !name.is_empty() {
135                    out.push(';');
136                }
137            }
138            continue;
139        }
140        out.push(c);
141    }
142    out
143}
144
145fn decode_charrefs(s: &str) -> String {
146    let mut out = String::new();
147    let mut rest = s;
148    while let Some(i) = rest.find("&#") {
149        out.push_str(&rest[..i]);
150        let after = &rest[i + 2..];
151        if let Some(hex) = after.strip_prefix('x').or_else(|| after.strip_prefix('X')) {
152            if let Some(end) = hex.find(';') {
153                if let Ok(v) = u32::from_str_radix(&hex[..end], 16) {
154                    if let Some(ch) = char::from_u32(v) {
155                        out.push(ch);
156                        rest = &hex[end + 1..];
157                        continue;
158                    }
159                }
160            }
161        } else if let Some(end) = after.find(';') {
162            if let Ok(v) = after[..end].parse::<u32>() {
163                if let Some(ch) = char::from_u32(v) {
164                    out.push(ch);
165                    rest = &after[end + 1..];
166                    continue;
167                }
168            }
169        }
170        out.push_str("&#");
171        rest = after;
172    }
173    out.push_str(rest);
174    out
175}
176
177struct DtdParser<'a> {
178    src: &'a str,
179    pos: usize,
180    dtd: &'a mut XmlDtd,
181}
182
183impl<'a> DtdParser<'a> {
184    fn rest(&self) -> &'a str {
185        &self.src[self.pos..]
186    }
187    fn skip_ws_and_comments(&mut self) {
188        loop {
189            let r = self.rest();
190            let trimmed = r.trim_start();
191            let n = r.len() - trimmed.len();
192            self.pos += n;
193            if self.rest().starts_with("<!--") {
194                if let Some(e) = self.rest().find("-->") {
195                    self.pos += e + 3;
196                    continue;
197                }
198            }
199            if self.rest().starts_with("<?") {
200                if let Some(e) = self.rest().find("?>") {
201                    self.pos += e + 2;
202                    continue;
203                }
204            }
205            break;
206        }
207    }
208    fn parse_markup(&mut self) -> Result<(), XmlError> {
209        loop {
210            self.skip_ws_and_comments();
211            if self.pos >= self.src.len() {
212                break;
213            }
214            if self.rest().starts_with("<!ELEMENT") {
215                self.parse_element()?;
216            } else if self.rest().starts_with("<!ATTLIST") {
217                self.parse_attlist()?;
218            } else if self.rest().starts_with("<!ENTITY") {
219                self.parse_entity()?;
220            } else if self.rest().starts_with("<!NOTATION") {
221                self.skip_decl()?;
222            } else if self.rest().starts_with("<![") {
223                self.skip_cond()?;
224            } else if self.rest().starts_with('<') {
225                self.skip_decl()?;
226            } else {
227                self.pos += self.rest().chars().next().unwrap().len_utf8();
228            }
229        }
230        Ok(())
231    }
232    fn skip_decl(&mut self) -> Result<(), XmlError> {
233        if let Some(i) = self.rest().find('>') {
234            self.pos += i + 1;
235            Ok(())
236        } else {
237            self.pos = self.src.len();
238            Ok(())
239        }
240    }
241    fn skip_cond(&mut self) -> Result<(), XmlError> {
242        let mut depth = 0i32;
243        let bytes = self.rest().as_bytes();
244        let mut i = 0;
245        while i < bytes.len() {
246            if bytes[i] == b'<' && bytes.get(i..i + 3) == Some(b"<![") {
247                depth += 1;
248                i += 3;
249                continue;
250            }
251            if bytes[i] == b']' && bytes.get(i..i + 3) == Some(b"]]>") {
252                depth -= 1;
253                i += 3;
254                if depth == 0 {
255                    self.pos += i;
256                    return Ok(());
257                }
258                continue;
259            }
260            i += 1;
261        }
262        self.pos = self.src.len();
263        Ok(())
264    }
265    fn bump(&mut self, n: usize) {
266        self.pos += n;
267    }
268    fn parse_name(&mut self) -> String {
269        self.skip_ws_and_comments();
270        let r = self.rest();
271        let mut n = 0;
272        for (i, c) in r.char_indices() {
273            if i == 0 {
274                if !(c.is_ascii_alphabetic() || c == '_' || c == ':') {
275                    break;
276                }
277            } else if !(c.is_ascii_alphanumeric() || "-._:".contains(c)) {
278                n = i;
279                break;
280            }
281            n = i + c.len_utf8();
282        }
283        let s = r[..n].to_string();
284        self.bump(n);
285        s
286    }
287    fn parse_quoted(&mut self) -> String {
288        self.skip_ws_and_comments();
289        let r = self.rest();
290        if r.starts_with('"') || r.starts_with('\'') {
291            let q = r.as_bytes()[0] as char;
292            self.bump(1);
293            if let Some(e) = self.rest().find(q) {
294                let s = decode_charrefs(&self.rest()[..e]);
295                self.bump(e + 1);
296                return s;
297            }
298        }
299        String::new()
300    }
301    fn parse_element(&mut self) -> Result<(), XmlError> {
302        self.bump("<!ELEMENT".len());
303        let name = self.parse_name();
304        self.skip_ws_and_comments();
305        let decl = if self.rest().starts_with("EMPTY") {
306            self.bump(5);
307            ElementDecl::Empty
308        } else if self.rest().starts_with("ANY") {
309            self.bump(3);
310            ElementDecl::Any
311        } else if self.rest().starts_with('(') {
312            let spec = self.take_until_gt_paren();
313            if spec.contains("#PCDATA") {
314                let mut names = Vec::new();
315                for part in spec.split('|') {
316                    let t = part.trim().trim_matches(|c: char| c == '(' || c == ')' || c == '*');
317                    if t != "#PCDATA" && !t.is_empty() {
318                        names.push(t.to_string());
319                    }
320                }
321                ElementDecl::Mixed(names)
322            } else {
323                ElementDecl::Children(spec)
324            }
325        } else {
326            self.skip_decl()?;
327            return Ok(());
328        };
329        self.dtd.elements.insert(name, decl);
330        self.skip_ws_and_comments();
331        if self.rest().starts_with('>') {
332            self.bump(1);
333        } else {
334            self.skip_decl()?;
335        }
336        Ok(())
337    }
338    fn take_until_gt_paren(&mut self) -> String {
339        let r = self.rest();
340        let mut depth = 0i32;
341        let mut i = 0;
342        for (off, c) in r.char_indices() {
343            match c {
344                '(' => depth += 1,
345                ')' => {
346                    depth -= 1;
347                    if depth == 0 {
348                        i = off + 1;
349                        break;
350                    }
351                }
352                '>' if depth == 0 => {
353                    i = off;
354                    break;
355                }
356                _ => {}
357            }
358            i = off + c.len_utf8();
359        }
360        let s = r[..i].to_string();
361        self.bump(i);
362        s
363    }
364    fn parse_attlist(&mut self) -> Result<(), XmlError> {
365        self.bump("<!ATTLIST".len());
366        let elem = self.parse_name();
367        loop {
368            self.skip_ws_and_comments();
369            if self.rest().starts_with('>') {
370                self.bump(1);
371                break;
372            }
373            if self.pos >= self.src.len() {
374                break;
375            }
376            let aname = self.parse_name();
377            if aname.is_empty() {
378                self.skip_decl()?;
379                break;
380            }
381            self.skip_ws_and_comments();
382            let mut enumerated = Vec::new();
383            let att_type = if self.rest().starts_with('(') {
384                let spec = self.take_until_gt_paren();
385                for part in spec.split('|') {
386                    let t = part.trim().trim_matches(|c: char| "()".contains(c));
387                    if !t.is_empty() {
388                        enumerated.push(t.to_string());
389                    }
390                }
391                "ENUMERATION".into()
392            } else {
393                self.parse_name()
394            };
395            self.skip_ws_and_comments();
396            let (default, default_value) = if self.rest().starts_with("#REQUIRED") {
397                self.bump(9);
398                (AttrDefault::Required, None)
399            } else if self.rest().starts_with("#IMPLIED") {
400                self.bump(8);
401                (AttrDefault::Implied, None)
402            } else if self.rest().starts_with("#FIXED") {
403                self.bump(6);
404                (AttrDefault::Fixed, Some(self.parse_quoted()))
405            } else {
406                (AttrDefault::Value, Some(self.parse_quoted()))
407            };
408            self.dtd.attributes.insert(
409                (elem.clone(), aname),
410                AttrDecl {
411                    att_type,
412                    default,
413                    default_value,
414                    enumerated,
415                },
416            );
417        }
418        Ok(())
419    }
420    fn parse_entity(&mut self) -> Result<(), XmlError> {
421        self.bump("<!ENTITY".len());
422        self.skip_ws_and_comments();
423        let pe = self.rest().starts_with('%');
424        if pe {
425            self.bump(1);
426            self.skip_ws_and_comments();
427        }
428        let name = self.parse_name();
429        self.skip_ws_and_comments();
430        if self.rest().starts_with("SYSTEM") || self.rest().starts_with("PUBLIC") {
431            self.skip_decl()?;
432            return Ok(());
433        }
434        let val = self.parse_quoted();
435        if pe {
436            self.dtd.parameter_entities.insert(name, val);
437        } else {
438            self.dtd.entities.insert(name, val);
439        }
440        self.skip_ws_and_comments();
441        if self.rest().starts_with('>') {
442            self.bump(1);
443        } else {
444            self.skip_decl()?;
445        }
446        Ok(())
447    }
448}
449
450/// Merge `src` into `dst` (external subset onto internal).
451pub fn merge_dtd(dst: &mut XmlDtd, src: XmlDtd) {
452    dst.entities.extend(src.entities);
453    dst.parameter_entities.extend(src.parameter_entities);
454    dst.elements.extend(src.elements);
455    dst.attributes.extend(src.attributes);
456    if dst.public_id.is_none() {
457        dst.public_id = src.public_id;
458    }
459    if dst.system_id.is_none() {
460        dst.system_id = src.system_id;
461    }
462}