use std::collections::BTreeMap;
use std::convert::TryInto;
use std::fs::File;
use std::io::BufReader;
use std::path::Path;
use rustc_hash::FxHashSet;
use crate::arc::Arc;
use crate::error::OpenFstError;
use crate::fst::Fst;
use crate::fst_header::{FstHeader, flags};
use crate::symbol_table::{K_NO_SYMBOL, SymbolTable};
pub fn prune_symbol_table<A, F>(
fst: &F,
syms: &SymbolTable,
input: bool,
) -> Result<SymbolTable, OpenFstError>
where
A: Arc,
F: Fst<A>,
A::Label: TryInto<i64>,
{
let mut seen = FxHashSet::default();
seen.insert(0);
for state in fst.states() {
for arc in fst.arcs(state) {
let sym_label = if input { arc.ilabel() } else { arc.olabel() };
let sym_i64: i64 = sym_label.try_into().map_err(|_| {
OpenFstError::SymbolTable("Failed to cast Arc::Label to i64".to_string())
})?;
seen.insert(sym_i64);
}
}
let mut pruned = SymbolTable::new(format!("{}_pruned", syms.name()));
for item in syms.iter() {
if seen.contains(&item.label) {
pruned.add_symbol(&item.symbol, item.label);
}
}
Ok(pruned)
}
pub fn compact_symbol_table(syms: &SymbolTable) -> SymbolTable {
let sorted: BTreeMap<i64, String> = syms.iter().map(|item| (item.label, item.symbol)).collect();
let mut compact = SymbolTable::new(format!("{}_compact", syms.name()));
for (new_key, (_, symbol)) in sorted.into_iter().enumerate() {
compact.add_symbol(&symbol, new_key as i64);
}
compact
}
pub fn merge_symbol_table(left: &SymbolTable, right: &SymbolTable) -> (SymbolTable, bool) {
let mut merged = SymbolTable::new(format!("merge_{}_{}", left.name(), right.name()));
let mut left_has_all = true;
let mut right_has_all = true;
let mut relabel = false;
for litem in left.iter() {
merged.add_symbol(&litem.symbol, litem.label);
if right_has_all {
let key = right.find_key(&litem.symbol);
if key == K_NO_SYMBOL {
right_has_all = false;
} else if key != litem.label {
right_has_all = false;
relabel = true;
}
}
}
if right_has_all {
return (right.clone(), relabel);
}
let mut conflicts = Vec::new();
for ritem in right.iter() {
let key = merged.find_key(&ritem.symbol);
if key != K_NO_SYMBOL {
if key != ritem.label {
relabel = true;
}
continue;
}
left_has_all = false;
if merged.find_symbol(ritem.label).is_some() {
relabel = true;
conflicts.push(ritem.symbol.clone());
continue;
}
merged.add_symbol(&ritem.symbol, ritem.label);
}
if left_has_all {
return (left.clone(), relabel);
}
for conflict in conflicts {
merged.add_symbol_auto(&conflict);
}
(merged, relabel)
}
pub fn fst_read_symbols(
source: impl AsRef<Path>,
input_symbols: bool,
) -> Result<Option<SymbolTable>, OpenFstError> {
let mut reader = BufReader::new(File::open(source.as_ref())?);
let header = FstHeader::read(&mut reader)?;
if header.flags & flags::HAS_ISYMBOLS != 0 {
let isymbols = SymbolTable::read(&mut reader)?;
if input_symbols {
return Ok(Some(isymbols));
}
}
if header.flags & flags::HAS_OSYMBOLS != 0 {
let osymbols = SymbolTable::read(&mut reader)?;
if !input_symbols {
return Ok(Some(osymbols));
}
}
Ok(None)
}
pub fn add_auxiliary_symbols(
prefix: &str,
start_label: i64,
nlabels: i64,
syms: &mut SymbolTable,
) -> Result<(), OpenFstError> {
for i in 0..nlabels {
let index = i + start_label;
let symbol_str = format!("{}{}", prefix, i);
if index != syms.add_symbol(&symbol_str, index) {
return Err(OpenFstError::SymbolTable(format!(
"AddAuxiliarySymbols: Symbol table clash for symbol '{}' at index {}",
symbol_str, index
)));
}
}
Ok(())
}