use std::collections::BTreeMap;
use super::ModuleError;
use super::codec::Reader;
use super::leb;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct StringTable {
strings: Vec<String>,
}
impl StringTable {
pub fn get(&self, index: usize) -> Option<&str> {
self.strings.get(index).map(String::as_str)
}
pub fn len(&self) -> usize {
self.strings.len()
}
pub fn is_empty(&self) -> bool {
self.strings.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = &str> {
self.strings.iter().map(String::as_str)
}
pub fn encode(&self) -> Vec<u8> {
let mut out = Vec::new();
leb::write(&mut out, self.strings.len() as u64);
for s in &self.strings {
leb::write(&mut out, s.len() as u64);
out.extend_from_slice(s.as_bytes());
}
out
}
pub fn decode(bytes: &[u8]) -> Result<Self, ModuleError> {
let none = Self::default();
let mut r = Reader::new("STRINGS", bytes, &none);
let count = r.count()?;
let mut strings = Vec::with_capacity(count);
let mut seen = BTreeMap::new();
for index in 0..count {
let at = r.position();
let len = r.count()?;
let text = std::str::from_utf8(r.bytes(len)?)
.map_err(|_| r.malformed(at, format!("string {index} is not UTF-8")))?;
if let Some(first) = seen.insert(text, index) {
return Err(r.malformed(at, format!("string {index} repeats string {first}")));
}
strings.push(text.to_owned());
}
r.finish()?;
Ok(Self { strings })
}
}
#[derive(Debug, Default)]
pub(crate) struct Interner {
table: StringTable,
index: BTreeMap<String, usize>,
}
impl Interner {
pub(crate) fn intern(&mut self, text: &str) -> usize {
if let Some(&index) = self.index.get(text) {
return index;
}
let index = self.table.strings.len();
self.table.strings.push(text.to_owned());
self.index.insert(text.to_owned(), index);
index
}
pub(crate) fn table(&self) -> &StringTable {
&self.table
}
}
#[cfg(test)]
mod tests {
use super::*;
fn table(texts: &[&str]) -> StringTable {
StringTable { strings: texts.iter().map(|s| (*s).to_owned()).collect() }
}
#[test]
fn strings_are_numbered_by_first_use() {
let mut interner = Interner::default();
let indices: Vec<usize> =
["PAYROLL", "WS-TOTAL", "PAYROLL", "", "WS-TOTAL", ""].map(|s| interner.intern(s)).into();
assert_eq!(indices, [0, 1, 0, 2, 1, 2]);
assert_eq!(interner.table(), &table(&["PAYROLL", "WS-TOTAL", ""]));
}
#[test]
fn the_table_body_round_trips() {
let t = table(&["", "A", "ÄÖÜ €", "WS-TOTAL"]);
let body = t.encode();
assert_eq!(&body[..4], [4, 0, 1, b'A']);
assert_eq!(StringTable::decode(&body), Ok(t));
assert_eq!(StringTable::decode(&[0]), Ok(StringTable::default()));
}
#[test]
fn a_bad_table_is_malformed() {
let malformed = |bytes: &[u8]| matches!(StringTable::decode(bytes), Err(ModuleError::Malformed { .. }));
assert!(malformed(&[1, 2, 0xC3, 0x28]), "invalid UTF-8");
assert!(malformed(&[2, 1, b'A', 1, b'A']), "a text twice");
assert!(malformed(&[1, 1, b'A', 0]), "trailing bytes");
assert!(malformed(&[2, 1, b'A']), "fewer strings than the count");
assert!(malformed(&[1, 5, b'A']), "a length beyond the body");
assert!(malformed(&[]), "no count");
}
}