use std::collections::HashSet;
use crate::phonetic::regex::ast::Regex;
use super::ast::SymbolTable;
use super::error::{LLreError, LLreErrorKind, LLreResult};
const MAX_EXPANSION_DEPTH: usize = 100;
pub fn expand_pattern_symbols(regex: &Regex, symbol_table: &SymbolTable) -> LLreResult<Regex> {
let mut visited = HashSet::new();
expand_recursive(regex, symbol_table, &mut visited, 0)
}
fn expand_recursive(
regex: &Regex,
symbol_table: &SymbolTable,
visited: &mut HashSet<String>,
depth: usize,
) -> LLreResult<Regex> {
if depth > MAX_EXPANSION_DEPTH {
return Err(LLreError::new(LLreErrorKind::RecursionDepthExceeded {
depth,
max: MAX_EXPANSION_DEPTH,
}));
}
match regex {
Regex::GroupRef(name) => {
if let Some(pattern) = symbol_table.patterns.get(name) {
if visited.contains(name) {
return Err(LLreError::new(LLreErrorKind::CyclicPatternReference {
name: name.clone(),
chain: visited.iter().cloned().collect(),
}));
}
visited.insert(name.clone());
let expanded = expand_recursive(pattern, symbol_table, visited, depth + 1)?;
visited.remove(name);
Ok(Regex::NonCapturingGroup(Box::new(expanded)))
} else {
Ok(regex.clone())
}
}
Regex::Concat(a, b) => {
let a_exp = expand_recursive(a, symbol_table, visited, depth)?;
let b_exp = expand_recursive(b, symbol_table, visited, depth)?;
Ok(Regex::Concat(Box::new(a_exp), Box::new(b_exp)))
}
Regex::Alt(a, b) => {
let a_exp = expand_recursive(a, symbol_table, visited, depth)?;
let b_exp = expand_recursive(b, symbol_table, visited, depth)?;
Ok(Regex::Alt(Box::new(a_exp), Box::new(b_exp)))
}
Regex::Star(inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::Star(Box::new(inner_exp)))
}
Regex::Plus(inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::Plus(Box::new(inner_exp)))
}
Regex::Optional(inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::Optional(Box::new(inner_exp)))
}
Regex::RepeatExact(inner, n) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::RepeatExact(Box::new(inner_exp), *n))
}
Regex::RepeatRange(inner, min, max) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::RepeatRange(Box::new(inner_exp), *min, *max))
}
Regex::CapturingGroup(num, inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::CapturingGroup(*num, Box::new(inner_exp)))
}
Regex::NonCapturingGroup(inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::NonCapturingGroup(Box::new(inner_exp)))
}
Regex::NamedGroup(name, inner) => {
let inner_exp = expand_recursive(inner, symbol_table, visited, depth)?;
Ok(Regex::NamedGroup(name.clone(), Box::new(inner_exp)))
}
Regex::FlagsGroup { flags, inner } => {
let inner_exp = inner
.as_ref()
.map(|i| expand_recursive(i, symbol_table, visited, depth))
.transpose()?;
Ok(Regex::FlagsGroup {
flags: flags.clone(),
inner: inner_exp.map(Box::new),
})
}
Regex::RewriteRule {
pattern,
replacement,
context,
weight,
} => {
let pattern_exp = expand_recursive(pattern, symbol_table, visited, depth)?;
let replacement_exp = expand_recursive(replacement, symbol_table, visited, depth)?;
let context_exp = context
.as_ref()
.map(|ctx| expand_context_predicate(ctx, symbol_table, visited, depth))
.transpose()?;
Ok(Regex::RewriteRule {
pattern: Box::new(pattern_exp),
replacement: Box::new(replacement_exp),
context: context_exp.map(Box::new),
weight: *weight,
})
}
Regex::Empty
| Regex::Char(_)
| Regex::CharClass(_)
| Regex::Any
| Regex::WordBoundary
| Regex::StartOfLine
| Regex::EndOfLine
| Regex::StartOfInput
| Regex::EndOfInput
| Regex::EndOfInputStrict => Ok(regex.clone()),
}
}
fn expand_context_predicate(
ctx: &crate::phonetic::regex::ast::ContextPredicate,
symbol_table: &SymbolTable,
visited: &mut HashSet<String>,
depth: usize,
) -> LLreResult<crate::phonetic::regex::ast::ContextPredicate> {
use crate::phonetic::regex::ast::ContextPredicate;
Ok(ContextPredicate {
left: ctx
.left
.as_ref()
.map(|e| expand_context_expr(e, symbol_table, visited, depth))
.transpose()?,
right: ctx
.right
.as_ref()
.map(|e| expand_context_expr(e, symbol_table, visited, depth))
.transpose()?,
syllable: ctx.syllable.clone(),
})
}
fn expand_context_expr(
expr: &crate::phonetic::regex::ast::ContextExpr,
symbol_table: &SymbolTable,
visited: &mut HashSet<String>,
depth: usize,
) -> LLreResult<crate::phonetic::regex::ast::ContextExpr> {
use crate::phonetic::regex::ast::ContextExpr;
match expr {
ContextExpr::Pattern(regex) => {
let expanded = expand_recursive(regex, symbol_table, visited, depth)?;
Ok(ContextExpr::Pattern(expanded))
}
ContextExpr::WordBoundary => Ok(ContextExpr::WordBoundary),
ContextExpr::And(a, b) => {
let a_exp = expand_context_expr(a, symbol_table, visited, depth)?;
let b_exp = expand_context_expr(b, symbol_table, visited, depth)?;
Ok(ContextExpr::And(Box::new(a_exp), Box::new(b_exp)))
}
ContextExpr::Or(a, b) => {
let a_exp = expand_context_expr(a, symbol_table, visited, depth)?;
let b_exp = expand_context_expr(b, symbol_table, visited, depth)?;
Ok(ContextExpr::Or(Box::new(a_exp), Box::new(b_exp)))
}
ContextExpr::Not(inner) => {
let inner_exp = expand_context_expr(inner, symbol_table, visited, depth)?;
Ok(ContextExpr::Not(Box::new(inner_exp)))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::phonetic::regex::parse;
#[test]
fn test_no_expansion_needed() {
let table = SymbolTable::new();
let regex = parse("[a-z]+").expect("test fixture: parse must be Ok");
let expanded =
expand_pattern_symbols(®ex, &table).expect("test fixture: expansion must be Ok");
assert_eq!(format!("{}", expanded), format!("{}", regex));
}
#[test]
fn test_simple_pattern_expansion() {
let mut table = SymbolTable::new();
let digit_pattern = parse("[0-9]").expect("test fixture: parse must be Ok");
table.add_pattern("DIGIT", digit_pattern, None);
let regex = Regex::GroupRef("DIGIT".to_string());
let expanded =
expand_pattern_symbols(®ex, &table).expect("test fixture: expansion must be Ok");
if let Regex::NonCapturingGroup(inner) = &expanded {
if let Regex::CharClass(_) = inner.as_ref() {
} else {
panic!(
"Expected CharClass inside NonCapturingGroup, got {:?}",
inner
);
}
} else {
panic!("Expected NonCapturingGroup, got {:?}", expanded);
}
}
#[test]
fn test_nested_pattern_expansion() {
let mut table = SymbolTable::new();
let digit_pattern = parse("[0-9]").expect("test fixture: parse must be Ok");
table.add_pattern("DIGIT", digit_pattern, None);
let number_pattern = Regex::Plus(Box::new(Regex::GroupRef("DIGIT".to_string())));
table.add_pattern("NUMBER", number_pattern, None);
let regex = Regex::GroupRef("NUMBER".to_string());
let expanded =
expand_pattern_symbols(®ex, &table).expect("test fixture: expansion must be Ok");
let formatted = format!("{}", expanded);
assert!(
formatted.contains("[0-9]") || formatted.contains("[0123456789]"),
"Expected expanded pattern to contain digit class, got: {}",
formatted
);
}
#[test]
fn test_cycle_detection_direct() {
let mut table = SymbolTable::new();
let a_pattern = Regex::GroupRef("A".to_string());
table.add_pattern("A", a_pattern, None);
let regex = Regex::GroupRef("A".to_string());
let result = expand_pattern_symbols(®ex, &table);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(
err.kind,
LLreErrorKind::CyclicPatternReference { .. }
));
}
#[test]
fn test_cycle_detection_indirect() {
let mut table = SymbolTable::new();
let a_pattern = Regex::GroupRef("B".to_string());
table.add_pattern("A", a_pattern, None);
let b_pattern = Regex::GroupRef("A".to_string());
table.add_pattern("B", b_pattern, None);
let regex = Regex::GroupRef("A".to_string());
let result = expand_pattern_symbols(®ex, &table);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(
err.kind,
LLreErrorKind::CyclicPatternReference { .. }
));
}
#[test]
fn test_undefined_pattern_preserved() {
let table = SymbolTable::new();
let regex = Regex::GroupRef("undefined".to_string());
let expanded =
expand_pattern_symbols(®ex, &table).expect("test fixture: expansion must be Ok");
assert!(matches!(expanded, Regex::GroupRef(_)));
}
#[test]
fn test_pattern_reuse_non_cyclic() {
let mut table = SymbolTable::new();
let digit_pattern = parse("[0-9]").expect("test fixture: parse must be Ok");
table.add_pattern("DIGIT", digit_pattern, None);
let regex = Regex::Concat(
Box::new(Regex::GroupRef("DIGIT".to_string())),
Box::new(Regex::Concat(
Box::new(Regex::Char('.')),
Box::new(Regex::GroupRef("DIGIT".to_string())),
)),
);
let result = expand_pattern_symbols(®ex, &table);
assert!(
result.is_ok(),
"Pattern reuse should not trigger cycle detection"
);
}
}