use super::ast::{ContextExpr, ContextPredicate, Regex, RegexFlags, UnicodeNormalization};
use crate::phonetic::features::expand_feature_based;
use crate::phonetic::nfa::types::CharClassChar;
use unicode_normalization::UnicodeNormalization as UnicodeNormTrait;
#[derive(Debug, Clone)]
pub struct TransformResult {
pub regex: Regex,
pub unicode_normalization: Option<UnicodeNormalization>,
pub multiline: bool,
pub dotall: bool,
pub local_distance: Option<u8>,
}
impl TransformResult {
pub fn new(regex: Regex) -> Self {
Self {
regex,
unicode_normalization: None,
multiline: false,
dotall: false,
local_distance: None,
}
}
}
pub fn apply_flags(regex: &Regex) -> TransformResult {
let mut result = TransformResult::new(Regex::Empty);
let transformed = apply_flags_with_context(regex, &RegexFlags::default(), &mut result);
result.regex = transformed;
result
}
fn extract_leftmost_inline_flags(regex: &Regex) -> (Option<RegexFlags>, Regex) {
match regex {
Regex::FlagsGroup { flags, inner: None } => {
(Some(flags.clone()), Regex::Empty)
}
Regex::Concat(a, b) => {
let (flags, remaining_a) = extract_leftmost_inline_flags(a);
if flags.is_some() {
match remaining_a {
Regex::Empty => {
(flags, (**b).clone())
}
_ => {
(flags, Regex::Concat(Box::new(remaining_a), b.clone()))
}
}
} else {
(None, regex.clone())
}
}
_ => (None, regex.clone()),
}
}
fn apply_flags_with_context(
regex: &Regex,
inherited: &RegexFlags,
result: &mut TransformResult,
) -> Regex {
match regex {
Regex::FlagsGroup { flags, inner } => {
let merged = inherited.merge(flags);
if let Some(norm) = merged.unicode_normalization {
result.unicode_normalization = Some(norm);
}
if merged.multiline == Some(true) {
result.multiline = true;
}
if merged.dotall == Some(true) {
result.dotall = true;
}
if let Some(dist) = merged.local_distance {
result.local_distance = Some(dist);
}
match inner {
Some(inner) => apply_flags_with_context(inner, &merged, result),
None => Regex::Empty,
}
}
Regex::Char(c) => {
let case_insensitive = inherited.case_insensitive == Some(true);
let accent_insensitive = inherited.accent_insensitive == Some(true);
let feature_based = inherited.feature_based == Some(true);
if case_insensitive || accent_insensitive || feature_based {
expand_char(*c, case_insensitive, accent_insensitive, feature_based)
} else {
Regex::Char(*c)
}
}
Regex::CharClass(class) => {
let case_insensitive = inherited.case_insensitive == Some(true);
let accent_insensitive = inherited.accent_insensitive == Some(true);
let feature_based = inherited.feature_based == Some(true);
if case_insensitive || accent_insensitive || feature_based {
expand_char_class(class, case_insensitive, accent_insensitive, feature_based)
} else {
Regex::CharClass(class.clone())
}
}
Regex::Concat(_, _) => {
let (leftmost_flags, remaining) = extract_leftmost_inline_flags(regex);
if let Some(flags) = leftmost_flags {
let merged = inherited.merge(&flags);
if let Some(norm) = merged.unicode_normalization {
result.unicode_normalization = Some(norm);
}
if merged.multiline == Some(true) {
result.multiline = true;
}
if merged.dotall == Some(true) {
result.dotall = true;
}
if let Some(dist) = merged.local_distance {
result.local_distance = Some(dist);
}
apply_flags_with_context(&remaining, &merged, result)
} else {
if let Regex::Concat(a, b) = regex {
let a_transformed = apply_flags_with_context(a, inherited, result);
let b_transformed = apply_flags_with_context(b, inherited, result);
match (&a_transformed, &b_transformed) {
(Regex::Empty, _) => b_transformed,
(_, Regex::Empty) => a_transformed,
_ => Regex::Concat(Box::new(a_transformed), Box::new(b_transformed)),
}
} else {
unreachable!("We matched Regex::Concat above")
}
}
}
Regex::Alt(a, b) => {
let a_transformed = apply_flags_with_context(a, inherited, result);
let b_transformed = apply_flags_with_context(b, inherited, result);
Regex::Alt(Box::new(a_transformed), Box::new(b_transformed))
}
Regex::Star(inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::Star(Box::new(inner_transformed))
}
Regex::Plus(inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::Plus(Box::new(inner_transformed))
}
Regex::Optional(inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::Optional(Box::new(inner_transformed))
}
Regex::RepeatExact(inner, n) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::RepeatExact(Box::new(inner_transformed), *n)
}
Regex::RepeatRange(inner, min, max) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::RepeatRange(Box::new(inner_transformed), *min, *max)
}
Regex::CapturingGroup(num, inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::CapturingGroup(*num, Box::new(inner_transformed))
}
Regex::NonCapturingGroup(inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::NonCapturingGroup(Box::new(inner_transformed))
}
Regex::NamedGroup(name, inner) => {
let inner_transformed = apply_flags_with_context(inner, inherited, result);
Regex::NamedGroup(name.clone(), Box::new(inner_transformed))
}
Regex::RewriteRule {
pattern,
replacement,
context,
weight,
} => {
let pattern_transformed = apply_flags_with_context(pattern, inherited, result);
let replacement_transformed = apply_flags_with_context(replacement, inherited, result);
let context_transformed = context
.as_ref()
.map(|ctx| Box::new(transform_context_predicate(ctx, inherited, result)));
Regex::RewriteRule {
pattern: Box::new(pattern_transformed),
replacement: Box::new(replacement_transformed),
context: context_transformed,
weight: *weight,
}
}
Regex::Empty => Regex::Empty,
Regex::Any => {
let dotall = inherited.dotall == Some(true);
if dotall {
Regex::Any
} else {
let mut class = CharClassChar::new();
class.add_char('\r');
class.add_char('\n');
class.negated = true;
Regex::CharClass(class)
}
}
Regex::GroupRef(name) => Regex::GroupRef(name.clone()),
Regex::WordBoundary => Regex::WordBoundary,
Regex::StartOfLine => Regex::StartOfLine,
Regex::EndOfLine => Regex::EndOfLine,
Regex::StartOfInput => Regex::StartOfInput,
Regex::EndOfInput => Regex::EndOfInput,
Regex::EndOfInputStrict => Regex::EndOfInputStrict,
}
}
fn transform_context_predicate(
ctx: &ContextPredicate,
inherited: &RegexFlags,
result: &mut TransformResult,
) -> ContextPredicate {
ContextPredicate {
left: ctx
.left
.as_ref()
.map(|e| transform_context_expr(e, inherited, result)),
right: ctx
.right
.as_ref()
.map(|e| transform_context_expr(e, inherited, result)),
syllable: ctx.syllable.clone(),
}
}
fn transform_context_expr(
expr: &ContextExpr,
inherited: &RegexFlags,
result: &mut TransformResult,
) -> ContextExpr {
match expr {
ContextExpr::Pattern(regex) => {
ContextExpr::Pattern(apply_flags_with_context(regex, inherited, result))
}
ContextExpr::WordBoundary => ContextExpr::WordBoundary,
ContextExpr::And(a, b) => ContextExpr::And(
Box::new(transform_context_expr(a, inherited, result)),
Box::new(transform_context_expr(b, inherited, result)),
),
ContextExpr::Or(a, b) => ContextExpr::Or(
Box::new(transform_context_expr(a, inherited, result)),
Box::new(transform_context_expr(b, inherited, result)),
),
ContextExpr::Not(inner) => {
ContextExpr::Not(Box::new(transform_context_expr(inner, inherited, result)))
}
}
}
fn expand_char(
c: char,
case_insensitive: bool,
accent_insensitive: bool,
feature_based: bool,
) -> Regex {
let mut chars = vec![c];
if case_insensitive {
add_case_variants(&mut chars);
}
if accent_insensitive {
add_accent_variants(&mut chars);
}
if feature_based {
add_feature_variants(&mut chars);
}
chars.sort();
chars.dedup();
if chars.len() == 1 {
Regex::Char(chars[0])
} else {
let class = CharClassChar::from_chars(&chars);
Regex::CharClass(class)
}
}
fn expand_char_class(
class: &CharClassChar,
case_insensitive: bool,
accent_insensitive: bool,
feature_based: bool,
) -> Regex {
let mut all_chars: Vec<char> = Vec::new();
for &(start, end) in &class.ranges {
for c in start..=end {
all_chars.push(c);
}
}
let mut expanded: Vec<char> = Vec::new();
for c in all_chars {
expanded.push(c);
if case_insensitive {
add_case_variants_for_char(c, &mut expanded);
}
if accent_insensitive {
add_accent_variants_for_char(c, &mut expanded);
}
if feature_based {
add_feature_variants_for_char(c, &mut expanded);
}
}
expanded.sort();
expanded.dedup();
let new_class = CharClassChar {
ranges: chars_to_ranges(&expanded),
negated: class.negated,
};
Regex::CharClass(new_class)
}
fn add_case_variants(chars: &mut Vec<char>) {
let original: Vec<char> = chars.clone();
for c in original {
add_case_variants_for_char(c, chars);
}
}
fn add_case_variants_for_char(c: char, chars: &mut Vec<char>) {
for lower in c.to_lowercase() {
if !chars.contains(&lower) {
chars.push(lower);
}
}
for upper in c.to_uppercase() {
if !chars.contains(&upper) {
chars.push(upper);
}
}
}
fn add_accent_variants(chars: &mut Vec<char>) {
let original: Vec<char> = chars.clone();
for c in original {
add_accent_variants_for_char(c, chars);
}
}
fn add_accent_variants_for_char(c: char, chars: &mut Vec<char>) {
let base = get_base_char(c);
if !chars.contains(&base) {
chars.push(base);
}
let variants = get_accent_variants(base);
for v in variants {
if !chars.contains(&v) {
chars.push(v);
}
}
}
fn add_feature_variants(chars: &mut Vec<char>) {
let original: Vec<char> = chars.clone();
for c in original {
add_feature_variants_for_char(c, chars);
}
}
fn add_feature_variants_for_char(c: char, chars: &mut Vec<char>) {
let variants = expand_feature_based(c);
for v in variants {
if !chars.contains(&v) {
chars.push(v);
}
}
}
fn get_base_char(c: char) -> char {
c.to_string().nfd().next().unwrap_or(c)
}
fn get_accent_variants(base: char) -> Vec<char> {
match base {
'a' => vec!['a', 'à', 'á', 'â', 'ã', 'ä', 'å', 'ā', 'ă', 'ą'],
'e' => vec!['e', 'è', 'é', 'ê', 'ë', 'ē', 'ĕ', 'ė', 'ę', 'ě'],
'i' => vec!['i', 'ì', 'í', 'î', 'ï', 'ĩ', 'ī', 'ĭ', 'į', 'ı'],
'o' => vec!['o', 'ò', 'ó', 'ô', 'õ', 'ö', 'ō', 'ŏ', 'ő', 'ø'],
'u' => vec!['u', 'ù', 'ú', 'û', 'ü', 'ũ', 'ū', 'ŭ', 'ů', 'ű', 'ų'],
'y' => vec!['y', 'ý', 'ÿ', 'ŷ'],
'c' => vec!['c', 'ç', 'ć', 'ĉ', 'č'],
'd' => vec!['d', 'ď', 'đ'],
'g' => vec!['g', 'ĝ', 'ğ', 'ġ', 'ģ'],
'h' => vec!['h', 'ĥ', 'ħ'],
'j' => vec!['j', 'ĵ'],
'k' => vec!['k', 'ķ'],
'l' => vec!['l', 'ĺ', 'ļ', 'ľ', 'ł'],
'n' => vec!['n', 'ñ', 'ń', 'ņ', 'ň'],
'r' => vec!['r', 'ŕ', 'ŗ', 'ř'],
's' => vec!['s', 'ś', 'ŝ', 'ş', 'š'],
't' => vec!['t', 'ţ', 'ť', 'ŧ'],
'w' => vec!['w', 'ŵ'],
'z' => vec!['z', 'ź', 'ż', 'ž'],
'A' => vec!['A', 'À', 'Á', 'Â', 'Ã', 'Ä', 'Å', 'Ā', 'Ă', 'Ą'],
'E' => vec!['E', 'È', 'É', 'Ê', 'Ë', 'Ē', 'Ĕ', 'Ė', 'Ę', 'Ě'],
'I' => vec!['I', 'Ì', 'Í', 'Î', 'Ï', 'Ĩ', 'Ī', 'Ĭ', 'Į'],
'O' => vec!['O', 'Ò', 'Ó', 'Ô', 'Õ', 'Ö', 'Ō', 'Ŏ', 'Ő', 'Ø'],
'U' => vec!['U', 'Ù', 'Ú', 'Û', 'Ü', 'Ũ', 'Ū', 'Ŭ', 'Ů', 'Ű', 'Ų'],
'Y' => vec!['Y', 'Ý', 'Ÿ', 'Ŷ'],
'C' => vec!['C', 'Ç', 'Ć', 'Ĉ', 'Č'],
'D' => vec!['D', 'Ď', 'Đ'],
'G' => vec!['G', 'Ĝ', 'Ğ', 'Ġ', 'Ģ'],
'H' => vec!['H', 'Ĥ', 'Ħ'],
'J' => vec!['J', 'Ĵ'],
'K' => vec!['K', 'Ķ'],
'L' => vec!['L', 'Ĺ', 'Ļ', 'Ľ', 'Ł'],
'N' => vec!['N', 'Ñ', 'Ń', 'Ņ', 'Ň'],
'R' => vec!['R', 'Ŕ', 'Ŗ', 'Ř'],
'S' => vec!['S', 'Ś', 'Ŝ', 'Ş', 'Š'],
'T' => vec!['T', 'Ţ', 'Ť', 'Ŧ'],
'W' => vec!['W', 'Ŵ'],
'Z' => vec!['Z', 'Ź', 'Ż', 'Ž'],
'æ' => vec!['æ', 'ǽ'],
'Æ' => vec!['Æ', 'Ǽ'],
'œ' => vec!['œ'],
'Œ' => vec!['Œ'],
'ß' => vec!['ß'],
_ => vec![base],
}
}
fn chars_to_ranges(chars: &[char]) -> Vec<(char, char)> {
if chars.is_empty() {
return Vec::new();
}
let mut ranges = Vec::new();
let mut start = chars[0];
let mut end = chars[0];
for &c in chars.iter().skip(1) {
if c as u32 == end as u32 + 1 {
end = c;
} else {
ranges.push((start, end));
start = c;
end = c;
}
}
ranges.push((start, end));
ranges
}
pub fn normalize_input(input: &str, form: UnicodeNormalization) -> String {
match form {
UnicodeNormalization::NFC => input.nfc().collect(),
UnicodeNormalization::NFD => input.nfd().collect(),
UnicodeNormalization::NFKC => input.nfkc().collect(),
UnicodeNormalization::NFKD => input.nfkd().collect(),
}
}
pub fn extract_flags(regex: &Regex) -> RegexFlags {
match regex {
Regex::FlagsGroup { flags, .. } => flags.clone(),
_ => RegexFlags::default(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::phonetic::regex::parse;
#[test]
fn test_case_insensitive_char() {
let expanded = expand_char('a', true, false, false);
match expanded {
Regex::CharClass(class) => {
assert!(class.matches('a'));
assert!(class.matches('A'));
}
_ => panic!("Expected CharClass"),
}
}
#[test]
fn test_accent_insensitive_char() {
let expanded = expand_char('e', false, true, false);
match expanded {
Regex::CharClass(class) => {
assert!(class.matches('e'));
assert!(class.matches('é'));
assert!(class.matches('è'));
assert!(class.matches('ê'));
assert!(class.matches('ë'));
}
_ => panic!("Expected CharClass"),
}
}
#[test]
fn test_combined_flags() {
let expanded = expand_char('e', true, true, false);
match expanded {
Regex::CharClass(class) => {
assert!(class.matches('e'));
assert!(class.matches('E'));
assert!(class.matches('é'));
assert!(class.matches('É'));
}
_ => panic!("Expected CharClass"),
}
}
#[test]
fn test_no_expansion_for_digit() {
let expanded = expand_char('5', true, true, false);
match expanded {
Regex::Char('5') => {}
_ => panic!("Expected unchanged Char('5')"),
}
}
#[test]
fn test_feature_based_char() {
let expanded = expand_char('p', false, false, true);
match expanded {
Regex::CharClass(class) => {
assert!(class.matches('p'));
assert!(class.matches('b')); }
_ => panic!("Expected CharClass"),
}
}
#[test]
fn test_feature_based_voiced_unvoiced_pairs() {
for (voiceless, voiced) in [('p', 'b'), ('t', 'd'), ('k', 'g'), ('f', 'v'), ('s', 'z')] {
let expanded = expand_char(voiceless, false, false, true);
match expanded {
Regex::CharClass(class) => {
assert!(class.matches(voiceless), "Expected {} in class", voiceless);
assert!(
class.matches(voiced),
"Expected {} in class for {}",
voiced,
voiceless
);
}
_ => panic!("Expected CharClass for {}", voiceless),
}
}
}
#[test]
fn test_apply_flags_case_insensitive() {
let regex = parse("(?i:abc)").expect("should parse");
let result = apply_flags(®ex);
let regex_str = format!("{}", result.regex);
assert!(regex_str.contains('[') || regex_str.len() > 3);
}
#[test]
fn test_apply_flags_unicode_normalization() {
let regex = parse("(?u:NFC:test)").expect("should parse");
let result = apply_flags(®ex);
assert_eq!(
result.unicode_normalization,
Some(UnicodeNormalization::NFC)
);
}
#[test]
fn test_normalize_input() {
let composed = "café"; let decomposed = "cafe\u{0301}";
let normalized_composed = normalize_input(composed, UnicodeNormalization::NFC);
let normalized_decomposed = normalize_input(decomposed, UnicodeNormalization::NFC);
assert_eq!(normalized_composed, normalized_decomposed);
}
#[test]
fn test_chars_to_ranges() {
let chars = vec!['a', 'b', 'c', 'x', 'y', 'z'];
let ranges = chars_to_ranges(&chars);
assert_eq!(ranges.len(), 2);
assert_eq!(ranges[0], ('a', 'c'));
assert_eq!(ranges[1], ('x', 'z'));
}
#[test]
fn test_multiline_flag_extracted() {
let regex = parse("(?m:^test$)").expect("should parse");
let result = apply_flags(®ex);
assert!(result.multiline);
}
#[test]
fn test_dotall_flag_extracted() {
let regex = parse("(?s:a.b)").expect("should parse");
let result = apply_flags(®ex);
assert!(result.dotall);
}
#[test]
fn test_inline_flags_propagate_to_subsequent_pattern() {
let regex = parse("(?i)abc").expect("should parse");
let result = apply_flags(®ex);
let regex_str = format!("{}", result.regex);
assert!(
regex_str.contains('['),
"Expected character classes in: {}",
regex_str
);
assert!(!regex_str.is_empty(), "Result should not be empty");
}
#[test]
fn test_inline_flags_case_insensitive() {
let inline_regex = parse("(?i)abc").expect("should parse");
let scoped_regex = parse("(?i:abc)").expect("should parse");
let inline_result = apply_flags(&inline_regex);
let scoped_result = apply_flags(&scoped_regex);
let inline_str = format!("{}", inline_result.regex);
let scoped_str = format!("{}", scoped_result.regex);
assert_eq!(inline_str, scoped_str,
"Inline (?i)abc and scoped (?i:abc) should produce same result.\nInline: {}\nScoped: {}",
inline_str, scoped_str);
}
#[test]
fn test_inline_flags_multiline() {
let regex = parse("(?m)^test$").expect("should parse");
let result = apply_flags(®ex);
assert!(
result.multiline,
"Multiline flag should be extracted from inline (?m)"
);
}
#[test]
fn test_inline_flags_dotall() {
let regex = parse("(?s)a.b").expect("should parse");
let result = apply_flags(®ex);
assert!(
result.dotall,
"Dotall flag should be extracted from inline (?s)"
);
}
#[test]
fn test_inline_unicode_normalization() {
let regex = parse("(?u:NFC)test").expect("should parse");
let result = apply_flags(®ex);
assert_eq!(
result.unicode_normalization,
Some(UnicodeNormalization::NFC),
"Unicode normalization should be extracted from inline (?u:NFC)"
);
}
#[test]
fn test_local_distance_scoped() {
let regex = parse("(?;2:test)").expect("should parse");
let result = apply_flags(®ex);
assert_eq!(
result.local_distance,
Some(2),
"Local distance should be extracted from scoped (?;2:...)"
);
}
#[test]
fn test_local_distance_inline() {
let regex = parse("(?;1)test").expect("should parse");
let result = apply_flags(®ex);
assert_eq!(
result.local_distance,
Some(1),
"Local distance should be extracted from inline (?;N)"
);
}
#[test]
fn test_local_distance_with_flags() {
let regex = parse("(?i;0:test)").expect("should parse");
let result = apply_flags(®ex);
assert_eq!(
result.local_distance,
Some(0),
"Local distance should be extracted from combined flags"
);
let display = format!("{}", result.regex);
assert!(
display.contains("[tT]") || display.contains("Tt"),
"Case insensitive flag should also apply: {}",
display
);
}
}