use crate::arc::Arc;
use crate::fst::{Fst, Label, MutableFst, StateId, VectorFst};
use crate::semiring::Semiring;
use crate::Result;
use std::collections::{HashMap, HashSet, VecDeque};
#[derive(Debug, Clone)]
pub struct ReplaceConfig {
pub max_depth: usize,
pub left_to_right: bool,
pub remove_epsilon: bool,
pub enable_cycle_detection: bool,
pub return_arc_type: ReturnArcType,
}
impl Default for ReplaceConfig {
fn default() -> Self {
Self {
max_depth: 100,
left_to_right: true,
remove_epsilon: false,
enable_cycle_detection: true,
return_arc_type: ReturnArcType::Epsilon,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReturnArcType {
Epsilon,
Explicit,
}
#[derive(Debug, Clone)]
pub struct ReplaceFst<W: Semiring> {
pub root: Label,
pub rules: HashMap<Label, VectorFst<W>>,
pub config: ReplaceConfig,
}
impl<W: Semiring> ReplaceFst<W> {
pub fn new(root: Label, rules: HashMap<Label, VectorFst<W>>) -> Self {
Self {
root,
rules,
config: ReplaceConfig::default(),
}
}
pub fn with_config(
root: Label,
rules: HashMap<Label, VectorFst<W>>,
config: ReplaceConfig,
) -> Self {
Self {
root,
rules,
config,
}
}
pub fn simple(root: Label) -> Self {
Self {
root,
rules: HashMap::new(),
config: ReplaceConfig::default(),
}
}
}
#[derive(Debug, Clone)]
struct ReplaceContext {
call_stack: Vec<Label>,
depth: usize,
state_mappings: Vec<HashMap<StateId, StateId>>,
}
impl ReplaceContext {
fn new() -> Self {
Self {
call_stack: Vec::new(),
depth: 0,
state_mappings: Vec::new(),
}
}
fn push_call(&mut self, label: Label) -> bool {
if self.call_stack.contains(&label) {
false } else {
self.call_stack.push(label);
self.depth += 1;
self.state_mappings.push(HashMap::new());
true
}
}
fn pop_call(&mut self) {
if !self.call_stack.is_empty() {
self.call_stack.pop();
self.depth -= 1;
self.state_mappings.pop();
}
}
}
pub fn replace<W, M>(replace_fst: &ReplaceFst<W>) -> Result<M>
where
W: Semiring + Clone,
M: MutableFst<W> + Default,
{
let mut result = M::default();
let mut context = ReplaceContext::new();
if replace_fst.rules.is_empty() {
let s0 = result.add_state();
result.set_start(s0);
result.set_final(s0, W::one());
return Ok(result);
}
if let Some(root_fst) = replace_fst.rules.get(&replace_fst.root) {
expand_grammar(
root_fst,
&mut result,
&replace_fst.rules,
&mut context,
&replace_fst.config,
)?;
} else {
let s0 = result.add_state();
result.set_start(s0);
result.set_final(s0, W::one());
}
if replace_fst.config.remove_epsilon {
result = remove_epsilon_transitions(&result)?;
}
Ok(result)
}
fn expand_grammar<W, M>(
root_fst: &VectorFst<W>,
result: &mut M,
rules: &HashMap<Label, VectorFst<W>>,
context: &mut ReplaceContext,
config: &ReplaceConfig,
) -> Result<HashMap<StateId, StateId>>
where
W: Semiring + Clone,
M: MutableFst<W>,
{
if context.depth >= config.max_depth {
return Err(crate::Error::Algorithm(format!(
"Maximum replacement depth {} exceeded",
config.max_depth
)));
}
let mut state_map = HashMap::new();
for state in root_fst.states() {
let new_state = result.add_state();
state_map.insert(state, new_state);
if let Some(weight) = root_fst.final_weight(state) {
result.set_final(new_state, weight.clone());
}
}
if context.depth == 0 {
if let Some(start) = root_fst.start() {
if let Some(&new_start) = state_map.get(&start) {
result.set_start(new_start);
}
}
}
if context.state_mappings.len() <= context.depth {
context.state_mappings.push(state_map.clone());
} else {
context.state_mappings[context.depth] = state_map.clone();
}
for state in root_fst.states() {
let mapped_state = state_map[&state];
for arc in root_fst.arcs(state) {
let mapped_nextstate = state_map[&arc.nextstate];
if let Some(replacement_fst) = rules.get(&arc.ilabel) {
expand_non_terminal(
replacement_fst,
result,
rules,
mapped_state,
mapped_nextstate,
&arc,
context,
config,
)?;
} else {
result.add_arc(
mapped_state,
Arc::new(arc.ilabel, arc.olabel, arc.weight.clone(), mapped_nextstate),
);
}
}
}
Ok(state_map)
}
fn expand_non_terminal<W, M>(
replacement_fst: &VectorFst<W>,
result: &mut M,
rules: &HashMap<Label, VectorFst<W>>,
from_state: StateId,
to_state: StateId,
original_arc: &Arc<W>,
context: &mut ReplaceContext,
config: &ReplaceConfig,
) -> Result<()>
where
W: Semiring + Clone,
M: MutableFst<W>,
{
if config.enable_cycle_detection && !context.push_call(original_arc.ilabel) {
result.add_arc(
from_state,
Arc::new(
0,
original_arc.olabel,
original_arc.weight.clone(),
to_state,
),
);
return Ok(());
}
let replacement_states = expand_grammar(replacement_fst, result, rules, context, config)?;
connect_replacement(
replacement_fst,
result,
&replacement_states,
from_state,
to_state,
original_arc,
config,
)?;
if config.enable_cycle_detection {
context.pop_call();
}
Ok(())
}
fn connect_replacement<W, M>(
replacement_fst: &VectorFst<W>,
result: &mut M,
replacement_states: &HashMap<StateId, StateId>,
from_state: StateId,
to_state: StateId,
original_arc: &Arc<W>,
config: &ReplaceConfig,
) -> Result<()>
where
W: Semiring + Clone,
M: MutableFst<W>,
{
if let Some(replacement_start) = replacement_fst.start() {
if let Some(&mapped_start) = replacement_states.get(&replacement_start) {
let entry_ilabel = match config.return_arc_type {
ReturnArcType::Epsilon => 0,
ReturnArcType::Explicit => original_arc.ilabel,
};
result.add_arc(
from_state,
Arc::new(
entry_ilabel,
original_arc.olabel,
original_arc.weight.clone(),
mapped_start,
),
);
}
}
for replacement_state in replacement_fst.states() {
if let Some(final_weight) = replacement_fst.final_weight(replacement_state) {
if let Some(&mapped_state) = replacement_states.get(&replacement_state) {
let exit_weight = match config.return_arc_type {
ReturnArcType::Epsilon => final_weight.clone(),
ReturnArcType::Explicit => final_weight.clone(),
};
result.add_arc(mapped_state, Arc::new(0, 0, exit_weight, to_state));
}
}
}
Ok(())
}
fn remove_epsilon_transitions<W, M>(fst: &M) -> Result<M>
where
W: Semiring + Clone,
M: MutableFst<W> + Default,
{
let mut result = M::default();
for _ in 0..fst.num_states() {
result.add_state();
}
if let Some(start) = fst.start() {
result.set_start(start);
}
for state in fst.states() {
if let Some(weight) = fst.final_weight(state) {
result.set_final(state, weight.clone());
}
for arc in fst.arcs(state) {
if arc.ilabel != 0 || arc.olabel != 0 {
result.add_arc(state, arc.clone());
} else {
}
}
}
Ok(result)
}
#[allow(dead_code)]
fn compute_epsilon_closure<W, F>(fst: &F, states: &HashSet<StateId>) -> Result<HashMap<StateId, W>>
where
W: Semiring + Clone,
F: Fst<W>,
{
let mut closure = HashMap::new();
let mut queue = VecDeque::new();
for &state in states {
closure.insert(state, W::one());
queue.push_back((state, W::one()));
}
while let Some((current_state, current_weight)) = queue.pop_front() {
for arc in fst.arcs(current_state) {
if arc.ilabel == 0 && arc.olabel == 0 {
let new_weight = current_weight.times(&arc.weight);
let should_update = match closure.get(&arc.nextstate) {
None => true,
Some(existing_weight) => {
let combined = existing_weight.plus(&new_weight);
if combined != *existing_weight {
closure.insert(arc.nextstate, combined);
false } else {
false
}
}
};
if should_update {
closure.insert(arc.nextstate, new_weight.clone());
queue.push_back((arc.nextstate, new_weight));
}
}
}
}
Ok(closure)
}
#[allow(dead_code)]
pub fn validate_grammar<W>(rules: &HashMap<Label, VectorFst<W>>) -> Result<()>
where
W: Semiring,
{
if rules.is_empty() {
return Ok(()); }
for (label, fst) in rules {
if fst.num_states() == 0 {
return Err(crate::Error::Algorithm(format!(
"Rule for label {label} has no states"
)));
}
if fst.start().is_none() {
return Err(crate::Error::Algorithm(format!(
"Rule for label {label} has no start state"
)));
}
let mut reachable = HashSet::new();
if let Some(start) = fst.start() {
compute_reachable_states(fst, start, &mut reachable);
}
if reachable.len() != fst.num_states() {
return Err(crate::Error::Algorithm(format!(
"Rule for label {label} has unreachable states"
)));
}
}
Ok(())
}
#[allow(dead_code)]
fn compute_reachable_states<W, F>(fst: &F, start: StateId, reachable: &mut HashSet<StateId>)
where
W: Semiring,
F: Fst<W>,
{
let mut stack = vec![start];
while let Some(state) = stack.pop() {
if reachable.insert(state) {
for arc in fst.arcs(state) {
stack.push(arc.nextstate);
}
}
}
}
#[allow(dead_code)]
pub fn from_string_rules<W>(
root: Label,
string_rules: HashMap<Label, String>,
) -> Result<ReplaceFst<W>>
where
W: Semiring + Clone,
{
let mut rules = HashMap::new();
for (label, string) in string_rules {
let fst = create_string_fst(&string)?;
rules.insert(label, fst);
}
Ok(ReplaceFst::new(root, rules))
}
#[allow(dead_code)]
fn create_string_fst<W>(string: &str) -> Result<VectorFst<W>>
where
W: Semiring + Clone,
{
let mut fst = VectorFst::new();
if string.is_empty() {
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, W::one());
return Ok(fst);
}
let chars: Vec<char> = string.chars().collect();
let mut states = Vec::new();
for _ in 0..=chars.len() {
states.push(fst.add_state());
}
fst.set_start(states[0]);
fst.set_final(states[chars.len()], W::one());
for (i, &ch) in chars.iter().enumerate() {
let label = ch as u32;
fst.add_arc(states[i], Arc::new(label, label, W::one(), states[i + 1]));
}
Ok(fst)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_replace_fst_simple() {
let replace_fst = ReplaceFst::<TropicalWeight>::simple(42);
assert_eq!(replace_fst.root, 42);
assert!(replace_fst.rules.is_empty());
}
#[test]
fn test_replace_config() {
let config = ReplaceConfig {
max_depth: 50,
left_to_right: false,
remove_epsilon: true,
enable_cycle_detection: false,
return_arc_type: ReturnArcType::Explicit,
};
assert_eq!(config.max_depth, 50);
assert!(!config.left_to_right);
assert!(config.remove_epsilon);
assert!(!config.enable_cycle_detection);
assert_eq!(config.return_arc_type, ReturnArcType::Explicit);
}
#[test]
fn test_replace_basic() {
let replace_fst = ReplaceFst::<TropicalWeight>::simple(1);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
assert_eq!(result.num_states(), 1);
let start_state = result.start().unwrap();
assert!(result.is_final(start_state));
assert_eq!(
result.final_weight(start_state),
Some(&TropicalWeight::one())
);
}
#[test]
fn test_replace_with_rules() {
let mut rules = HashMap::new();
let mut rule_fst = VectorFst::<TropicalWeight>::new();
let s0 = rule_fst.add_state();
let s1 = rule_fst.add_state();
rule_fst.set_start(s0);
rule_fst.set_final(s1, TropicalWeight::one());
rule_fst.add_arc(
s0,
Arc::new('a' as u32, 'a' as u32, TropicalWeight::one(), s1),
);
rules.insert(1, rule_fst);
let replace_fst = ReplaceFst::new(1, rules);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
assert!(result.num_states() >= 2);
}
#[test]
fn test_replace_with_config() {
let mut rules = HashMap::new();
let mut rule_fst = VectorFst::<TropicalWeight>::new();
let s0 = rule_fst.add_state();
let s1 = rule_fst.add_state();
rule_fst.set_start(s0);
rule_fst.set_final(s1, TropicalWeight::one());
rule_fst.add_arc(
s0,
Arc::new('b' as u32, 'b' as u32, TropicalWeight::one(), s1),
);
rules.insert(2, rule_fst);
let config = ReplaceConfig {
max_depth: 10,
remove_epsilon: true,
..Default::default()
};
let replace_fst = ReplaceFst::with_config(2, rules, config);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
}
#[test]
fn test_replace_context() {
let mut context = ReplaceContext::new();
assert_eq!(context.depth, 0);
assert!(context.call_stack.is_empty());
assert!(context.push_call(1));
assert_eq!(context.depth, 1);
assert_eq!(context.call_stack.len(), 1);
assert!(!context.push_call(1));
context.pop_call();
assert_eq!(context.depth, 0);
}
#[test]
fn test_validate_grammar() {
let mut rules = HashMap::new();
let mut rule_fst = VectorFst::<TropicalWeight>::new();
let s0 = rule_fst.add_state();
let s1 = rule_fst.add_state();
rule_fst.set_start(s0);
rule_fst.set_final(s1, TropicalWeight::one());
rule_fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
rules.insert(1, rule_fst);
assert!(validate_grammar(&rules).is_ok());
let mut invalid_fst = VectorFst::<TropicalWeight>::new();
invalid_fst.add_state();
rules.insert(2, invalid_fst);
assert!(validate_grammar(&rules).is_err());
}
#[test]
fn test_create_string_fst() {
let fst: VectorFst<TropicalWeight> = create_string_fst("abc").unwrap();
assert!(fst.start().is_some());
assert_eq!(fst.num_states(), 4);
let empty_fst: VectorFst<TropicalWeight> = create_string_fst("").unwrap();
assert_eq!(empty_fst.num_states(), 1);
assert!(empty_fst.is_final(empty_fst.start().unwrap()));
}
#[test]
fn test_from_string_rules() {
let mut string_rules = HashMap::new();
string_rules.insert(1, "hello".to_string());
string_rules.insert(2, "world".to_string());
let replace_fst: ReplaceFst<TropicalWeight> = from_string_rules(1, string_rules).unwrap();
assert_eq!(replace_fst.root, 1);
assert_eq!(replace_fst.rules.len(), 2);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
}
#[test]
fn test_return_arc_types() {
assert_eq!(ReturnArcType::Epsilon, ReturnArcType::Epsilon);
assert_ne!(ReturnArcType::Epsilon, ReturnArcType::Explicit);
}
#[test]
fn test_epsilon_closure() {
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(0, 0, TropicalWeight::one(), s1));
fst.add_arc(s1, Arc::new(0, 0, TropicalWeight::new(0.5), s2));
let mut states = HashSet::new();
states.insert(s0);
let closure = compute_epsilon_closure(&fst, &states).unwrap();
assert!(closure.contains_key(&s0));
assert!(closure.contains_key(&s1));
assert!(closure.contains_key(&s2));
}
#[test]
fn test_complex_replacement() {
let mut rules = HashMap::new();
let mut rule_a = VectorFst::<TropicalWeight>::new();
let s0 = rule_a.add_state();
let s1 = rule_a.add_state();
let s2 = rule_a.add_state();
rule_a.set_start(s0);
rule_a.set_final(s2, TropicalWeight::one());
rule_a.add_arc(
s0,
Arc::new('a' as u32, 'a' as u32, TropicalWeight::one(), s1),
);
rule_a.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
let mut rule_b = VectorFst::<TropicalWeight>::new();
let s0 = rule_b.add_state();
let s1 = rule_b.add_state();
rule_b.set_start(s0);
rule_b.set_final(s1, TropicalWeight::one());
rule_b.add_arc(
s0,
Arc::new('b' as u32, 'b' as u32, TropicalWeight::one(), s1),
);
rules.insert(1, rule_a); rules.insert(2, rule_b);
let replace_fst = ReplaceFst::new(1, rules);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
assert!(result.num_states() > 2); }
#[test]
fn test_cycle_detection() {
let mut rules = HashMap::new();
let mut rule_a = VectorFst::<TropicalWeight>::new();
let s0 = rule_a.add_state();
let s1 = rule_a.add_state();
let s2 = rule_a.add_state();
rule_a.set_start(s0);
rule_a.set_final(s2, TropicalWeight::one());
rule_a.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1)); rule_a.add_arc(
s1,
Arc::new('a' as u32, 'a' as u32, TropicalWeight::one(), s2),
);
rules.insert(1, rule_a);
let config = ReplaceConfig {
enable_cycle_detection: true,
..Default::default()
};
let replace_fst = ReplaceFst::with_config(1, rules, config);
let result: VectorFst<TropicalWeight> = replace(&replace_fst).unwrap();
assert!(result.start().is_some());
}
#[test]
fn test_max_depth_limit() {
let mut rules = HashMap::new();
for i in 1..=10 {
let mut rule = VectorFst::<TropicalWeight>::new();
let s0 = rule.add_state();
let s1 = rule.add_state();
rule.set_start(s0);
rule.set_final(s1, TropicalWeight::one());
if i < 10 {
rule.add_arc(s0, Arc::new(i + 1, i + 1, TropicalWeight::one(), s1));
} else {
rule.add_arc(
s0,
Arc::new('x' as u32, 'x' as u32, TropicalWeight::one(), s1),
);
}
rules.insert(i, rule);
}
let config = ReplaceConfig {
max_depth: 5, ..Default::default()
};
let replace_fst = ReplaceFst::with_config(1, rules, config);
let result = replace::<TropicalWeight, VectorFst<TropicalWeight>>(&replace_fst);
assert!(result.is_err()); }
}