use crate::arc::Arc;
use crate::fst::{Fst, MutableFst, StateId};
use crate::semiring::Semiring;
use crate::{Error, Result};
use rustc_hash::{FxHashMap, FxHashSet};
pub fn connect<W, F, M>(fst: &F) -> Result<M>
where
W: Semiring,
F: Fst<W>,
M: MutableFst<W> + Default,
{
let start = fst
.start()
.ok_or_else(|| Error::Algorithm("FST has no start state".into()))?;
let accessible = find_accessible_states(fst, start);
let coaccessible = find_coaccessible_states(fst);
let keep: FxHashSet<StateId> = accessible.intersection(&coaccessible).cloned().collect();
if keep.is_empty() {
return Ok(M::default());
}
let mut result = M::default();
let mut state_map = vec![None; fst.num_states()];
for &state in &keep {
let new_state = result.add_state();
state_map[state as usize] = Some(new_state);
}
if let Some(new_start) = state_map[start as usize] {
result.set_start(new_start);
}
for &state in &keep {
if let Some(new_state) = state_map[state as usize] {
if let Some(weight) = fst.final_weight(state) {
result.set_final(new_state, weight.clone());
}
for arc in fst.arcs(state) {
if keep.contains(&arc.nextstate) {
if let Some(new_nextstate) = state_map[arc.nextstate as usize] {
result.add_arc(
new_state,
Arc::new(arc.ilabel, arc.olabel, arc.weight.clone(), new_nextstate),
);
}
}
}
}
}
Ok(result)
}
fn find_accessible_states<W: Semiring, F: Fst<W>>(fst: &F, start: StateId) -> FxHashSet<StateId> {
let mut accessible = FxHashSet::default();
let mut stack = vec![start];
while let Some(state) = stack.pop() {
if accessible.insert(state) {
for arc in fst.arcs(state) {
stack.push(arc.nextstate);
}
}
}
accessible
}
fn find_coaccessible_states<W: Semiring, F: Fst<W>>(fst: &F) -> FxHashSet<StateId> {
let mut predecessors: FxHashMap<StateId, Vec<StateId>> = FxHashMap::default();
for state in fst.states() {
for arc in fst.arcs(state) {
predecessors.entry(arc.nextstate).or_default().push(state);
}
}
let mut coaccessible = FxHashSet::default();
let mut stack = Vec::new();
for state in fst.states() {
if fst.is_final(state) {
stack.push(state);
}
}
while let Some(state) = stack.pop() {
if coaccessible.insert(state) {
if let Some(preds) = predecessors.get(&state) {
for &pred in preds {
if !coaccessible.contains(&pred) {
stack.push(pred);
}
}
}
}
}
coaccessible
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use num_traits::One;
#[test]
fn test_connect_removes_unreachable() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
let s3 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s2, Arc::new(2, 2, TropicalWeight::new(1.0), s3));
let connected: VectorFst<TropicalWeight> = connect(&fst).unwrap();
assert!(connected.num_states() < fst.num_states());
assert!(connected.start().is_some());
}
#[test]
fn test_connect_already_connected() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
let connected: VectorFst<TropicalWeight> = connect(&fst).unwrap();
assert_eq!(connected.num_states(), fst.num_states());
assert!(connected.start().is_some());
}
}