use std::collections::BTreeMap;
use std::fmt::Write as FmtWrite;
use std::hash::{Hash, Hasher};
use std::io::{BufRead, Read, Write};
use md5::{Digest, Md5};
use rustc_hash::FxHasher;
use crate::AtomicRc;
use crate::error::OpenFstError;
use crate::utils::io::{read_scalar, read_string, write_scalar, write_string};
pub const K_NO_SYMBOL: i64 = -1;
pub const K_SYMBOL_TABLE_MAGIC_NUMBER: i32 = 2125658996;
struct CheckSummer(Md5);
impl CheckSummer {
fn new() -> Self {
Self(Md5::new())
}
fn update(&mut self, data: impl AsRef<[u8]>) {
self.0.update(data);
}
fn digest(self) -> String {
let digest = self.0.finalize();
let mut hex_string = String::with_capacity(32);
for byte in digest {
write!(&mut hex_string, "{:02x}", byte).expect("String formatting should never fail");
}
hex_string
}
}
const K_EMPTY_BUCKET: isize = -1;
const K_MAX_OCCUPANCY_RATIO: f32 = 0.75;
#[derive(Debug, Clone)]
struct DenseSymbolMap {
symbols: Vec<String>,
buckets: Vec<isize>,
hash_mask: usize,
}
impl DenseSymbolMap {
fn new() -> Self {
let initial_buckets = 16;
Self {
symbols: Vec::new(),
buckets: vec![K_EMPTY_BUCKET; initial_buckets],
hash_mask: initial_buckets - 1,
}
}
fn get_hash(key: &str) -> u64 {
let mut hasher = FxHasher::default();
key.hash(&mut hasher);
hasher.finish()
}
fn insert_or_find(&mut self, key: &str) -> (usize, bool) {
if self.symbols.len() as f32 >= (self.buckets.len() as f32 * K_MAX_OCCUPANCY_RATIO) {
self.rehash(self.buckets.len() * 2);
}
let mut idx = (Self::get_hash(key) as usize) & self.hash_mask;
loop {
let stored_value = self.buckets[idx];
if stored_value == K_EMPTY_BUCKET {
break;
}
if self.symbols[stored_value as usize] == key {
return (stored_value as usize, false);
}
idx = (idx + 1) & self.hash_mask;
}
let next = self.symbols.len();
self.buckets[idx] = next as isize;
self.symbols.push(key.to_string());
(next, true)
}
fn find(&self, key: &str) -> isize {
let mut idx = (Self::get_hash(key) as usize) & self.hash_mask;
loop {
let stored_value = self.buckets[idx];
if stored_value == K_EMPTY_BUCKET {
return K_EMPTY_BUCKET;
}
if self.symbols[stored_value as usize] == key {
return stored_value;
}
idx = (idx + 1) & self.hash_mask;
}
}
fn size(&self) -> usize {
self.symbols.len()
}
fn get_symbol(&self, idx: usize) -> &str {
&self.symbols[idx]
}
fn rehash(&mut self, num_buckets: usize) {
assert!(num_buckets.is_power_of_two());
self.buckets.clear();
self.buckets.resize(num_buckets, K_EMPTY_BUCKET);
self.hash_mask = num_buckets - 1;
for (i, symbol) in self.symbols.iter().enumerate() {
let mut idx = (Self::get_hash(symbol) as usize) & self.hash_mask;
while self.buckets[idx] != K_EMPTY_BUCKET {
idx = (idx + 1) & self.hash_mask;
}
self.buckets[idx] = i as isize;
}
}
fn remove_symbol(&mut self, idx: usize) {
self.symbols.remove(idx);
self.rehash(self.buckets.len());
}
fn shrink_to_fit(&mut self) {
self.symbols.shrink_to_fit();
let required_capacity = (self.symbols.len() as f32 / K_MAX_OCCUPANCY_RATIO) as usize;
let mut new_buckets = 16;
while new_buckets < required_capacity {
new_buckets *= 2;
}
if new_buckets < self.buckets.len() {
self.rehash(new_buckets);
}
self.buckets.shrink_to_fit();
}
}
#[derive(Debug, Clone)]
struct SymbolTableImpl {
name: String,
available_key: i64,
dense_key_limit: i64,
symbols: DenseSymbolMap,
idx_key: Vec<i64>,
key_map: BTreeMap<i64, i64>,
check_sum_finalized: bool,
check_sum_string: String,
labeled_check_sum_string: String,
}
impl SymbolTableImpl {
fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
available_key: 0,
dense_key_limit: 0,
symbols: DenseSymbolMap::new(),
idx_key: Vec::new(),
key_map: BTreeMap::new(),
check_sum_finalized: false,
check_sum_string: String::new(),
labeled_check_sum_string: String::new(),
}
}
fn add_symbol(&mut self, symbol: &str, key: i64) -> i64 {
if key == K_NO_SYMBOL {
return key;
}
let (insert_idx, inserted) = self.symbols.insert_or_find(symbol);
if !inserted {
let key_already = self.get_nth_key(insert_idx as isize);
return if key_already == key { key } else { key_already };
}
if key + 1 == self.symbols.size() as i64 && key == self.dense_key_limit {
self.dense_key_limit += 1;
} else {
self.idx_key.push(key);
self.key_map.insert(key, self.symbols.size() as i64 - 1);
}
if key >= self.available_key {
self.available_key = key + 1;
}
self.check_sum_finalized = false;
key
}
fn add_symbol_auto(&mut self, symbol: &str) -> i64 {
self.add_symbol(symbol, self.available_key)
}
fn remove_symbol(&mut self, key: i64) {
let mut idx = key;
if key < 0 || key >= self.dense_key_limit {
if let Some(&mapped_idx) = self.key_map.get(&key) {
idx = mapped_idx;
self.key_map.remove(&key);
} else {
return;
}
}
if idx < 0 || idx >= self.symbols.size() as i64 {
return;
}
self.symbols.remove_symbol(idx as usize);
for mapped_idx in self.key_map.values_mut() {
if *mapped_idx > idx {
*mapped_idx -= 1;
}
}
if key >= 0 && key < self.dense_key_limit {
let new_dense_key_limit = key;
for i in (key + 1)..self.dense_key_limit {
self.key_map.insert(i, i - 1);
}
let old_idx_key_len = self.idx_key.len();
self.idx_key
.resize(self.symbols.size() - new_dense_key_limit as usize, 0);
for i in (self.dense_key_limit as usize..=self.symbols.size()).rev() {
let dest_idx = i - new_dense_key_limit as usize - 1;
let src_idx = i - self.dense_key_limit as usize;
if src_idx < old_idx_key_len {
self.idx_key[dest_idx] = self.idx_key[src_idx];
} else {
self.idx_key[dest_idx] = 0;
}
}
for i in new_dense_key_limit..(self.dense_key_limit - 1) {
self.idx_key[(i - new_dense_key_limit) as usize] = i + 1;
}
self.dense_key_limit = new_dense_key_limit;
} else {
let start_idx = (idx - self.dense_key_limit) as usize;
for i in start_idx..(self.idx_key.len() - 1) {
self.idx_key[i] = self.idx_key[i + 1];
}
self.idx_key.pop();
}
if key == self.available_key - 1 {
self.available_key = key;
}
self.check_sum_finalized = false;
}
fn find_symbol(&self, key: i64) -> Option<&str> {
let idx = if key < 0 || key >= self.dense_key_limit {
*self.key_map.get(&key)?
} else {
key
};
if idx >= 0 && (idx as usize) < self.symbols.size() {
Some(self.symbols.get_symbol(idx as usize))
} else {
None
}
}
fn find_key(&self, symbol: &str) -> i64 {
let idx = self.symbols.find(symbol);
if idx == K_EMPTY_BUCKET {
return K_NO_SYMBOL;
}
if idx < self.dense_key_limit as isize {
return idx as i64;
}
self.idx_key[(idx - self.dense_key_limit as isize) as usize]
}
fn get_nth_key(&self, pos: isize) -> i64 {
if pos < 0 || pos as usize >= self.symbols.size() {
K_NO_SYMBOL
} else if pos < self.dense_key_limit as isize {
pos as i64
} else {
self.find_key(self.symbols.get_symbol(pos as usize))
}
}
fn maybe_recompute_check_sum(&mut self) {
if self.check_sum_finalized {
return;
}
let mut check_sum = CheckSummer::new();
for i in 0..self.symbols.size() {
check_sum.update(self.symbols.get_symbol(i).as_bytes());
check_sum.update(b"\0");
}
self.check_sum_string = check_sum.digest();
let mut labeled_check_sum = CheckSummer::new();
for i in 0..self.dense_key_limit {
labeled_check_sum.update(format!("{}\t{}", self.symbols.get_symbol(i as usize), i));
}
for (&key, &idx) in &self.key_map {
if (0..self.dense_key_limit).contains(&key) {
continue;
}
labeled_check_sum.update(format!(
"{}\t{}",
self.symbols.get_symbol(idx as usize),
key
));
}
self.labeled_check_sum_string = labeled_check_sum.digest();
self.check_sum_finalized = true;
}
fn write<W: Write>(&self, w: &mut W) -> Result<(), OpenFstError> {
write_scalar(w, K_SYMBOL_TABLE_MAGIC_NUMBER)?;
write_string(w, &self.name)?;
write_scalar(w, self.available_key)?;
write_scalar(w, self.symbols.size() as i64)?;
for i in 0..self.dense_key_limit {
write_string(w, self.symbols.get_symbol(i as usize))?;
write_scalar(w, i)?;
}
for (&key, &idx) in &self.key_map {
write_string(w, self.symbols.get_symbol(idx as usize))?;
write_scalar(w, key)?;
}
w.flush()?;
Ok(())
}
fn read<R: Read>(r: &mut R) -> Result<Self, OpenFstError> {
let magic: i32 = read_scalar(r)?;
if magic != K_SYMBOL_TABLE_MAGIC_NUMBER {
return Err(OpenFstError::SymbolTable(format!(
"Invalid symbol table magic number: expected {}, found {}",
K_SYMBOL_TABLE_MAGIC_NUMBER, magic
)));
}
let name = read_string(r)?;
let mut table = Self::new(name);
table.available_key = read_scalar(r)?;
let size: i64 = read_scalar(r)?;
if size < 0 {
return Err(OpenFstError::SymbolTable(format!(
"Invalid symbol table size: {size}"
)));
}
table.check_sum_finalized = false;
for _ in 0..size {
let symbol = read_string(r)?;
let key: i64 = read_scalar(r)?;
table.add_symbol(&symbol, key);
}
table.symbols.shrink_to_fit();
Ok(table)
}
}
#[derive(Debug, Clone)]
pub struct SymbolTable {
inner: AtomicRc<SymbolTableImpl>,
}
impl SymbolTable {
pub fn new(name: impl Into<String>) -> Self {
Self {
inner: AtomicRc::new(SymbolTableImpl::new(name)),
}
}
fn make_mut(&mut self) -> &mut SymbolTableImpl {
std::sync::Arc::make_mut(&mut self.inner)
}
pub fn read<R: Read>(reader: &mut R) -> Result<Self, OpenFstError> {
let inner = SymbolTableImpl::read(reader)?;
Ok(Self {
inner: AtomicRc::new(inner),
})
}
pub fn write<W: Write>(&self, writer: &mut W) -> Result<(), OpenFstError> {
self.inner.write(writer)
}
pub fn read_text<R: BufRead>(
reader: &mut R,
name: impl Into<String>,
sep: &str,
) -> Result<Self, OpenFstError> {
let mut table = SymbolTableImpl::new(name);
let default_sep = "\t ";
let separator = if sep.is_empty() { default_sep } else { sep };
let mut line = String::new();
let mut nline = 0;
while let Ok(bytes) = reader.read_line(&mut line) {
if bytes == 0 {
break;
}
nline += 1;
let trimmed = line.trim_end_matches(&['\n', '\r'][..]);
if trimmed.is_empty() {
line.clear();
continue;
}
let parts: Vec<&str> = trimmed
.split(|c| separator.contains(c))
.filter(|s| !s.is_empty())
.collect();
if parts.len() != 2 {
return Err(OpenFstError::SymbolTable(format!(
"ReadText: Bad number of columns ({}), line = {}",
parts.len(),
nline
)));
}
let symbol = parts[0];
let key = parts[1].parse::<i64>().map_err(|_| {
OpenFstError::SymbolTable(format!(
"ReadText: Invalid integer label ({}), line = {}",
parts[1], nline
))
})?;
table.add_symbol(symbol, key);
line.clear();
}
table.symbols.shrink_to_fit();
Ok(Self {
inner: AtomicRc::new(table),
})
}
pub fn write_text<W: Write>(&self, writer: &mut W, sep: &str) -> Result<(), OpenFstError> {
let default_sep = "\t";
let separator = if sep.is_empty() { default_sep } else { sep };
for item in self.iter() {
writeln!(writer, "{}{}{}", item.symbol, separator, item.label)?;
}
writer.flush()?;
Ok(())
}
pub fn add_symbol(&mut self, symbol: &str, key: i64) -> i64 {
self.make_mut().add_symbol(symbol, key)
}
pub fn add_symbol_auto(&mut self, symbol: &str) -> i64 {
self.make_mut().add_symbol_auto(symbol)
}
pub fn remove_symbol(&mut self, key: i64) {
self.make_mut().remove_symbol(key)
}
pub fn find_symbol(&self, key: i64) -> Option<&str> {
self.inner.find_symbol(key)
}
pub fn find_key(&self, symbol: &str) -> i64 {
self.inner.find_key(symbol)
}
pub fn member_key(&self, key: i64) -> bool {
self.find_symbol(key).is_some()
}
pub fn member_symbol(&self, symbol: &str) -> bool {
self.find_key(symbol) != K_NO_SYMBOL
}
pub fn available_key(&self) -> i64 {
self.inner.available_key
}
pub fn name(&self) -> &str {
&self.inner.name
}
pub fn set_name(&mut self, name: impl Into<String>) {
self.make_mut().name = name.into();
}
pub fn num_symbols(&self) -> usize {
self.inner.symbols.size()
}
fn get_nth_key(&self, pos: isize) -> i64 {
self.inner.get_nth_key(pos)
}
pub fn check_sum(&mut self) -> &str {
self.make_mut().maybe_recompute_check_sum();
&self.inner.check_sum_string
}
pub fn labeled_check_sum(&mut self) -> &str {
self.make_mut().maybe_recompute_check_sum();
&self.inner.labeled_check_sum_string
}
pub fn add_table(&mut self, table: &SymbolTable) {
let mut_impl = self.make_mut();
for item in table.iter() {
mut_impl.add_symbol_auto(&item.symbol);
}
}
pub fn iter(&self) -> SymbolTableIterator<'_> {
SymbolTableIterator {
table: self,
pos: 0,
nsymbols: self.num_symbols(),
}
}
}
pub struct SymbolTableItem {
pub label: i64,
pub symbol: String,
}
pub struct SymbolTableIterator<'a> {
table: &'a SymbolTable,
pos: usize,
nsymbols: usize,
}
impl<'a> Iterator for SymbolTableIterator<'a> {
type Item = SymbolTableItem;
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.nsymbols {
let key = self.table.get_nth_key(self.pos as isize);
let symbol = self.table.find_symbol(key).unwrap().to_string();
self.pos += 1;
Some(SymbolTableItem { label: key, symbol })
} else {
None
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
let rem = self.nsymbols - self.pos;
(rem, Some(rem))
}
}
pub fn compat_symbols(syms1: Option<&mut SymbolTable>, syms2: Option<&mut SymbolTable>) -> bool {
if let (Some(s1), Some(s2)) = (syms1, syms2)
&& s1.labeled_check_sum() != s2.labeled_check_sum()
{
return false;
}
true
}
pub fn compat_symbols_rc(
mut syms1: Option<AtomicRc<SymbolTable>>,
mut syms2: Option<AtomicRc<SymbolTable>>,
) -> bool {
compat_symbols(
syms1.as_mut().map(AtomicRc::make_mut),
syms2.as_mut().map(AtomicRc::make_mut),
)
}
pub fn compat_symbols_with_warn(
syms1: Option<&mut SymbolTable>,
syms2: Option<&mut SymbolTable>,
warning: &mut dyn Write,
) -> Result<bool, std::io::Error> {
if let (Some(s1), Some(s2)) = (syms1, syms2)
&& s1.labeled_check_sum() != s2.labeled_check_sum()
{
writeln!(
warning,
"WARNING: CompatSymbols: Symbol table checksums do not match. Table sizes are {} and {}",
s1.num_symbols(),
s2.num_symbols()
)?;
return Ok(false);
}
Ok(true)
}