ironwork_rt/module/
strings.rs1use std::collections::BTreeMap;
2
3use super::ModuleError;
4use super::codec::Reader;
5use super::leb;
6
7#[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 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 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#[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}