Skip to main content

ironwork_rt/module/
strings.rs

1use std::collections::BTreeMap;
2
3use super::ModuleError;
4use super::codec::Reader;
5use super::leb;
6
7/// The one table of strings a module's sections refer to by index (load-module.md §4.2).
8#[derive(Clone, Debug, Default, PartialEq, Eq)]
9pub struct StringTable {
10    strings: Vec<String>,
11}
12
13impl StringTable {
14    pub fn get(&self, index: usize) -> Option<&str> {
15        self.strings.get(index).map(String::as_str)
16    }
17
18    pub fn len(&self) -> usize {
19        self.strings.len()
20    }
21
22    pub fn is_empty(&self) -> bool {
23        self.strings.is_empty()
24    }
25
26    pub fn iter(&self) -> impl Iterator<Item = &str> {
27        self.strings.iter().map(String::as_str)
28    }
29
30    /// The `STRINGS` section body: a count, then each string's byte length and UTF-8 bytes.
31    pub fn encode(&self) -> Vec<u8> {
32        let mut out = Vec::new();
33        leb::write(&mut out, self.strings.len() as u64);
34        for s in &self.strings {
35            leb::write(&mut out, s.len() as u64);
36            out.extend_from_slice(s.as_bytes());
37        }
38        out
39    }
40
41    /// Refuses invalid UTF-8, a text stored twice, and bytes after the last string.
42    pub fn decode(bytes: &[u8]) -> Result<Self, ModuleError> {
43        let none = Self::default();
44        let mut r = Reader::new("STRINGS", bytes, &none);
45        let count = r.count()?;
46        let mut strings = Vec::with_capacity(count);
47        let mut seen = BTreeMap::new();
48        for index in 0..count {
49            let at = r.position();
50            let len = r.count()?;
51            let text = std::str::from_utf8(r.bytes(len)?)
52                .map_err(|_| r.malformed(at, format!("string {index} is not UTF-8")))?;
53            if let Some(first) = seen.insert(text, index) {
54                return Err(r.malformed(at, format!("string {index} repeats string {first}")));
55            }
56            strings.push(text.to_owned());
57        }
58        r.finish()?;
59        Ok(Self { strings })
60    }
61}
62
63/// Numbers each text by its first use.
64#[derive(Debug, Default)]
65pub(crate) struct Interner {
66    table: StringTable,
67    index: BTreeMap<String, usize>,
68}
69
70impl Interner {
71    pub(crate) fn intern(&mut self, text: &str) -> usize {
72        if let Some(&index) = self.index.get(text) {
73            return index;
74        }
75        let index = self.table.strings.len();
76        self.table.strings.push(text.to_owned());
77        self.index.insert(text.to_owned(), index);
78        index
79    }
80
81    pub(crate) fn table(&self) -> &StringTable {
82        &self.table
83    }
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89
90    fn table(texts: &[&str]) -> StringTable {
91        StringTable { strings: texts.iter().map(|s| (*s).to_owned()).collect() }
92    }
93
94    #[test]
95    fn strings_are_numbered_by_first_use() {
96        let mut interner = Interner::default();
97        let indices: Vec<usize> =
98            ["PAYROLL", "WS-TOTAL", "PAYROLL", "", "WS-TOTAL", ""].map(|s| interner.intern(s)).into();
99        assert_eq!(indices, [0, 1, 0, 2, 1, 2]);
100        assert_eq!(interner.table(), &table(&["PAYROLL", "WS-TOTAL", ""]));
101    }
102
103    #[test]
104    fn the_table_body_round_trips() {
105        let t = table(&["", "A", "ÄÖÜ €", "WS-TOTAL"]);
106        let body = t.encode();
107        assert_eq!(&body[..4], [4, 0, 1, b'A']);
108        assert_eq!(StringTable::decode(&body), Ok(t));
109        assert_eq!(StringTable::decode(&[0]), Ok(StringTable::default()));
110    }
111
112    #[test]
113    fn a_bad_table_is_malformed() {
114        let malformed = |bytes: &[u8]| matches!(StringTable::decode(bytes), Err(ModuleError::Malformed { .. }));
115        assert!(malformed(&[1, 2, 0xC3, 0x28]), "invalid UTF-8");
116        assert!(malformed(&[2, 1, b'A', 1, b'A']), "a text twice");
117        assert!(malformed(&[1, 1, b'A', 0]), "trailing bytes");
118        assert!(malformed(&[2, 1, b'A']), "fewer strings than the count");
119        assert!(malformed(&[1, 5, b'A']), "a length beyond the body");
120        assert!(malformed(&[]), "no count");
121    }
122}