use std::cell::RefCell;
use std::hash::Hash;
use std::io::{Read, Write};
use std::rc::Rc;
use crate::AtomicRc;
use crate::algorithms::arc_map::{ArcMapper, MapFinalAction, MapSymbolsAction, arc_map};
use crate::algorithms::rmfinalepsilon::rm_final_epsilon;
use crate::arc::{Arc, ArcLabel, ArcStateId};
use crate::data_structures::bi_table::CompactHashBiTable;
use crate::error::OpenFstError;
use crate::fst::{Fst, MutableFst};
use crate::fst_type::ArcType;
use crate::properties::{
K_ACCEPTOR, K_ADD_SUPER_FINAL_PROPERTIES, K_ERROR, K_FST_PROPERTIES, K_I_DETERMINISTIC,
K_I_LABEL_INVARIANT_PROPERTIES, K_O_LABEL_INVARIANT_PROPERTIES, K_RM_SUPER_FINAL_PROPERTIES,
K_UNWEIGHTED, K_UNWEIGHTED_CYCLES, K_WEIGHT_INVARIANT_PROPERTIES,
};
use crate::symbol_table::SymbolTable;
use crate::utils::io::{FstScalar, read_scalar, read_string, write_scalar, write_string};
use crate::weight::{Weight, WeightIo};
pub const ENCODE_LABELS: u8 = 0x01;
pub const ENCODE_WEIGHTS: u8 = 0x02;
pub const ENCODE_FLAGS: u8 = ENCODE_LABELS | ENCODE_WEIGHTS;
const ENCODE_HAS_ISYMBOLS: u8 = 0x04;
const ENCODE_HAS_OSYMBOLS: u8 = 0x08;
pub const ENCODE_MAGIC_NUMBER: i32 = 2128178506;
const ENCODE_DEPRECATED_MAGIC_NUMBER: i32 = 2129983209;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EncodeType {
Encode,
Decode,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct Triple<L, W> {
pub ilabel: L,
pub olabel: L,
pub weight: W,
}
impl<L: ArcLabel, W: Weight> Triple<L, W> {
fn from_arc<A: Arc<Label = L, Weight = W>>(arc: &A, flags: u8) -> Self {
Self {
ilabel: arc.ilabel(),
olabel: if flags & ENCODE_LABELS != 0 {
arc.olabel()
} else {
L::epsilon()
},
weight: if flags & ENCODE_WEIGHTS != 0 {
arc.weight().clone()
} else {
W::one()
},
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EncodeTableHeader {
pub arc_type: String,
pub flags: u8,
pub size: u64,
}
impl EncodeTableHeader {
pub fn read<R: Read>(reader: &mut R) -> Result<Self, OpenFstError> {
let magic: i32 = read_scalar(reader)?;
match magic {
ENCODE_MAGIC_NUMBER => Ok(Self {
arc_type: read_string(reader)?,
flags: read_scalar(reader)?,
size: read_scalar(reader)?,
}),
ENCODE_DEPRECATED_MAGIC_NUMBER => {
let flags: u32 = read_scalar(reader)?;
let size: i64 = read_scalar(reader)?;
Ok(Self {
arc_type: String::new(),
flags: flags as u8,
size: size as u64,
})
}
_ => Err(OpenFstError::InvalidFstHeader(format!(
"EncodeTableHeader::read: bad magic number {magic}"
))),
}
}
pub fn write<W: Write>(&self, writer: &mut W) -> Result<(), OpenFstError> {
write_scalar(writer, ENCODE_MAGIC_NUMBER)?;
write_string(writer, &self.arc_type)?;
write_scalar(writer, self.flags)?;
write_scalar(writer, self.size)?;
Ok(())
}
}
pub struct EncodeTable<L, W> {
flags: u8,
triples: CompactHashBiTable<usize, Triple<L, W>>,
isymbols: Option<AtomicRc<SymbolTable>>,
osymbols: Option<AtomicRc<SymbolTable>>,
}
impl<L, W> EncodeTable<L, W>
where
L: ArcLabel,
W: Weight + Hash + Eq,
{
pub fn new(flags: u8) -> Self {
Self {
flags,
triples: CompactHashBiTable::new(1024),
isymbols: None,
osymbols: None,
}
}
pub fn encode<A: Arc<Label = L, Weight = W>>(&mut self, arc: &A) -> L {
let triple =
if arc.nextstate() == A::StateId::no_state() && self.flags & ENCODE_WEIGHTS != 0 {
Triple {
ilabel: L::no_label(),
olabel: L::no_label(),
weight: arc.weight().clone(),
}
} else {
Triple::from_arc(arc, self.flags)
};
self.encode_triple(triple)
}
fn encode_triple(&mut self, triple: Triple<L, W>) -> L {
let id = self
.triples
.find_id(&triple, true)
.expect("find_id inserts when asked to");
L::from_i64(id as i64 + 1).unwrap_or_else(L::no_label)
}
pub fn decode(&self, label: L) -> Option<&Triple<L, W>> {
let index = label.to_i64()?.checked_sub(1)?;
self.triples.find_entry(usize::try_from(index).ok()?)
}
pub fn len(&self) -> usize {
self.triples.size()
}
pub fn is_empty(&self) -> bool {
self.triples.size() == 0
}
pub fn flags(&self) -> u8 {
self.flags & ENCODE_FLAGS
}
pub fn input_symbols(&self) -> Option<AtomicRc<SymbolTable>> {
self.isymbols.clone()
}
pub fn output_symbols(&self) -> Option<AtomicRc<SymbolTable>> {
self.osymbols.clone()
}
pub fn set_input_symbols(&mut self, syms: Option<AtomicRc<SymbolTable>>) {
match syms {
Some(syms) => {
self.isymbols = Some(syms);
self.flags |= ENCODE_HAS_ISYMBOLS;
}
None => {
self.isymbols = None;
self.flags &= !ENCODE_HAS_ISYMBOLS;
}
}
}
pub fn set_output_symbols(&mut self, syms: Option<AtomicRc<SymbolTable>>) {
match syms {
Some(syms) => {
self.osymbols = Some(syms);
self.flags |= ENCODE_HAS_OSYMBOLS;
}
None => {
self.osymbols = None;
self.flags &= !ENCODE_HAS_OSYMBOLS;
}
}
}
}
impl<L, W> EncodeTable<L, W>
where
L: ArcLabel + FstScalar,
W: Weight + WeightIo + Hash + Eq,
{
pub fn read<R: Read>(reader: &mut R) -> Result<Self, OpenFstError> {
let header = EncodeTableHeader::read(reader)?;
let mut table = Self::new(header.flags);
for _ in 0..header.size {
let triple = Triple {
ilabel: read_scalar(reader)?,
olabel: read_scalar(reader)?,
weight: W::read(reader)?,
};
table.encode_triple(triple);
}
if header.flags & ENCODE_HAS_ISYMBOLS != 0 {
table.isymbols = Some(AtomicRc::new(SymbolTable::read(reader)?));
}
if header.flags & ENCODE_HAS_OSYMBOLS != 0 {
table.osymbols = Some(AtomicRc::new(SymbolTable::read(reader)?));
}
Ok(table)
}
pub fn write<Wr: Write>(&self, writer: &mut Wr, arc_type: ArcType) -> Result<(), OpenFstError> {
EncodeTableHeader {
arc_type: arc_type.to_string(),
flags: self.flags,
size: self.len() as u64,
}
.write(writer)?;
for index in 0..self.len() {
let triple = self.triples.find_entry(index).expect("index is in range");
write_scalar(writer, triple.ilabel)?;
write_scalar(writer, triple.olabel)?;
triple.weight.write(writer)?;
}
if self.flags & ENCODE_HAS_ISYMBOLS != 0
&& let Some(syms) = &self.isymbols
{
syms.write(writer)?;
}
if self.flags & ENCODE_HAS_OSYMBOLS != 0
&& let Some(syms) = &self.osymbols
{
syms.write(writer)?;
}
Ok(())
}
}
pub struct EncodeMapper<A: Arc> {
flags: u8,
encode_type: EncodeType,
table: Rc<RefCell<EncodeTable<A::Label, A::Weight>>>,
error: bool,
}
impl<A: Arc> EncodeMapper<A>
where
A::Weight: Hash + Eq,
{
pub fn new(flags: u8) -> Self {
Self {
flags: flags & ENCODE_FLAGS,
encode_type: EncodeType::Encode,
table: Rc::new(RefCell::new(EncodeTable::new(flags & ENCODE_FLAGS))),
error: false,
}
}
pub fn inverse(&self) -> Self {
Self {
flags: self.flags,
encode_type: match self.encode_type {
EncodeType::Encode => EncodeType::Decode,
EncodeType::Decode => EncodeType::Encode,
},
table: Rc::clone(&self.table),
error: self.error,
}
}
pub fn from_table(
table: Rc<RefCell<EncodeTable<A::Label, A::Weight>>>,
encode_type: EncodeType,
) -> Self {
let flags = table.borrow().flags();
Self {
flags,
encode_type,
table,
error: false,
}
}
pub fn encode_type(&self) -> EncodeType {
self.encode_type
}
pub fn flags(&self) -> u8 {
self.flags
}
pub fn table(&self) -> &Rc<RefCell<EncodeTable<A::Label, A::Weight>>> {
&self.table
}
pub fn error(&self) -> bool {
self.error
}
pub fn input_symbols(&self) -> Option<AtomicRc<SymbolTable>> {
self.table.borrow().input_symbols()
}
pub fn output_symbols(&self) -> Option<AtomicRc<SymbolTable>> {
self.table.borrow().output_symbols()
}
pub fn set_input_symbols(&self, syms: Option<AtomicRc<SymbolTable>>) {
self.table.borrow_mut().set_input_symbols(syms);
}
pub fn set_output_symbols(&self, syms: Option<AtomicRc<SymbolTable>>) {
self.table.borrow_mut().set_output_symbols(syms);
}
fn is_superfinal(arc: &A) -> bool {
arc.nextstate() == A::StateId::no_state()
}
fn encode_arc(&mut self, arc: &A) -> A {
if Self::is_superfinal(arc)
&& (self.flags & ENCODE_WEIGHTS == 0 || *arc.weight() == A::Weight::zero())
{
return arc.clone();
}
let label = self.table.borrow_mut().encode(arc);
A::new(
label,
if self.flags & ENCODE_LABELS != 0 {
label
} else {
arc.olabel()
},
if self.flags & ENCODE_WEIGHTS != 0 {
A::Weight::one()
} else {
arc.weight().clone()
},
arc.nextstate(),
)
}
fn decode_arc(&mut self, arc: &A) -> A {
if Self::is_superfinal(arc) || arc.ilabel() == A::Label::epsilon() {
return arc.clone();
}
if self.flags & ENCODE_LABELS != 0 && arc.ilabel() != arc.olabel() {
self.error = true;
}
if self.flags & ENCODE_WEIGHTS != 0 && *arc.weight() != A::Weight::one() {
self.error = true;
}
let table = self.table.borrow();
let Some(triple) = table.decode(arc.ilabel()) else {
self.error = true;
return A::new(
A::Label::no_label(),
A::Label::no_label(),
A::Weight::no_weight(),
arc.nextstate(),
);
};
if triple.ilabel == A::Label::no_label() {
return A::new(
A::Label::epsilon(),
A::Label::epsilon(),
triple.weight.clone(),
arc.nextstate(),
);
}
A::new(
triple.ilabel,
if self.flags & ENCODE_LABELS != 0 {
triple.olabel
} else {
arc.olabel()
},
if self.flags & ENCODE_WEIGHTS != 0 {
triple.weight.clone()
} else {
arc.weight().clone()
},
arc.nextstate(),
)
}
}
impl<A: Arc> ArcMapper<A, A> for EncodeMapper<A>
where
A::Weight: Hash + Eq,
{
fn map(&mut self, arc: &A) -> A {
match self.encode_type {
EncodeType::Encode => self.encode_arc(arc),
EncodeType::Decode => self.decode_arc(arc),
}
}
fn final_action(&self) -> MapFinalAction {
if self.encode_type == EncodeType::Encode && self.flags & ENCODE_WEIGHTS != 0 {
MapFinalAction::RequireSuperfinal
} else {
MapFinalAction::NoSuperfinal
}
}
fn input_symbols_action(&self) -> MapSymbolsAction {
MapSymbolsAction::Clear
}
fn output_symbols_action(&self) -> MapSymbolsAction {
MapSymbolsAction::Clear
}
fn properties(&self, inprops: u64) -> u64 {
let mut outprops = inprops;
if self.error {
outprops |= K_ERROR;
}
let mut mask = K_FST_PROPERTIES;
if self.flags & ENCODE_LABELS != 0 {
mask &= K_I_LABEL_INVARIANT_PROPERTIES & K_O_LABEL_INVARIANT_PROPERTIES;
}
if self.flags & ENCODE_WEIGHTS != 0 {
mask &= K_I_LABEL_INVARIANT_PROPERTIES
& K_WEIGHT_INVARIANT_PROPERTIES
& if self.encode_type == EncodeType::Encode {
K_ADD_SUPER_FINAL_PROPERTIES
} else {
K_RM_SUPER_FINAL_PROPERTIES
};
}
if self.encode_type == EncodeType::Encode {
mask |= K_I_DETERMINISTIC;
}
outprops &= mask;
if self.encode_type == EncodeType::Encode {
if self.flags & ENCODE_LABELS != 0 {
outprops |= K_ACCEPTOR;
}
if self.flags & ENCODE_WEIGHTS != 0 {
outprops |= K_UNWEIGHTED | K_UNWEIGHTED_CYCLES;
}
}
outprops
}
}
pub fn encode<A, F>(fst: &mut F, mapper: &mut EncodeMapper<A>) -> Result<(), OpenFstError>
where
A: Arc,
A::Weight: Hash + Eq,
F: MutableFst<A>,
{
mapper.set_input_symbols(fst.input_symbols());
mapper.set_output_symbols(fst.output_symbols());
arc_map(fst, mapper)
}
pub fn decode<A, F>(fst: &mut F, mapper: &EncodeMapper<A>) -> Result<(), OpenFstError>
where
A: Arc,
A::Weight: Hash + Eq,
F: MutableFst<A>,
{
let mut decoder = mapper.inverse();
debug_assert_eq!(decoder.encode_type(), EncodeType::Decode);
check_decodable(fst, &decoder)?;
arc_map(fst, &mut decoder)?;
rm_final_epsilon(fst);
fst.set_input_symbols(mapper.input_symbols());
fst.set_output_symbols(mapper.output_symbols());
Ok(())
}
fn check_decodable<A, F>(fst: &F, decoder: &EncodeMapper<A>) -> Result<(), OpenFstError>
where
A: Arc,
A::Weight: Hash + Eq,
F: Fst<A>,
{
let flags = decoder.flags();
let table = decoder.table.borrow();
for state in fst.states() {
for arc in fst.arcs(state) {
if arc.ilabel() == A::Label::epsilon() {
continue;
}
if flags & ENCODE_LABELS != 0 && arc.ilabel() != arc.olabel() {
return Err(OpenFstError::InvalidOperation(format!(
"Decode: label-encoded arc from state {:?} has different input and output \
labels: {} and {}",
state,
arc.ilabel(),
arc.olabel()
)));
}
if flags & ENCODE_WEIGHTS != 0 && *arc.weight() != A::Weight::one() {
return Err(OpenFstError::InvalidOperation(format!(
"Decode: weight-encoded arc from state {:?} has non-trivial weight {}",
state,
arc.weight()
)));
}
if table.decode(arc.ilabel()).is_none() {
return Err(OpenFstError::InvalidOperation(format!(
"Decode: arc from state {:?} carries label {}, which the encode table does \
not have",
state,
arc.ilabel()
)));
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::algorithms::test_support::{Rng, paths, random_acyclic_fst, string_weights};
use crate::arc::StdArc;
use crate::fst::ExpandedFst as _;
use crate::fsts::vector_fst::StdVectorFst;
use crate::properties::{K_ACCEPTOR, K_FST_PROPERTIES, K_UNWEIGHTED};
use crate::weights::float_weight::TropicalWeight;
fn mapper(flags: u8) -> EncodeMapper<StdArc> {
EncodeMapper::new(flags)
}
fn chain() -> StdVectorFst {
let mut fst = StdVectorFst::new();
for _ in 0..3 {
fst.add_state();
}
fst.set_start(0);
fst.add_arc(0, StdArc::new(1, 10, TropicalWeight(1.0), 1));
fst.add_arc(1, StdArc::new(2, 20, TropicalWeight(2.0), 2));
fst.set_final(2, TropicalWeight(3.0));
fst
}
fn observable(fst: &StdVectorFst) -> Vec<(Vec<i32>, Vec<i32>, String)> {
string_weights(paths(fst, 12))
}
#[test]
fn encoding_and_decoding_gives_back_the_same_paths() {
let mut rng = Rng::new(0x_E4C0_DEED);
for round in 0..200 {
let fst = random_acyclic_fst(&mut rng, 6);
let before = observable(&fst);
for flags in [ENCODE_LABELS, ENCODE_WEIGHTS, ENCODE_FLAGS] {
let mut encoder = mapper(flags);
let mut copy = fst.clone();
encode(&mut copy, &mut encoder).unwrap();
decode(&mut copy, &encoder).unwrap();
assert_eq!(observable(©), before, "round {round}, flags {flags}");
}
}
}
#[test]
fn encoding_labels_makes_an_acceptor() {
let mut fst = chain();
assert_eq!(fst.properties(K_ACCEPTOR, true) & K_ACCEPTOR, 0);
let mut encoder = mapper(ENCODE_LABELS);
encode(&mut fst, &mut encoder).unwrap();
for state in 0..fst.num_states() as i32 {
for arc in fst.arcs(state) {
assert_eq!(arc.ilabel(), arc.olabel());
}
}
assert_ne!(fst.properties(K_ACCEPTOR, true) & K_ACCEPTOR, 0);
assert_eq!(fst.final_weight(2), TropicalWeight(3.0));
}
#[test]
fn encoding_weights_makes_it_unweighted() {
let mut fst = chain();
let mut encoder = mapper(ENCODE_WEIGHTS);
encode(&mut fst, &mut encoder).unwrap();
for state in 0..fst.num_states() as i32 {
for arc in fst.arcs(state) {
assert_eq!(*arc.weight(), TropicalWeight::one());
}
}
assert_ne!(fst.properties(K_UNWEIGHTED, true) & K_UNWEIGHTED, 0);
let olabels: Vec<i32> = (0..fst.num_states() as i32)
.flat_map(|s| fst.arcs(s).map(|a| a.olabel()).collect::<Vec<_>>())
.collect();
assert_eq!(olabels, vec![10, 20, 0]);
}
#[test]
fn a_final_weight_becomes_an_arc_and_comes_back() {
let mut fst = chain();
let before = fst.num_states();
let mut encoder = mapper(ENCODE_WEIGHTS);
encode(&mut fst, &mut encoder).unwrap();
assert_eq!(fst.num_states(), before + 1, "a superfinal state was added");
assert_eq!(
fst.final_weight(2),
TropicalWeight::zero(),
"state 2 is no longer final; its weight left on an arc"
);
assert_eq!(fst.num_arcs(2), 1);
decode(&mut fst, &encoder).unwrap();
assert_eq!(fst.final_weight(2), TropicalWeight(3.0));
assert_eq!(fst.num_arcs(2), 0);
}
#[test]
fn a_final_state_at_weight_one_keeps_its_final_weight() {
let mut fst = chain();
fst.set_final(2, TropicalWeight::one());
let mut encoder = mapper(ENCODE_WEIGHTS);
encode(&mut fst, &mut encoder).unwrap();
assert_eq!(encoder.table().borrow().len(), 3);
decode(&mut fst, &encoder).unwrap();
assert_eq!(fst.final_weight(2), TropicalWeight::one());
}
#[test]
fn arcs_share_a_label_exactly_when_what_is_encoded_matches() {
let mut fst = StdVectorFst::new();
for _ in 0..2 {
fst.add_state();
}
fst.set_start(0);
fst.add_arc(0, StdArc::new(1, 10, TropicalWeight(1.0), 1));
fst.add_arc(0, StdArc::new(1, 10, TropicalWeight(1.0), 1)); fst.add_arc(0, StdArc::new(1, 10, TropicalWeight(2.0), 1)); fst.add_arc(0, StdArc::new(1, 99, TropicalWeight(1.0), 1));
let mut labels = fst.clone();
let mut encoder = mapper(ENCODE_LABELS);
encode(&mut labels, &mut encoder).unwrap();
let got: Vec<i32> = labels.arcs(0).map(|a| a.ilabel()).collect();
assert_eq!(
got,
vec![1, 1, 1, 2],
"only the output label separates them when weights are not encoded"
);
let mut weights = fst.clone();
let mut encoder = mapper(ENCODE_WEIGHTS);
encode(&mut weights, &mut encoder).unwrap();
let got: Vec<i32> = weights.arcs(0).map(|a| a.ilabel()).collect();
assert_eq!(
got,
vec![1, 1, 2, 1],
"only the weight separates them when labels are not encoded"
);
let mut both = fst;
let mut encoder = mapper(ENCODE_FLAGS);
encode(&mut both, &mut encoder).unwrap();
let got: Vec<i32> = both.arcs(0).map(|a| a.ilabel()).collect();
assert_eq!(got, vec![1, 1, 2, 3]);
}
#[test]
fn symbol_tables_are_kept_aside_and_put_back() {
let mut syms = SymbolTable::new("input");
syms.add_symbol("a", 1);
let mut osyms = SymbolTable::new("output");
osyms.add_symbol("A", 1);
let mut fst = chain();
fst.set_input_symbols(Some(AtomicRc::new(syms)));
fst.set_output_symbols(Some(AtomicRc::new(osyms)));
let mut encoder = mapper(ENCODE_FLAGS);
encode(&mut fst, &mut encoder).unwrap();
assert!(fst.input_symbols().is_none());
assert!(fst.output_symbols().is_none());
decode(&mut fst, &encoder).unwrap();
assert_eq!(fst.input_symbols().unwrap().name(), "input");
assert_eq!(fst.output_symbols().unwrap().name(), "output");
}
#[test]
fn decoding_a_label_the_table_does_not_have_is_refused() {
let mut fst = chain();
let mut encoder = mapper(ENCODE_FLAGS);
encode(&mut fst, &mut encoder).unwrap();
let past_the_end = encoder.table().borrow().len() as i32 + 1;
fst.add_arc(
0,
StdArc::new(past_the_end, past_the_end, TropicalWeight::one(), 1),
);
let before: Vec<StdArc> = fst.arcs(0).collect();
let err = decode(&mut fst, &encoder).unwrap_err();
assert!(format!("{err}").contains("encode table"), "{err}");
assert_eq!(
fst.arcs(0).collect::<Vec<_>>(),
before,
"nothing was rewritten"
);
}
#[test]
fn decoding_refuses_an_arc_that_was_never_encoded_that_way() {
let mut fst = chain();
let mut encoder = mapper(ENCODE_FLAGS);
encode(&mut fst, &mut encoder).unwrap();
fst.add_arc(0, StdArc::new(1, 2, TropicalWeight::one(), 1));
assert!(decode(&mut fst, &encoder).is_err());
let mut fst = chain();
let mut encoder = mapper(ENCODE_WEIGHTS);
encode(&mut fst, &mut encoder).unwrap();
fst.add_arc(0, StdArc::new(1, 1, TropicalWeight(7.0), 1));
assert!(decode(&mut fst, &encoder).is_err());
}
#[test]
fn a_decoder_shares_the_table_it_was_made_from() {
let mut encoder = mapper(ENCODE_FLAGS);
let decoder = encoder.inverse();
assert_eq!(decoder.encode_type(), EncodeType::Decode);
assert!(decoder.table().borrow().is_empty());
let mut fst = chain();
encode(&mut fst, &mut encoder).unwrap();
assert_eq!(
decoder.table().borrow().len(),
encoder.table().borrow().len()
);
assert!(!decoder.table().borrow().is_empty());
}
#[test]
fn the_claimed_properties_are_the_ones_the_result_has() {
let mut rng = Rng::new(0x_9E0_9E0);
for round in 0..100 {
let fst = random_acyclic_fst(&mut rng, 5);
let inprops = fst.properties(K_FST_PROPERTIES, true);
for flags in [ENCODE_LABELS, ENCODE_WEIGHTS, ENCODE_FLAGS] {
let mut encoder = mapper(flags);
let mut copy = fst.clone();
let claimed = encoder.properties(inprops);
encode(&mut copy, &mut encoder).unwrap();
let actual = copy.properties(K_FST_PROPERTIES, true);
assert_eq!(
claimed & !actual & K_FST_PROPERTIES,
0,
"round {round}, flags {flags}: claimed a property the result does not have: \
{:#x}",
claimed & !actual
);
}
}
}
#[test]
fn the_table_matches_the_bytes_openfst_writes() {
let mut table: EncodeTable<i32, TropicalWeight> = EncodeTable::new(ENCODE_FLAGS);
table.encode(&StdArc::new(1, 2, TropicalWeight(0.5), 1));
table.encode(&StdArc::new(3, 4, TropicalWeight(1.5), 1));
let mut bytes = Vec::new();
table.write(&mut bytes, ArcType::STANDARD).unwrap();
#[rustfmt::skip]
let golden: [u8; 49] = [
0x4a, 0x6d, 0xd9, 0x7e, 0x08, 0x00, 0x00, 0x00, b's', b't', b'a', b'n', b'd', b'a', b'r', b'd', 0x03, 0x02, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x02, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x3f, 0x03, 0x00, 0x00, 0x00, 0x04, 0x00, 0x00, 0x00, 0x00, 0x00, 0xc0, 0x3f, ];
assert_eq!(bytes, golden);
}
#[test]
fn a_table_round_trips_through_bytes() {
let mut table: EncodeTable<i32, TropicalWeight> = EncodeTable::new(ENCODE_FLAGS);
for (ilabel, olabel, weight) in [(1, 2, 0.5), (3, 4, 1.5), (1, 2, 2.5)] {
table.encode(&StdArc::new(ilabel, olabel, TropicalWeight(weight), 1));
}
let mut syms = SymbolTable::new("input");
syms.add_symbol("a", 1);
table.set_input_symbols(Some(AtomicRc::new(syms)));
let mut bytes = Vec::new();
table.write(&mut bytes, ArcType::STANDARD).unwrap();
let read: EncodeTable<i32, TropicalWeight> =
EncodeTable::read(&mut bytes.as_slice()).unwrap();
assert_eq!(read.len(), table.len());
assert_eq!(read.flags(), table.flags());
for label in 1..=table.len() as i32 {
assert_eq!(read.decode(label), table.decode(label), "label {label}");
}
assert_eq!(read.input_symbols().unwrap().name(), "input");
assert!(read.output_symbols().is_none());
}
#[test]
fn the_deprecated_header_is_still_readable() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&ENCODE_DEPRECATED_MAGIC_NUMBER.to_le_bytes());
bytes.extend_from_slice(&(ENCODE_FLAGS as u32).to_le_bytes());
bytes.extend_from_slice(&1i64.to_le_bytes());
bytes.extend_from_slice(&7i32.to_le_bytes());
bytes.extend_from_slice(&8i32.to_le_bytes());
bytes.extend_from_slice(&0.25f32.to_le_bytes());
let table: EncodeTable<i32, TropicalWeight> =
EncodeTable::read(&mut bytes.as_slice()).unwrap();
assert_eq!(table.len(), 1);
assert_eq!(
table.decode(1),
Some(&Triple {
ilabel: 7,
olabel: 8,
weight: TropicalWeight(0.25)
})
);
}
#[test]
fn a_stream_that_is_not_an_encode_table_is_refused() {
let bytes = 12345i32.to_le_bytes();
assert!(EncodeTable::<i32, TropicalWeight>::read(&mut bytes.as_slice()).is_err());
}
#[test]
fn a_table_claiming_more_triples_than_it_has_is_refused() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&ENCODE_MAGIC_NUMBER.to_le_bytes());
bytes.extend_from_slice(&8i32.to_le_bytes());
bytes.extend_from_slice(b"standard");
bytes.push(ENCODE_FLAGS);
bytes.extend_from_slice(&u64::MAX.to_le_bytes());
assert!(EncodeTable::<i32, TropicalWeight>::read(&mut bytes.as_slice()).is_err());
}
}