use alloc::vec::Vec;
use crate::write::{Error, Result, WritableBuffer};
#[cfg(feature = "write_std")]
type IndexSet<K> = indexmap::IndexSet<K>;
#[cfg(not(feature = "write_std"))]
type IndexSet<K> = indexmap::IndexSet<K, hashbrown::DefaultHashBuilder>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StringId(u32);
#[derive(Debug, Default)]
pub struct StringTable<'a> {
strings: IndexSet<&'a [u8]>,
offsets: Vec<u32>,
size: u64,
in_order: bool,
base: u32,
}
impl<'a> StringTable<'a> {
pub fn new() -> Self {
StringTable::default()
}
pub fn new_in_order(base: u32) -> Self {
StringTable {
strings: IndexSet::default(),
offsets: Vec::new(),
size: base.into(),
in_order: true,
base,
}
}
pub fn is_empty(&self) -> bool {
self.strings.is_empty()
}
pub fn add(&mut self, string: &'a [u8]) -> StringId {
debug_assert!(self.in_order || self.offsets.is_empty());
debug_assert!(!string.contains(&0));
let (id, new) = self.strings.insert_full(string);
if new && self.in_order {
self.offsets.push(self.size as u32);
self.size += string.len() as u64 + 1;
}
StringId(id as u32)
}
pub fn get_id(&self, string: &[u8]) -> StringId {
let id = self.strings.get_index_of(string).unwrap();
StringId(id as u32)
}
pub fn get_string(&self, id: StringId) -> &'a [u8] {
self.strings.get_index(id.0 as usize).unwrap()
}
pub fn get_offset(&self, id: StringId) -> u32 {
self.offsets[id.0 as usize]
}
pub fn maybe_get_offset(&self, id: Option<StringId>) -> u32 {
let Some(id) = id else {
return 0;
};
self.get_offset(id)
}
pub fn write<W: WritableBuffer + ?Sized>(&mut self, w: &mut W, base: u32) -> Result<u32> {
if !self.in_order && !self.offsets.is_empty() {
return Err(Error("string table already written".into()));
}
if self.strings.iter().any(|s| s.contains(&0)) {
return Err(Error("string table entry contains null byte".into()));
}
if self.in_order {
debug_assert_eq!(self.base, base);
for string in &self.strings {
w.write_bytes(string);
w.write_bytes(&[0]);
}
return u32::try_from(self.size)
.map_err(|_| Error("string table size overflow".into()));
}
let mut ids: Vec<_> = (0..self.strings.len()).collect();
sort(&mut ids, 1, &self.strings);
self.offsets = vec![0; ids.len()];
let mut offset = u64::from(base);
let mut previous = &[][..];
for id in ids {
let string = self.strings.get_index(id).unwrap();
let len = string.len() as u64 + 1;
if previous.ends_with(string) {
self.offsets[id] = (offset - len) as u32;
} else {
self.offsets[id] = offset as u32;
w.write_bytes(string);
w.write_bytes(&[0]);
offset += len;
previous = string;
}
}
u32::try_from(offset).map_err(|_| Error("string table size overflow".into()))
}
#[cfg(all(feature = "build_core", feature = "elf"))]
pub(crate) fn size(&self, base: u32) -> u64 {
if self.in_order {
debug_assert_eq!(self.base, base);
return self.size;
}
let mut ids: Vec<_> = (0..self.strings.len()).collect();
sort(&mut ids, 1, &self.strings);
let mut size = base as u64;
let mut previous = &[][..];
for id in ids {
let string = self.strings.get_index(id).unwrap();
if !previous.ends_with(string) {
size += string.len() as u64 + 1;
previous = string;
}
}
size
}
}
fn sort(mut ids: &mut [usize], mut pos: usize, strings: &IndexSet<&[u8]>) {
loop {
if ids.len() <= 1 {
return;
}
let pivot = byte(ids[0], pos, strings);
let mut lower = 0;
let mut upper = ids.len();
let mut i = 1;
while i < upper {
let b = byte(ids[i], pos, strings);
if b > pivot {
ids.swap(lower, i);
lower += 1;
i += 1;
} else if b < pivot {
upper -= 1;
ids.swap(upper, i);
} else {
i += 1;
}
}
sort(&mut ids[..lower], pos, strings);
sort(&mut ids[upper..], pos, strings);
if pivot == 0 {
return;
}
ids = &mut ids[lower..upper];
pos += 1;
}
}
fn byte(id: usize, pos: usize, strings: &IndexSet<&[u8]>) -> u8 {
let string = strings.get_index(id).unwrap();
let len = string.len();
if len >= pos {
string[len - pos]
} else {
0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn string_table() {
let mut table = StringTable::default();
let id0 = table.add(b"");
let id1 = table.add(b"foo");
let id2 = table.add(b"bar");
let id3 = table.add(b"foobar");
let mut data = vec![0];
assert_eq!(table.write(&mut data, 1), Ok(12));
assert_eq!(data, b"\0foobar\0foo\0");
assert_eq!(table.get_offset(id0), 11);
assert_eq!(table.get_offset(id1), 8);
assert_eq!(table.get_offset(id2), 4);
assert_eq!(table.get_offset(id3), 1);
let mut data = Vec::new();
assert!(table.write(&mut data, 1).is_err());
}
#[test]
fn string_table_in_order() {
let mut table = StringTable::new_in_order(1);
let id0 = table.add(b"");
let id1 = table.add(b"foo");
let id2 = table.add(b"bar");
let id3 = table.add(b"foobar");
let mut data = vec![0];
assert_eq!(table.write(&mut data, 1), Ok(17));
assert_eq!(data, b"\0\0foo\0bar\0foobar\0");
assert_eq!(table.get_offset(id0), 1);
assert_eq!(table.get_offset(id1), 2);
assert_eq!(table.get_offset(id2), 6);
assert_eq!(table.get_offset(id3), 10);
let mut data2 = vec![0];
assert_eq!(table.write(&mut data2, 1), Ok(17));
assert_eq!(data, data2);
}
}