use crate::algorithms::connect::condense;
use crate::algorithms::dfs_visit::{DfsVisitor, dfs_visit_any};
use crate::arc::{Arc, ArcStateId};
use crate::data_structures::interval_set::{IntInterval, IntervalSet};
use crate::error::OpenFstError;
use crate::fst::Fst;
use crate::fsts::vector_fst::VectorFst;
use crate::properties::K_ACYCLIC;
use crate::weight::Weight;
pub type Index = i64;
const NO_INDEX: Index = -1;
struct IntervalReachVisitor<'f, A: Arc, F: Fst<A>> {
fst: &'f F,
isets: Vec<IntervalSet<Index>>,
state2index: Vec<Index>,
next_index: Option<Index>,
error: Option<&'static str>,
_marker: std::marker::PhantomData<A>,
}
impl<'f, A: Arc, F: Fst<A>> IntervalReachVisitor<'f, A, F> {
fn new(fst: &'f F) -> Self {
Self {
fst,
isets: Vec::new(),
state2index: Vec::new(),
next_index: Some(1),
error: None,
_marker: std::marker::PhantomData,
}
}
fn ensure(&mut self, s: usize) {
if self.isets.len() <= s {
self.isets.resize_with(s + 1, IntervalSet::new);
}
if self.state2index.len() <= s {
self.state2index.resize(s + 1, NO_INDEX);
}
}
}
impl<A: Arc, F: Fst<A>> DfsVisitor<A> for IntervalReachVisitor<'_, A, F> {
fn init_visit<G: Fst<A>>(&mut self, _fst: &G) {
self.error = None;
}
fn init_state(&mut self, s: A::StateId, _root: A::StateId) -> bool {
let idx = s.as_usize();
self.ensure(idx);
if self.fst.final_weight(s) == A::Weight::zero() {
return true;
}
match self.next_index {
Some(index) => {
self.isets[idx]
.intervals_mut()
.push(IntInterval::new(index, index + 1));
self.state2index[idx] = index;
self.next_index = Some(index + 1);
}
None => {
if self.fst.num_arcs(s) > 0 {
self.error =
Some("a supplied numbering requires the final states to have no arcs");
return false;
}
let index = self.state2index[idx];
if index == NO_INDEX {
self.error = Some("the supplied numbering is incomplete");
return false;
}
self.isets[idx]
.intervals_mut()
.push(IntInterval::new(index, index + 1));
}
}
true
}
#[inline]
fn tree_arc(&mut self, _s: A::StateId, _arc: &A) -> bool {
true
}
fn back_arc(&mut self, _s: A::StateId, _arc: &A) -> bool {
self.error = Some("the FST has a cycle");
false
}
fn forward_or_cross_arc(&mut self, s: A::StateId, arc: &A) -> bool {
let (from, to) = (s.as_usize(), arc.nextstate().as_usize());
self.ensure(from.max(to));
let reached = std::mem::take(&mut self.isets[to]);
self.isets[from].union(&reached);
self.isets[to] = reached;
true
}
fn finish_state(&mut self, s: A::StateId, parent: Option<A::StateId>, _arc: Option<&A>) {
let idx = s.as_usize();
self.ensure(idx);
if let Some(index) = self.next_index
&& self.fst.final_weight(s) != A::Weight::zero()
{
self.isets[idx].intervals_mut()[0].end = index;
}
self.isets[idx].normalize();
if let Some(parent) = parent {
let parent_idx = parent.as_usize();
self.ensure(parent_idx);
let reached = std::mem::take(&mut self.isets[idx]);
self.isets[parent_idx].union(&reached);
self.isets[idx] = reached;
}
}
fn finish_visit(&mut self) {}
}
pub struct StateReachable {
isets: Vec<IntervalSet<Index>>,
state2index: Vec<Index>,
}
impl StateReachable {
pub fn new<A: Arc, F: Fst<A>>(fst: &F) -> Result<Self, OpenFstError> {
if fst.properties(K_ACYCLIC, true) & K_ACYCLIC != 0 {
Self::acyclic(fst)
} else {
Self::cyclic(fst)
}
}
fn acyclic<A: Arc, F: Fst<A>>(fst: &F) -> Result<Self, OpenFstError> {
let mut visitor = IntervalReachVisitor::new(fst);
dfs_visit_any(fst, &mut visitor);
if let Some(reason) = visitor.error {
return Err(OpenFstError::InvalidOperation(format!(
"StateReachable: {reason}"
)));
}
Ok(Self {
isets: visitor.isets,
state2index: visitor.state2index,
})
}
fn cyclic<A: Arc, F: Fst<A>>(fst: &F) -> Result<Self, OpenFstError> {
let mut condensed = VectorFst::<A>::new();
let mut scc: Vec<A::StateId> = Vec::new();
condense(fst, &mut condensed, &mut scc);
let reachable = Self::new(&condensed)?;
let mut component_size: Vec<usize> = Vec::new();
for &c in &scc {
let c = c.as_usize();
if component_size.len() <= c {
component_size.resize(c + 1, 0);
}
component_size[c] += 1;
}
let mut isets = vec![IntervalSet::new(); scc.len()];
let mut state2index = vec![NO_INDEX; scc.len()];
for (s, &c) in scc.iter().enumerate() {
let c = c.as_usize();
isets[s] = reachable.isets[c].clone();
state2index[s] = reachable.state2index[c];
if condensed.final_weight(A::StateId::from_usize(c)) != A::Weight::zero()
&& component_size[c] > 1
{
return Err(OpenFstError::InvalidOperation(
"StateReachable: a final state is contained in a cycle".to_string(),
));
}
}
Ok(Self { isets, state2index })
}
pub fn reach<S: ArcStateId>(&self, from: S, to: S) -> bool {
let (from, to) = (from.as_usize(), to.as_usize());
let Some(&index) = self.state2index.get(to) else {
return false;
};
if index == NO_INDEX {
return false;
}
self.isets.get(from).is_some_and(|iset| iset.member(index))
}
pub fn state2index(&self) -> &[Index] {
&self.state2index
}
pub fn interval_sets(&self) -> &[IntervalSet<Index>] {
&self.isets
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::arc::StdArc;
use crate::fst::{ExpandedFst as _, MutableFst};
use crate::fsts::vector_fst::StdVectorFst;
use crate::weights::float_weight::TropicalWeight;
fn build(nstates: usize, edges: &[(i32, i32)], finals: &[i32]) -> StdVectorFst {
let mut fst = StdVectorFst::new();
for _ in 0..nstates {
fst.add_state();
}
fst.set_start(0);
for &(from, to) in edges {
fst.add_arc(from, StdArc::new(1, 1, TropicalWeight::one(), to));
}
for &s in finals {
fst.set_final(s, TropicalWeight::one());
}
fst
}
fn brute_force(fst: &StdVectorFst, from: usize, to: usize) -> bool {
let n = fst.num_states();
let mut seen = vec![false; n];
let mut stack = vec![from];
seen[from] = true;
while let Some(s) = stack.pop() {
if s == to {
return true;
}
for arc in fst.arcs(s as i32) {
let next = arc.nextstate() as usize;
if !seen[next] {
seen[next] = true;
stack.push(next);
}
}
}
false
}
fn assert_matches_reachability(fst: &StdVectorFst) {
let reachable = StateReachable::new(fst).expect("acyclic or condensable");
let n = fst.num_states();
for from in 0..n {
for to in 0..n {
let is_final = fst.final_weight(to as i32) != TropicalWeight::zero();
let want = is_final && brute_force(fst, from, to);
assert_eq!(
reachable.reach(from as i32, to as i32),
want,
"{from} -> {to}"
);
}
}
}
#[test]
fn a_chain_reaches_the_final_states_after_it() {
assert_matches_reachability(&build(4, &[(0, 1), (1, 2), (2, 3)], &[1, 3]));
}
#[test]
fn a_branching_fst_reaches_the_final_states_down_each_branch() {
assert_matches_reachability(&build(5, &[(0, 1), (0, 2), (1, 3), (2, 4)], &[3, 4]));
}
#[test]
fn a_shared_final_state_is_reachable_from_both_sides() {
assert_matches_reachability(&build(4, &[(0, 1), (0, 2), (1, 3), (2, 3)], &[3]));
}
#[test]
fn a_state_that_is_not_final_is_never_reported_reachable() {
let fst = build(3, &[(0, 1), (1, 2)], &[2]);
let reachable = StateReachable::new(&fst).unwrap();
assert!(reachable.reach(0, 2));
assert!(!reachable.reach(0, 1));
assert_eq!(reachable.state2index()[1], NO_INDEX);
}
#[test]
fn a_cyclic_fst_is_answered_through_its_components() {
let fst = build(4, &[(0, 1), (1, 2), (2, 1), (2, 3)], &[3]);
let reachable = StateReachable::new(&fst).unwrap();
assert!(reachable.reach(0, 3));
assert!(reachable.reach(1, 3));
assert!(reachable.reach(2, 3));
assert!(reachable.reach(3, 3));
assert!(!reachable.reach(3, 0), "nothing leads back out of state 3");
}
#[test]
fn a_final_state_inside_a_cycle_is_refused() {
let fst = build(3, &[(0, 1), (1, 2), (2, 1)], &[1]);
assert!(StateReachable::new(&fst).is_err());
}
#[test]
fn an_fst_with_no_final_states_reaches_nothing() {
let fst = build(3, &[(0, 1), (1, 2)], &[]);
let reachable = StateReachable::new(&fst).unwrap();
for from in 0..3 {
for to in 0..3 {
assert!(!reachable.reach(from, to));
}
}
}
}