use crate::fst::Label;
use std::collections::HashMap;
#[derive(Debug, Clone, Default)]
pub struct SymbolTable {
symbols: Vec<String>,
symbol_map: HashMap<String, Label>,
}
impl SymbolTable {
pub fn new() -> Self {
let mut table = Self::default();
table.add_symbol("<eps>");
table
}
pub fn add_symbol(&mut self, symbol: &str) -> Label {
if let Some(&id) = self.symbol_map.get(symbol) {
id
} else {
let id = self.symbols.len() as Label;
self.symbols.push(symbol.to_string());
self.symbol_map.insert(symbol.to_string(), id);
id
}
}
pub fn find(&self, id: Label) -> Option<&str> {
self.symbols.get(id as usize).map(|s| s.as_str())
}
pub fn find_id(&self, symbol: &str) -> Option<Label> {
self.symbol_map.get(symbol).copied()
}
pub fn len(&self) -> usize {
self.symbols.len()
}
pub fn size(&self) -> usize {
self.len()
}
pub fn find_symbol(&self, symbol: &str) -> Option<Label> {
self.find_id(symbol)
}
pub fn find_key(&self, id: Label) -> Option<&str> {
self.find(id)
}
pub fn contains_symbol(&self, symbol: &str) -> bool {
self.symbol_map.contains_key(symbol)
}
pub fn contains_key(&self, id: Label) -> bool {
(id as usize) < self.symbols.len()
}
pub fn clear(&mut self) {
self.symbols.clear();
self.symbol_map.clear();
self.add_symbol("<eps>");
}
pub fn symbols(&self) -> impl Iterator<Item = &str> {
self.symbols.iter().map(|s| s.as_str())
}
pub fn keys(&self) -> impl Iterator<Item = Label> {
0..self.symbols.len() as Label
}
pub fn is_empty(&self) -> bool {
self.symbols.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
#[test]
fn test_symbol_table_creation() {
let table = SymbolTable::new();
assert_eq!(table.size(), 1);
assert!(!table.is_empty());
assert_eq!(table.find_key(0), Some("<eps>"));
}
#[test]
fn test_symbol_table_add_symbol() {
let mut table = SymbolTable::new();
let id1 = table.add_symbol("hello");
let id2 = table.add_symbol("world");
let id3 = table.add_symbol("hello");
assert_eq!(table.size(), 3);
assert_eq!(id1, id3); assert_ne!(id1, id2); }
#[test]
fn test_symbol_table_find_symbol() {
let mut table = SymbolTable::new();
let id = table.add_symbol("test");
assert_eq!(table.find_symbol("test"), Some(id));
assert_eq!(table.find_symbol("nonexistent"), None);
}
#[test]
fn test_symbol_table_find_key() {
let mut table = SymbolTable::new();
let id = table.add_symbol("example");
assert_eq!(table.find_key(id), Some("example"));
assert_eq!(table.find_key(999), None); }
#[test]
fn test_symbol_table_contains() {
let mut table = SymbolTable::new();
table.add_symbol("exists");
assert!(table.contains_symbol("exists"));
assert!(!table.contains_symbol("does_not_exist"));
let id = table.find_symbol("exists").unwrap();
assert!(table.contains_key(id));
assert!(!table.contains_key(999));
}
#[test]
fn test_symbol_table_clear() {
let mut table = SymbolTable::new();
table.add_symbol("test1");
table.add_symbol("test2");
assert_eq!(table.size(), 3);
table.clear();
assert_eq!(table.size(), 1);
assert!(!table.is_empty());
assert_eq!(table.find_symbol("test1"), None);
assert_eq!(table.find_key(0), Some("<eps>"));
}
#[test]
fn test_symbol_table_iteration() {
let mut table = SymbolTable::new();
table.add_symbol("apple");
table.add_symbol("banana");
table.add_symbol("cherry");
let symbols: HashSet<_> = table.symbols().collect();
let keys: HashSet<_> = table.keys().collect();
assert_eq!(symbols.len(), 4); assert!(symbols.contains("<eps>"));
assert!(symbols.contains("apple"));
assert!(symbols.contains("banana"));
assert!(symbols.contains("cherry"));
assert_eq!(keys.len(), 4);
assert!(keys.contains(&0)); assert!(keys.contains(&1));
assert!(keys.contains(&2));
assert!(keys.contains(&3));
}
#[test]
fn test_symbol_table_len() {
let mut table = SymbolTable::new();
assert_eq!(table.len(), 1);
table.add_symbol("a");
assert_eq!(table.len(), 2);
table.add_symbol("b");
assert_eq!(table.len(), 3);
table.add_symbol("a"); assert_eq!(table.len(), 3); }
#[test]
fn test_symbol_table_special_symbols() {
let mut table = SymbolTable::new();
let eps_id = table.add_symbol("<eps>");
assert_eq!(eps_id, 0);
let special_id = table.add_symbol("<unk>");
assert_ne!(special_id, 0);
assert_eq!(table.find_symbol("<unk>"), Some(special_id));
}
#[test]
fn test_symbol_table_debug() {
let mut table = SymbolTable::new();
table.add_symbol("hello");
table.add_symbol("world");
let debug = format!("{table:?}");
assert!(debug.contains("hello"));
assert!(debug.contains("world"));
}
#[test]
fn test_symbol_table_aliases() {
let mut table = SymbolTable::new();
let id = table.add_symbol("test");
assert_eq!(table.find_symbol("test"), Some(id));
assert_eq!(table.find_key(id), Some("test"));
}
}