#[cfg(feature = "serialization")]
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use rustc_hash::FxHashSet;
use super::state_set::StateSet;
use super::types::{StateId, Transition, TransitionChar, TransitionLabel, TransitionLabelChar};
use super::{NFAChar, NFA};
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serialization", derive(Serialize, Deserialize))]
pub struct OptimizationConfig {
pub eliminate_epsilon: bool,
pub remove_unreachable: bool,
pub remove_dead: bool,
pub deduplicate_transitions: bool,
}
impl Default for OptimizationConfig {
fn default() -> Self {
Self::full()
}
}
impl OptimizationConfig {
pub fn full() -> Self {
Self {
eliminate_epsilon: true,
remove_unreachable: true,
remove_dead: true,
deduplicate_transitions: true,
}
}
pub fn quick() -> Self {
Self {
eliminate_epsilon: false,
remove_unreachable: true,
remove_dead: true,
deduplicate_transitions: false,
}
}
pub fn none() -> Self {
Self {
eliminate_epsilon: false,
remove_unreachable: false,
remove_dead: false,
deduplicate_transitions: false,
}
}
}
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "serialization", derive(Serialize, Deserialize))]
pub struct OptimizationStats {
pub original_states: usize,
pub original_transitions: usize,
pub original_epsilon_count: usize,
pub final_states: usize,
pub final_transitions: usize,
pub states_removed: usize,
pub epsilon_transitions_eliminated: usize,
pub unreachable_states_removed: usize,
pub dead_states_removed: usize,
pub duplicate_transitions_removed: usize,
}
impl OptimizationStats {
pub fn state_reduction_percent(&self) -> f64 {
if self.original_states == 0 {
0.0
} else {
100.0 * (self.original_states - self.final_states) as f64 / self.original_states as f64
}
}
pub fn transition_reduction_percent(&self) -> f64 {
if self.original_transitions == 0 {
0.0
} else {
100.0 * (self.original_transitions - self.final_transitions) as f64
/ self.original_transitions as f64
}
}
}
#[derive(Debug, Clone)]
pub struct NfaOptimizerChar {
config: OptimizationConfig,
}
impl NfaOptimizerChar {
pub fn new(config: OptimizationConfig) -> Self {
Self { config }
}
pub fn optimize(&self, nfa: NFAChar) -> (NFAChar, OptimizationStats) {
let mut nfa = nfa;
nfa.finalize();
let mut stats = OptimizationStats {
original_states: nfa.num_states(),
original_transitions: nfa.num_transitions(),
original_epsilon_count: count_epsilon_transitions_char(&nfa),
..Default::default()
};
let mut result = nfa;
if self.config.eliminate_epsilon {
let before_transitions = result.num_transitions();
result = eliminate_epsilon_char(result);
result.finalize();
let epsilon_after = count_epsilon_transitions_char(&result);
stats.epsilon_transitions_eliminated =
stats.original_epsilon_count.saturating_sub(epsilon_after);
let _ = before_transitions; }
if self.config.remove_unreachable {
let before_states = result.num_states();
result = remove_unreachable_char(result);
result.finalize();
stats.unreachable_states_removed = before_states.saturating_sub(result.num_states());
}
if self.config.remove_dead {
let before_states = result.num_states();
result = remove_dead_char(result);
result.finalize();
stats.dead_states_removed = before_states.saturating_sub(result.num_states());
}
if self.config.deduplicate_transitions {
let before_transitions = result.num_transitions();
result = deduplicate_transitions_char(result);
result.finalize();
stats.duplicate_transitions_removed =
before_transitions.saturating_sub(result.num_transitions());
}
stats.final_states = result.num_states();
stats.final_transitions = result.num_transitions();
stats.states_removed = stats.original_states.saturating_sub(stats.final_states);
(result, stats)
}
}
#[derive(Debug, Clone)]
pub struct NfaOptimizer {
config: OptimizationConfig,
}
impl NfaOptimizer {
pub fn new(config: OptimizationConfig) -> Self {
Self { config }
}
pub fn optimize(&self, nfa: NFA) -> (NFA, OptimizationStats) {
let mut nfa = nfa;
nfa.finalize();
let mut stats = OptimizationStats {
original_states: nfa.num_states(),
original_transitions: nfa.num_transitions(),
original_epsilon_count: count_epsilon_transitions(&nfa),
..Default::default()
};
let mut result = nfa;
if self.config.eliminate_epsilon {
result = eliminate_epsilon(result);
result.finalize();
let epsilon_after = count_epsilon_transitions(&result);
stats.epsilon_transitions_eliminated =
stats.original_epsilon_count.saturating_sub(epsilon_after);
}
if self.config.remove_unreachable {
let before_states = result.num_states();
result = remove_unreachable(result);
result.finalize();
stats.unreachable_states_removed = before_states.saturating_sub(result.num_states());
}
if self.config.remove_dead {
let before_states = result.num_states();
result = remove_dead(result);
result.finalize();
stats.dead_states_removed = before_states.saturating_sub(result.num_states());
}
if self.config.deduplicate_transitions {
let before_transitions = result.num_transitions();
result = deduplicate_transitions(result);
result.finalize();
stats.duplicate_transitions_removed =
before_transitions.saturating_sub(result.num_transitions());
}
stats.final_states = result.num_states();
stats.final_transitions = result.num_transitions();
stats.states_removed = stats.original_states.saturating_sub(stats.final_states);
(result, stats)
}
}
fn count_epsilon_transitions_char(nfa: &NFAChar) -> usize {
nfa.transitions()
.iter()
.filter(|t| t.label.is_epsilon())
.count()
}
fn remove_unreachable_char(nfa: NFAChar) -> NFAChar {
let mut reachable = FxHashSet::default();
let mut queue = VecDeque::new();
reachable.insert(nfa.start());
queue.push_back(nfa.start());
while let Some(state) = queue.pop_front() {
for trans in nfa.transitions_from(state) {
if reachable.insert(trans.to) {
queue.push_back(trans.to);
}
}
}
if reachable.len() == nfa.num_states() {
return nfa;
}
build_nfa_with_states_char(&nfa, &reachable)
}
fn remove_dead_char(nfa: NFAChar) -> NFAChar {
let mut reverse_index: HashMap<StateId, Vec<StateId>> = HashMap::new();
for trans in nfa.transitions() {
reverse_index.entry(trans.to).or_default().push(trans.from);
}
let mut can_reach_final = FxHashSet::default();
let mut queue: VecDeque<StateId> = nfa.finals().iter().copied().collect();
for &final_state in nfa.finals() {
can_reach_final.insert(final_state);
}
while let Some(state) = queue.pop_front() {
if let Some(predecessors) = reverse_index.get(&state) {
for &pred in predecessors {
if can_reach_final.insert(pred) {
queue.push_back(pred);
}
}
}
}
can_reach_final.insert(nfa.start());
if can_reach_final.len() == nfa.num_states() {
return nfa;
}
build_nfa_with_states_char(&nfa, &can_reach_final)
}
fn eliminate_epsilon_char(nfa: NFAChar) -> NFAChar {
let closures: Vec<StateSet> = (0..nfa.num_states() as StateId)
.map(|s| nfa.epsilon_closure_single(s))
.collect();
let mut new_nfa = NFAChar::new();
for _ in 1..nfa.num_states() {
new_nfa.add_state(false);
}
for state_id in 0..nfa.num_states() as StateId {
let is_final = closures[state_id as usize].iter().any(|s| nfa.is_final(s));
new_nfa.set_final(state_id, is_final);
}
for state_id in 0..nfa.num_states() as StateId {
let closure = &closures[state_id as usize];
for reachable in closure.iter() {
for trans in nfa.transitions_from(reachable) {
if trans.label.is_epsilon() {
continue;
}
let dest_closure = &closures[trans.to as usize];
for dest in dest_closure.iter() {
new_nfa.add_transition_weighted(
state_id,
trans.label.clone(),
dest,
trans.weight,
);
}
}
}
}
new_nfa
}
fn deduplicate_transitions_char(nfa: NFAChar) -> NFAChar {
let mut seen: FxHashSet<(StateId, StateId, u64)> = FxHashSet::default();
let mut unique_transitions: Vec<TransitionChar> = Vec::new();
for trans in nfa.transitions() {
let label_hash = hash_label_char(&trans.label);
let key = (trans.from, trans.to, label_hash);
if seen.insert(key) {
unique_transitions.push(trans.clone());
}
}
if unique_transitions.len() == nfa.num_transitions() {
return nfa;
}
let mut new_nfa = NFAChar::new();
for _i in 1..nfa.num_states() {
new_nfa.add_state(false);
}
for &final_state in nfa.finals() {
new_nfa.set_final(final_state, true);
}
for trans in unique_transitions {
new_nfa.add_transition_weighted(trans.from, trans.label, trans.to, trans.weight);
}
new_nfa
}
fn build_nfa_with_states_char(nfa: &NFAChar, keep_states: &FxHashSet<StateId>) -> NFAChar {
let mut old_to_new: HashMap<StateId, StateId> = HashMap::new();
let mut sorted_states: Vec<StateId> = keep_states.iter().copied().collect();
sorted_states.sort_unstable();
for (new_id, &old_id) in sorted_states.iter().enumerate() {
old_to_new.insert(old_id, new_id as StateId);
}
let mut new_nfa = NFAChar::new();
for _i in 1..sorted_states.len() {
new_nfa.add_state(false);
}
let _new_start = *old_to_new.get(&nfa.start()).unwrap_or(&0);
for &old_id in nfa.finals() {
if let Some(&new_id) = old_to_new.get(&old_id) {
new_nfa.set_final(new_id, true);
}
}
for trans in nfa.transitions() {
if let (Some(&new_from), Some(&new_to)) =
(old_to_new.get(&trans.from), old_to_new.get(&trans.to))
{
new_nfa.add_transition_weighted(new_from, trans.label.clone(), new_to, trans.weight);
}
}
new_nfa
}
fn hash_label_char(label: &TransitionLabelChar) -> u64 {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
match label {
TransitionLabelChar::Epsilon => 0u8.hash(&mut hasher),
TransitionLabelChar::Char(c) => {
1u8.hash(&mut hasher);
c.hash(&mut hasher);
}
TransitionLabelChar::CharClass(class) => {
2u8.hash(&mut hasher);
format!("{:?}", class).hash(&mut hasher);
}
TransitionLabelChar::Any => 3u8.hash(&mut hasher),
TransitionLabelChar::StartOfLine => 4u8.hash(&mut hasher),
TransitionLabelChar::EndOfLine => 5u8.hash(&mut hasher),
TransitionLabelChar::StartOfInput => 6u8.hash(&mut hasher),
TransitionLabelChar::EndOfInput => 7u8.hash(&mut hasher),
TransitionLabelChar::EndOfInputStrict => 8u8.hash(&mut hasher),
}
hasher.finish()
}
fn count_epsilon_transitions(nfa: &NFA) -> usize {
nfa.transitions()
.iter()
.filter(|t| t.label.is_epsilon())
.count()
}
fn remove_unreachable(nfa: NFA) -> NFA {
let mut reachable = FxHashSet::default();
let mut queue = VecDeque::new();
reachable.insert(nfa.start());
queue.push_back(nfa.start());
while let Some(state) = queue.pop_front() {
for trans in nfa.transitions_from(state) {
if reachable.insert(trans.to) {
queue.push_back(trans.to);
}
}
}
if reachable.len() == nfa.num_states() {
return nfa;
}
build_nfa_with_states(&nfa, &reachable)
}
fn remove_dead(nfa: NFA) -> NFA {
let mut reverse_index: HashMap<StateId, Vec<StateId>> = HashMap::new();
for trans in nfa.transitions() {
reverse_index.entry(trans.to).or_default().push(trans.from);
}
let mut can_reach_final = FxHashSet::default();
let mut queue: VecDeque<StateId> = nfa.finals().iter().copied().collect();
for &final_state in nfa.finals() {
can_reach_final.insert(final_state);
}
while let Some(state) = queue.pop_front() {
if let Some(predecessors) = reverse_index.get(&state) {
for &pred in predecessors {
if can_reach_final.insert(pred) {
queue.push_back(pred);
}
}
}
}
can_reach_final.insert(nfa.start());
if can_reach_final.len() == nfa.num_states() {
return nfa;
}
build_nfa_with_states(&nfa, &can_reach_final)
}
fn eliminate_epsilon(nfa: NFA) -> NFA {
let closures: Vec<StateSet> = (0..nfa.num_states() as StateId)
.map(|s| nfa.epsilon_closure_single(s))
.collect();
let mut new_nfa = NFA::new();
for _ in 1..nfa.num_states() {
new_nfa.add_state(false);
}
for state_id in 0..nfa.num_states() as StateId {
let is_final = closures[state_id as usize].iter().any(|s| nfa.is_final(s));
new_nfa.set_final(state_id, is_final);
}
for state_id in 0..nfa.num_states() as StateId {
let closure = &closures[state_id as usize];
for reachable in closure.iter() {
for trans in nfa.transitions_from(reachable) {
if trans.label.is_epsilon() {
continue;
}
let dest_closure = &closures[trans.to as usize];
for dest in dest_closure.iter() {
new_nfa.add_transition_weighted(
state_id,
trans.label.clone(),
dest,
trans.weight,
);
}
}
}
}
new_nfa
}
fn deduplicate_transitions(nfa: NFA) -> NFA {
let mut seen: FxHashSet<(StateId, StateId, u64)> = FxHashSet::default();
let mut unique_transitions: Vec<Transition> = Vec::new();
for trans in nfa.transitions() {
let label_hash = hash_label(&trans.label);
let key = (trans.from, trans.to, label_hash);
if seen.insert(key) {
unique_transitions.push(trans.clone());
}
}
if unique_transitions.len() == nfa.num_transitions() {
return nfa;
}
let mut new_nfa = NFA::new();
for _ in 1..nfa.num_states() {
new_nfa.add_state(false);
}
for &final_state in nfa.finals() {
new_nfa.set_final(final_state, true);
}
for trans in unique_transitions {
new_nfa.add_transition_weighted(trans.from, trans.label, trans.to, trans.weight);
}
new_nfa
}
fn build_nfa_with_states(nfa: &NFA, keep_states: &FxHashSet<StateId>) -> NFA {
let mut old_to_new: HashMap<StateId, StateId> = HashMap::new();
let mut sorted_states: Vec<StateId> = keep_states.iter().copied().collect();
sorted_states.sort_unstable();
for (new_id, &old_id) in sorted_states.iter().enumerate() {
old_to_new.insert(old_id, new_id as StateId);
}
let mut new_nfa = NFA::new();
for _ in 1..sorted_states.len() {
new_nfa.add_state(false);
}
for &old_id in nfa.finals() {
if let Some(&new_id) = old_to_new.get(&old_id) {
new_nfa.set_final(new_id, true);
}
}
for trans in nfa.transitions() {
if let (Some(&new_from), Some(&new_to)) =
(old_to_new.get(&trans.from), old_to_new.get(&trans.to))
{
new_nfa.add_transition_weighted(new_from, trans.label.clone(), new_to, trans.weight);
}
}
new_nfa
}
fn hash_label(label: &TransitionLabel) -> u64 {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
match label {
TransitionLabel::Epsilon => 0u8.hash(&mut hasher),
TransitionLabel::Byte(b) => {
1u8.hash(&mut hasher);
b.hash(&mut hasher);
}
TransitionLabel::CharClass(class) => {
2u8.hash(&mut hasher);
format!("{:?}", class).hash(&mut hasher);
}
TransitionLabel::Any => 3u8.hash(&mut hasher),
TransitionLabel::StartOfLine => 4u8.hash(&mut hasher),
TransitionLabel::EndOfLine => 5u8.hash(&mut hasher),
TransitionLabel::StartOfInput => 6u8.hash(&mut hasher),
TransitionLabel::EndOfInput => 7u8.hash(&mut hasher),
TransitionLabel::EndOfInputStrict => 8u8.hash(&mut hasher),
}
hasher.finish()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::phonetic::nfa::{compile, ThompsonBuilderChar};
use crate::phonetic::regex::parse;
#[test]
fn test_optimization_config_full() {
let config = OptimizationConfig::full();
assert!(config.eliminate_epsilon);
assert!(config.remove_unreachable);
assert!(config.remove_dead);
assert!(config.deduplicate_transitions);
}
#[test]
fn test_optimization_config_quick() {
let config = OptimizationConfig::quick();
assert!(!config.eliminate_epsilon);
assert!(config.remove_unreachable);
assert!(config.remove_dead);
assert!(!config.deduplicate_transitions);
}
#[test]
fn test_optimization_config_none() {
let config = OptimizationConfig::none();
assert!(!config.eliminate_epsilon);
assert!(!config.remove_unreachable);
assert!(!config.remove_dead);
assert!(!config.deduplicate_transitions);
}
#[test]
fn test_count_epsilon_transitions() {
let builder = ThompsonBuilderChar::new();
let a = builder.single_char('a');
let b = builder.single_char('b');
let nfa = builder.alternation(a, b);
let count = count_epsilon_transitions_char(&nfa);
assert!(count > 0, "alternation should have epsilon transitions");
}
#[test]
fn test_remove_unreachable_no_change() {
let builder = ThompsonBuilderChar::new();
let nfa = builder.single_char('a');
let optimized = remove_unreachable_char(nfa.clone());
assert_eq!(nfa.num_states(), optimized.num_states());
assert_eq!(nfa.num_transitions(), optimized.num_transitions());
}
#[test]
fn test_epsilon_elimination_simple() {
let builder = ThompsonBuilderChar::new();
let a = builder.single_char('a');
let nfa = builder.kleene_star(a);
let before_epsilon = count_epsilon_transitions_char(&nfa);
assert!(
before_epsilon > 0,
"kleene_star should have epsilon transitions"
);
let optimized = eliminate_epsilon_char(nfa.clone());
let after_epsilon = count_epsilon_transitions_char(&optimized);
assert_eq!(
after_epsilon, 0,
"epsilon elimination should remove all epsilon transitions"
);
assert!(nfa.accepts(""));
assert!(optimized.accepts(""));
assert!(nfa.accepts("a"));
assert!(optimized.accepts("a"));
assert!(nfa.accepts("aaa"));
assert!(optimized.accepts("aaa"));
assert!(!nfa.accepts("b"));
assert!(!optimized.accepts("b"));
}
#[test]
fn test_epsilon_elimination_alternation() {
let builder = ThompsonBuilderChar::new();
let a = builder.single_char('a');
let b = builder.single_char('b');
let nfa = builder.alternation(a, b);
let optimized = eliminate_epsilon_char(nfa.clone());
assert!(nfa.accepts("a"));
assert!(optimized.accepts("a"));
assert!(nfa.accepts("b"));
assert!(optimized.accepts("b"));
assert!(!nfa.accepts("c"));
assert!(!optimized.accepts("c"));
assert!(!nfa.accepts("ab"));
assert!(!optimized.accepts("ab"));
}
#[test]
fn test_full_optimization_preserves_language() {
use crate::phonetic::nfa::compiler::NFACompilerChar;
let patterns = ["a", "ab", "a|b", "a*", "a+", "a?", "(ab)+", "(a|b)*c"];
for pattern in patterns {
let regex = parse(pattern).expect("parse");
let mut compiler = NFACompilerChar::new().without_optimization();
let nfa = compiler.compile(®ex).expect("compile");
let optimizer = NfaOptimizerChar::new(OptimizationConfig::full());
let (optimized, stats) = optimizer.optimize(nfa.clone());
let test_inputs = ["", "a", "b", "c", "ab", "abc", "aaa", "bbb", "abababc"];
for input in test_inputs {
assert_eq!(
nfa.accepts(input),
optimized.accepts(input),
"language mismatch for pattern '{}' on input '{}'",
pattern,
input
);
}
if pattern.contains('|') || pattern.contains('*') || pattern.contains('+') {
assert!(
stats.original_epsilon_count > 0,
"pattern '{}' should have epsilon transitions",
pattern
);
}
}
}
#[test]
fn test_optimization_stats() {
let builder = ThompsonBuilderChar::new();
let a = builder.single_char('a');
let nfa = builder.kleene_star(a);
let optimizer = NfaOptimizerChar::new(OptimizationConfig::full());
let (_optimized, stats) = optimizer.optimize(nfa);
assert!(stats.original_states > 0);
assert!(stats.original_transitions > 0);
assert!(stats.original_epsilon_count > 0);
assert!(stats.epsilon_transitions_eliminated > 0);
assert!(stats.final_states > 0);
}
#[test]
fn test_deduplicate_transitions() {
let mut nfa = NFAChar::new();
let q1 = nfa.add_state(true);
nfa.add_transition_char(0, 'a', q1);
nfa.add_transition_char(0, 'a', q1);
nfa.add_transition_char(0, 'a', q1);
nfa.finalize();
assert_eq!(nfa.num_transitions(), 3);
let mut optimized = deduplicate_transitions_char(nfa);
optimized.finalize();
assert_eq!(optimized.num_transitions(), 1);
}
#[test]
fn test_optimization_with_anchors() {
let regex = parse("^hello$").expect("parse");
let nfa = compile(®ex).expect("compile");
let optimizer = NfaOptimizerChar::new(OptimizationConfig::full());
let (optimized, _stats) = optimizer.optimize(nfa.clone());
assert!(optimized.num_states() > 0);
assert!(optimized.num_transitions() > 0);
}
}