1use std::cmp::min;
35use std::collections::HashMap;
36
37#[derive(Debug, Clone, PartialEq, Eq)]
39pub enum EditOp {
40 Delete(char),
42 Insert(char),
44 Substitute(char, char),
46 Transpose(char, char),
48}
49
50#[derive(Debug, Clone)]
52pub struct ErrorModel {
53 pub p_deletion: f64,
55 pub p_insertion: f64,
57 pub p_substitution: f64,
59 pub p_transposition: f64,
61 _char_confusion: HashMap<(char, char), f64>,
63 max_edit_distance: usize,
65}
66
67impl Default for ErrorModel {
68 fn default() -> Self {
69 Self {
70 p_deletion: 0.25,
71 p_insertion: 0.25,
72 p_substitution: 0.25,
73 p_transposition: 0.25,
74 _char_confusion: HashMap::new(),
75 max_edit_distance: 2, }
77 }
78}
79
80impl ErrorModel {
81 pub fn new(
83 p_deletion: f64,
84 p_insertion: f64,
85 p_substitution: f64,
86 p_transposition: f64,
87 ) -> Self {
88 let total = p_deletion + p_insertion + p_substitution + p_transposition;
90 Self {
91 p_deletion: p_deletion / total,
92 p_insertion: p_insertion / total,
93 p_substitution: p_substitution / total,
94 p_transposition: p_transposition / total,
95 _char_confusion: HashMap::new(),
96 max_edit_distance: 2,
97 }
98 }
99
100 pub fn with_max_distance(mut self, maxdistance: usize) -> Self {
102 self.max_edit_distance = maxdistance;
103 self
104 }
105
106 pub fn error_probability(&self, typo: &str, correct: &str) -> f64 {
108 if typo == correct {
110 return 1.0;
111 }
112
113 let edit_distance = self.min_edit_operations(typo, correct);
115
116 match edit_distance.len() {
117 0 => 1.0, 1 => {
119 match edit_distance[0] {
121 EditOp::Delete(_) => self.p_deletion,
122 EditOp::Insert(_) => self.p_insertion,
123 EditOp::Substitute(_, _) => self.p_substitution,
124 EditOp::Transpose(_, _) => self.p_transposition,
125 }
126 }
127 n => {
128 let base_prob = 0.1f64.powi(n as i32 - 1);
130 let mut prob = base_prob;
131
132 for op in &edit_distance {
133 match op {
134 EditOp::Delete(_) => prob *= self.p_deletion,
135 EditOp::Insert(_) => prob *= self.p_insertion,
136 EditOp::Substitute(_, _) => prob *= self.p_substitution,
137 EditOp::Transpose(_, _) => prob *= self.p_transposition,
138 }
139 }
140
141 prob
142 }
143 }
144 }
145
146 pub fn min_edit_operations(&self, typo: &str, correct: &str) -> Vec<EditOp> {
148 let typo_chars: Vec<char> = typo.chars().collect();
149 let correct_chars: Vec<char> = correct.chars().collect();
150
151 if typo == correct {
153 return vec![];
154 }
155
156 let length_difference_exceeds_threshold =
164 (typo_chars.len() as isize - correct_chars.len() as isize).abs()
165 > self.max_edit_distance as isize;
166 if length_difference_exceeds_threshold {
167 let mut operations = Vec::new();
168 let _distance = self.levenshtein_with_ops_efficient(correct, typo, &mut operations);
169 return operations;
170 }
171
172 if correct_chars.len() == typo_chars.len() + 1 {
174 for i in 0..correct_chars.len() {
176 let mut test_chars = correct_chars.clone();
177 test_chars.remove(i);
178 if test_chars == typo_chars {
179 return vec![EditOp::Delete(correct_chars[i])];
180 }
181 }
182 } else if correct_chars.len() + 1 == typo_chars.len() {
183 for i in 0..typo_chars.len() {
185 let mut test_chars = typo_chars.clone();
186 test_chars.remove(i);
187 if test_chars == correct_chars {
188 return vec![EditOp::Insert(typo_chars[i])];
189 }
190 }
191 } else if correct_chars.len() == typo_chars.len() {
192 let mut diff_positions = Vec::new();
194
195 for i in 0..correct_chars.len() {
196 if correct_chars[i] != typo_chars[i] {
197 diff_positions.push(i);
198 }
199 }
200
201 if diff_positions.len() == 1 {
202 let i = diff_positions[0];
204 return vec![EditOp::Substitute(correct_chars[i], typo_chars[i])];
205 } else if diff_positions.len() == 2 && diff_positions[0] + 1 == diff_positions[1] {
206 let i = diff_positions[0];
207
208 if correct_chars[i] == typo_chars[i + 1] && correct_chars[i + 1] == typo_chars[i] {
210 return vec![EditOp::Transpose(correct_chars[i], correct_chars[i + 1])];
211 }
212 }
213 }
214
215 let mut operations = Vec::new();
217 let _distance = self.levenshtein_with_ops_efficient(correct, typo, &mut operations);
218 operations
219 }
220
221 fn levenshtein_with_ops_efficient(
224 &self,
225 s1: &str,
226 s2: &str,
227 operations: &mut Vec<EditOp>,
228 ) -> usize {
229 let chars1: Vec<char> = s1.chars().collect();
230 let chars2: Vec<char> = s2.chars().collect();
231 let len1 = chars1.len();
232 let len2 = chars2.len();
233
234 if s1 == s2 {
236 return 0;
237 }
238
239 if (len1 as isize - len2 as isize).abs() > self.max_edit_distance as isize {
241 return self.max_edit_distance + 1; }
243
244 let mut prev_row = (0..=len2).collect::<Vec<_>>();
246 let mut curr_row = vec![0; len2 + 1];
247
248 let mut op_matrix = vec![vec![0; len2 + 1]; len1 + 1];
251
252 for j in 1..=len2 {
254 op_matrix[0][j] = 1; }
256
257 for i in 1..=len1 {
258 curr_row[0] = i;
259 op_matrix[i][0] = 2; for j in 1..=len2 {
262 let cost = if chars1[i - 1] == chars2[j - 1] { 0 } else { 1 };
263
264 let del_cost = prev_row[j] + 1;
266 let ins_cost = curr_row[j - 1] + 1;
267 let sub_cost = prev_row[j - 1] + cost;
268
269 curr_row[j] = min(min(del_cost, ins_cost), sub_cost);
271
272 if curr_row[j] == del_cost {
274 op_matrix[i][j] = 2; } else if curr_row[j] == ins_cost {
276 op_matrix[i][j] = 1; } else if cost > 0 {
278 op_matrix[i][j] = 3; } else {
280 op_matrix[i][j] = 0; }
282
283 if i > 1
285 && j > 1
286 && chars1[i - 1] == chars2[j - 2]
287 && chars1[i - 2] == chars2[j - 1]
288 {
289 let trans_cost = prev_row[j - 2] + 1;
290 if trans_cost < curr_row[j] {
291 curr_row[j] = trans_cost;
292 op_matrix[i][j] = 4; }
294 }
295 }
296
297 if curr_row.iter().all(|&c| c > self.max_edit_distance) {
299 return self.max_edit_distance + 1;
300 }
301
302 std::mem::swap(&mut prev_row, &mut curr_row);
304 }
305
306 let mut i = len1;
308 let mut j = len2;
309 let mut backtrack_ops = Vec::new();
310
311 while i > 0 || j > 0 {
312 match if i == 0 || j == 0 {
313 if i == 0 {
314 1
315 } else {
316 2
317 } } else {
319 op_matrix[i][j]
320 } {
321 0 => {
322 i -= 1;
324 j -= 1;
325 }
326 1 => {
327 j -= 1;
329 backtrack_ops.push(EditOp::Insert(chars2[j]));
330 }
331 2 => {
332 i -= 1;
334 backtrack_ops.push(EditOp::Delete(chars1[i]));
335 }
336 3 => {
337 i -= 1;
339 j -= 1;
340 backtrack_ops.push(EditOp::Substitute(chars1[i], chars2[j]));
341 }
342 4 => {
343 i -= 2;
345 j -= 2;
346 backtrack_ops.push(EditOp::Transpose(chars1[i + 1], chars1[i + 2]));
347 }
348 _ => break, }
350 }
351
352 backtrack_ops.reverse();
354 operations.extend(backtrack_ops);
355
356 prev_row[len2]
358 }
359
360 pub fn levenshtein_with_ops(&self, s1: &str, s2: &str, operations: &mut Vec<EditOp>) -> usize {
362 let chars1: Vec<char> = s1.chars().collect();
363 let chars2: Vec<char> = s2.chars().collect();
364 let len1 = chars1.len();
365 let len2 = chars2.len();
366
367 let mut matrix = vec![vec![0; len2 + 1]; len1 + 1];
369
370 for (i, row) in matrix.iter_mut().enumerate().take(len1 + 1) {
372 row[0] = i;
373 }
374
375 for j in 0..=len2 {
376 matrix[0][j] = j;
377 }
378
379 for i in 1..=len1 {
381 for j in 1..=len2 {
382 let cost = if chars1[i - 1] == chars2[j - 1] { 0 } else { 1 };
383
384 matrix[i][j] = min(
385 min(
386 matrix[i - 1][j] + 1, matrix[i][j - 1] + 1, ),
389 matrix[i - 1][j - 1] + cost, );
391
392 if i > 1
394 && j > 1
395 && chars1[i - 1] == chars2[j - 2]
396 && chars1[i - 2] == chars2[j - 1]
397 {
398 matrix[i][j] = min(
399 matrix[i][j],
400 matrix[i - 2][j - 2] + 1, );
402 }
403 }
404 }
405
406 let mut i = len1;
408 let mut j = len2;
409
410 let mut temp_ops = Vec::new();
412
413 while i > 0 || j > 0 {
414 if i > 0 && j > 0 && chars1[i - 1] == chars2[j - 1] {
415 i -= 1;
417 j -= 1;
418 } else if i > 1
419 && j > 1
420 && chars1[i - 1] == chars2[j - 2]
421 && chars1[i - 2] == chars2[j - 1]
422 && matrix[i][j] == matrix[i - 2][j - 2] + 1
423 {
424 temp_ops.push(EditOp::Transpose(chars1[i - 2], chars1[i - 1]));
426 i -= 2;
427 j -= 2;
428 } else if i > 0 && j > 0 && matrix[i][j] == matrix[i - 1][j - 1] + 1 {
429 temp_ops.push(EditOp::Substitute(chars1[i - 1], chars2[j - 1]));
431 i -= 1;
432 j -= 1;
433 } else if i > 0 && matrix[i][j] == matrix[i - 1][j] + 1 {
434 temp_ops.push(EditOp::Delete(chars1[i - 1]));
436 i -= 1;
437 } else if j > 0 && matrix[i][j] == matrix[i][j - 1] + 1 {
438 temp_ops.push(EditOp::Insert(chars2[j - 1]));
440 j -= 1;
441 } else {
442 break;
444 }
445 }
446
447 temp_ops.reverse();
449 operations.extend(temp_ops);
450
451 matrix[len1][len2]
452 }
453}
454
455#[cfg(test)]
456mod tests {
457 use super::*;
458
459 #[test]
460 fn test_error_model() {
461 let error_model = ErrorModel::default();
462
463 let p_deletion = error_model.error_probability("cat", "cart"); let p_insertion = error_model.error_probability("cart", "cat"); let p_substitution = error_model.error_probability("cat", "cut"); let p_transposition = error_model.error_probability("form", "from"); assert!(p_deletion > 0.0);
471 assert!(p_insertion > 0.0);
472 assert!(p_substitution > 0.0);
473 assert!(p_transposition > 0.0);
474
475 assert_eq!(error_model.error_probability("word", "word"), 1.0);
477 }
478
479 #[test]
480 fn test_edit_operations() {
481 let error_model = ErrorModel::default();
482
483 let ops = error_model.min_edit_operations("cat", "cart");
485 assert_eq!(ops.len(), 1);
486 assert!(matches!(ops[0], EditOp::Delete('r')));
487
488 let ops = error_model.min_edit_operations("cart", "cat");
490 assert_eq!(ops.len(), 1);
491 assert!(matches!(ops[0], EditOp::Insert('r')));
492
493 let ops = error_model.min_edit_operations("cut", "cat");
495 assert_eq!(ops.len(), 1);
496 assert!(matches!(ops[0], EditOp::Substitute('a', 'u')));
497
498 let ops = error_model.min_edit_operations("from", "form");
500 assert_eq!(ops.len(), 1);
501 assert!(matches!(ops[0], EditOp::Transpose('o', 'r')));
502 }
503
504 #[test]
505 fn test_efficient_levenshtein() {
506 let error_model = ErrorModel::default();
507
508 let mut ops1 = Vec::new();
510 let mut ops2 = Vec::new();
511 let dist1 = error_model.levenshtein_with_ops("hello", "hello", &mut ops1);
512 let dist2 = error_model.levenshtein_with_ops_efficient("hello", "hello", &mut ops2);
513 assert_eq!(dist1, 0);
514 assert_eq!(dist2, 0);
515 assert!(ops1.is_empty());
516 assert!(ops2.is_empty());
517
518 let test_cases = [
520 ("cat", "bat"), ("cat", "cats"), ("cats", "cat"), ];
524
525 for (s1, s2) in test_cases {
526 let mut ops1 = Vec::new();
527 let mut ops2 = Vec::new();
528 let dist1 = error_model.levenshtein_with_ops(s1, s2, &mut ops1);
529 let dist2 = error_model.levenshtein_with_ops_efficient(s1, s2, &mut ops2);
530
531 assert_eq!(dist1, 1);
533 assert_eq!(dist2, 1);
534 }
535
536 let mut ops1 = Vec::new();
539 let mut ops2 = Vec::new();
540 error_model.levenshtein_with_ops("abc", "acb", &mut ops1);
541 error_model.levenshtein_with_ops_efficient("abc", "acb", &mut ops2);
542 assert!(ops1.len() <= 2); assert!(ops2.len() <= 2);
544
545 let mut ops1 = Vec::new();
547 let mut ops2 = Vec::new();
548 let dist1 = error_model.levenshtein_with_ops("programming", "programmer", &mut ops1);
549 let dist2 =
550 error_model.levenshtein_with_ops_efficient("programming", "programmer", &mut ops2);
551 assert!(dist1 <= 3); assert!(dist2 <= 3);
553 }
554
555 #[test]
556 fn test_early_termination() {
557 let error_model = ErrorModel::default().with_max_distance(1);
559
560 let ops = error_model.min_edit_operations("cat", "dog");
562
563 if !ops.is_empty() {
566 assert!(matches!(ops[0], EditOp::Substitute(_, _)) || ops.len() > 1);
568 }
569
570 let error_model = ErrorModel::default().with_max_distance(3);
572
573 let ops = error_model.min_edit_operations("kitten", "sitting");
575 assert!(!ops.is_empty()); let ops = error_model.min_edit_operations("algorithm", "logarithm");
579 if ops.len() == 1 {
582 assert!(matches!(ops[0], EditOp::Substitute(_, _)));
584 } else {
585 assert!(!ops.is_empty());
587 }
588 }
589}