use crate::arc::{Arc, ArcIterator};
use crate::fst::{Fst, Label, StateId};
use crate::semiring::Semiring;
use std::collections::{HashMap, HashSet};
#[derive(Debug)]
pub struct FailureFst<F: Fst<W>, W: Semiring> {
inner: F,
failure_map: HashMap<StateId, StateId>, _phantom: std::marker::PhantomData<W>,
}
pub struct FailureArcIterator<'a, W: Semiring, F: Fst<W>> {
inner: &'a F,
_failure_map: &'a HashMap<StateId, StateId>, current_state: StateId,
inner_arcs: <F as Fst<W>>::ArcIter<'a>,
failure_chain: Vec<StateId>, _visited_failures: HashSet<StateId>, failure_arcs: Option<<F as Fst<W>>::ArcIter<'a>>,
failure_state_idx: usize,
}
pub struct FailureMatchingArcIterator<'a, W: Semiring, F: Fst<W>> {
inner: &'a F,
#[allow(dead_code)] failure_map: &'a HashMap<StateId, StateId>,
current_state: StateId,
target_ilabel: Label,
inner_arcs: <F as Fst<W>>::ArcIter<'a>,
failure_chain: Vec<StateId>,
failure_state_idx: usize,
found_match: bool,
}
impl<'a, W: Semiring, F: Fst<W>> std::fmt::Debug for FailureMatchingArcIterator<'a, W, F> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FailureMatchingArcIterator")
.field("current_state", &self.current_state)
.field("target_ilabel", &self.target_ilabel)
.field("failure_chain", &self.failure_chain)
.field("failure_state_idx", &self.failure_state_idx)
.field("found_match", &self.found_match)
.finish_non_exhaustive()
}
}
impl<'a, W: Semiring, F: Fst<W>> std::fmt::Debug for FailureArcIterator<'a, W, F> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FailureArcIterator")
.field("current_state", &self.current_state)
.field("failure_chain", &self.failure_chain)
.field("failure_state_idx", &self.failure_state_idx)
.finish_non_exhaustive()
}
}
impl<'a, W: Semiring, F: Fst<W>> FailureArcIterator<'a, W, F> {
fn new(inner: &'a F, failure_map: &'a HashMap<StateId, StateId>, state: StateId) -> Self {
let inner_arcs = inner.arcs(state);
let mut failure_chain = Vec::new();
let mut visited = HashSet::new();
let mut current = state;
while let Some(&failure_state) = failure_map.get(¤t) {
if visited.contains(&failure_state) {
break; }
visited.insert(failure_state);
failure_chain.push(failure_state);
current = failure_state;
}
Self {
inner,
_failure_map: failure_map,
current_state: state,
inner_arcs,
failure_chain,
_visited_failures: visited,
failure_arcs: None,
failure_state_idx: 0,
}
}
}
impl<'a, W: Semiring, F: Fst<W>> Iterator for FailureArcIterator<'a, W, F> {
type Item = Arc<W>;
fn next(&mut self) -> Option<Self::Item> {
if let Some(arc) = self.inner_arcs.next() {
return Some(arc);
}
while self.failure_state_idx < self.failure_chain.len() {
let failure_state = self.failure_chain[self.failure_state_idx];
if self.failure_arcs.is_none() {
self.failure_arcs = Some(self.inner.arcs(failure_state));
}
if let Some(ref mut iter) = self.failure_arcs {
if let Some(arc) = iter.next() {
return Some(arc);
}
}
self.failure_arcs = None;
self.failure_state_idx += 1;
}
None
}
}
impl<'a, W: Semiring, F: Fst<W>> ArcIterator<W> for FailureArcIterator<'a, W, F> {
fn reset(&mut self) {
self.inner_arcs = self.inner.arcs(self.current_state);
self.failure_arcs = None;
self.failure_state_idx = 0;
}
}
impl<'a, W: Semiring, F: Fst<W>> FailureMatchingArcIterator<'a, W, F> {
fn new(
inner: &'a F,
failure_map: &'a HashMap<StateId, StateId>,
state: StateId,
target_ilabel: Label,
) -> Self {
let inner_arcs = inner.arcs(state);
let mut failure_chain = Vec::new();
let mut visited = HashSet::new();
let mut current = state;
while let Some(&failure_state) = failure_map.get(¤t) {
if visited.contains(&failure_state) {
break; }
visited.insert(failure_state);
failure_chain.push(failure_state);
current = failure_state;
}
Self {
inner,
failure_map,
current_state: state,
target_ilabel,
inner_arcs,
failure_chain,
failure_state_idx: 0,
found_match: false,
}
}
}
impl<'a, W: Semiring, F: Fst<W>> Iterator for FailureMatchingArcIterator<'a, W, F> {
type Item = Arc<W>;
fn next(&mut self) -> Option<Self::Item> {
for arc in self.inner_arcs.by_ref() {
if arc.ilabel == self.target_ilabel {
self.found_match = true;
return Some(arc);
}
}
if self.found_match {
return None;
}
while self.failure_state_idx < self.failure_chain.len() {
let failure_state = self.failure_chain[self.failure_state_idx];
let failure_arcs = self.inner.arcs(failure_state);
for arc in failure_arcs {
if arc.ilabel == self.target_ilabel {
self.found_match = true;
return Some(arc);
}
}
self.failure_state_idx += 1;
}
None
}
}
impl<F: Fst<W>, W: Semiring> FailureFst<F, W> {
pub fn new(fst: F) -> Self {
Self {
inner: fst,
failure_map: HashMap::new(),
_phantom: std::marker::PhantomData,
}
}
pub fn set_failure(&mut self, state: StateId, failure_state: StateId) {
self.failure_map.insert(state, failure_state);
}
pub fn failure_state(&self, state: StateId) -> Option<StateId> {
self.failure_map.get(&state).copied()
}
pub fn arcs_matching(
&self,
state: StateId,
ilabel: Label,
) -> FailureMatchingArcIterator<'_, W, F> {
FailureMatchingArcIterator::new(&self.inner, &self.failure_map, state, ilabel)
}
}
impl<F: Fst<W>, W: Semiring> Fst<W> for FailureFst<F, W> {
type ArcIter<'a>
= FailureArcIterator<'a, W, F>
where
Self: 'a;
fn start(&self) -> Option<StateId> {
self.inner.start()
}
fn final_weight(&self, state: StateId) -> Option<&W> {
self.inner.final_weight(state)
}
fn num_arcs(&self, state: StateId) -> usize {
let mut count = self.inner.num_arcs(state);
let mut visited = HashSet::new();
let mut current = state;
while let Some(&failure_state) = self.failure_map.get(¤t) {
if visited.contains(&failure_state) {
break; }
visited.insert(failure_state);
count += self.inner.num_arcs(failure_state);
current = failure_state;
}
count
}
fn num_states(&self) -> usize {
self.inner.num_states()
}
fn properties(&self) -> crate::properties::FstProperties {
self.inner.properties()
}
fn arcs(&self, state: StateId) -> Self::ArcIter<'_> {
FailureArcIterator::new(&self.inner, &self.failure_map, state)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_failure_fst_new() {
let fst = VectorFst::<TropicalWeight>::new();
let failure_fst = FailureFst::new(fst);
assert_eq!(failure_fst.num_states(), 0);
}
#[test]
fn test_failure_fst_set_failure() {
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());
let mut failure_fst = FailureFst::new(fst);
failure_fst.set_failure(s1, s0);
assert_eq!(failure_fst.failure_state(s1), Some(s0));
}
#[test]
fn test_failure_fst_delegates_to_inner() {
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::new(0.5));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
let failure_fst = FailureFst::new(fst);
assert_eq!(failure_fst.start(), Some(s0));
assert!(failure_fst.is_final(s1));
assert_eq!(failure_fst.num_arcs(s0), 1);
}
#[test]
fn test_aho_corasick_conditional_failure() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
fst.set_start(s0);
fst.set_final(s2, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
let mut failure_fst = FailureFst::new(fst);
failure_fst.set_failure(s1, s0);
let matching: Vec<_> = failure_fst.arcs_matching(s1, 1).collect();
assert_eq!(matching.len(), 1);
assert_eq!(matching[0].ilabel, 1);
assert_eq!(matching[0].nextstate, s1);
let matching: Vec<_> = failure_fst.arcs_matching(s1, 2).collect();
assert_eq!(matching.len(), 1);
assert_eq!(matching[0].ilabel, 2);
assert_eq!(matching[0].nextstate, s2);
}
#[test]
fn test_aho_corasick_no_failure_when_match_exists() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::new(2.0), s0));
let mut failure_fst = FailureFst::new(fst);
failure_fst.set_failure(s1, s0);
let matching: Vec<_> = failure_fst.arcs_matching(s1, 1).collect();
assert_eq!(matching.len(), 1);
assert_eq!(matching[0].ilabel, 1);
assert_eq!(*matching[0].weight.value(), 2.0); }
#[test]
fn test_aho_corasick_no_match_found() {
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::one(), s1));
let mut failure_fst = FailureFst::new(fst);
failure_fst.set_failure(s1, s0);
let matching: Vec<_> = failure_fst.arcs_matching(s1, 999).collect();
assert_eq!(matching.len(), 0);
}
#[test]
fn test_aho_corasick_multiple_matches_in_current_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 10, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 20, TropicalWeight::new(2.0), s2));
let failure_fst = FailureFst::new(fst);
let matching: Vec<_> = failure_fst.arcs_matching(s0, 1).collect();
assert_eq!(matching.len(), 2);
let outputs: Vec<u32> = matching.iter().map(|a| a.olabel).collect();
assert!(outputs.contains(&10));
assert!(outputs.contains(&20));
}
}