use crate::units::{BitPeq, Operands, Unit, UnitMap, dispatch};
use rustc_hash::FxHashMap;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Options {
pub insertion_cost: f64,
pub deletion_cost: f64,
pub substitution_cost: f64,
pub transposition_cost: f64,
pub restricted: bool,
}
impl Default for Options {
fn default() -> Self {
Self {
insertion_cost: 1.0,
deletion_cost: 1.0,
substitution_cost: 1.0,
transposition_cost: 1.0,
restricted: false,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct SearchResult {
pub substring: String,
pub distance: f64,
pub offset: isize,
}
pub fn levenshtein(source: &str, target: &str, opts: &Options) -> f64 {
distance_impl(source, target, opts, false)
}
pub fn damerau_levenshtein(source: &str, target: &str, opts: &Options) -> f64 {
distance_impl(source, target, opts, true)
}
pub fn levenshtein_search(source: &str, target: &str, opts: &Options) -> SearchResult {
search_impl(source, target, opts, false)
}
pub fn damerau_levenshtein_search(source: &str, target: &str, opts: &Options) -> SearchResult {
search_impl(source, target, opts, true)
}
#[cfg(feature = "parallel")]
pub fn par_levenshtein_batch(pairs: &[(&str, &str)], opts: &Options) -> Vec<f64> {
use rayon::prelude::*;
pairs
.par_iter()
.map(|(a, b)| levenshtein(a, b, opts))
.collect()
}
#[cfg(feature = "parallel")]
pub fn par_damerau_levenshtein_batch(pairs: &[(&str, &str)], opts: &Options) -> Vec<f64> {
use rayon::prelude::*;
pairs
.par_iter()
.map(|(a, b)| damerau_levenshtein(a, b, opts))
.collect()
}
fn distance_impl(source: &str, target: &str, opts: &Options, damerau: bool) -> f64 {
if !damerau {
if source.is_empty() {
let units = utf16_len(target);
return if opts.insertion_cost == 1.0 {
units as f64
} else {
repeated_cost(units, opts.insertion_cost)
};
}
if target.is_empty() {
let units = utf16_len(source);
return if opts.deletion_cost == 1.0 {
units as f64
} else {
repeated_cost(units, opts.deletion_cost)
};
}
if source.is_ascii() && target.is_ascii() {
return plain_levenshtein(source.as_bytes(), target.as_bytes(), opts);
}
let (source, target) = if is_unit_cost(opts) && source.len().min(target.len()) > 64 {
trim_common_utf8_affixes(source, target)
} else {
(source, target)
};
if source.is_empty() {
return utf16_len(target) as f64;
}
if target.is_empty() {
return utf16_len(source) as f64;
}
if source.is_ascii() && target.is_ascii() {
return plain_levenshtein(source.as_bytes(), target.as_bytes(), opts);
}
const STACK_UNITS: usize = 64;
if source.len() <= STACK_UNITS && target.len() <= STACK_UNITS {
let mut source_units = [0u16; STACK_UNITS];
let mut target_units = [0u16; STACK_UNITS];
let source_len = encode_utf16_into(source, &mut source_units);
let target_len = encode_utf16_into(target, &mut target_units);
return plain_levenshtein(
&source_units[..source_len],
&target_units[..target_len],
opts,
);
}
let source_units: Vec<u16> = source.encode_utf16().collect();
let target_units: Vec<u16> = target.encode_utf16().collect();
return plain_levenshtein(&source_units, &target_units, opts);
}
dispatch(source, target, |ops| match ops {
Operands::Bytes(s, t) => distance_generic(s, t, opts, damerau),
Operands::Units(s, t) => distance_generic(s, t, opts, damerau),
})
}
#[inline]
fn repeated_cost(count: usize, cost: f64) -> f64 {
(0..count).fold(0.0, |total, _| total + cost)
}
#[inline]
fn utf16_len(input: &str) -> usize {
if input.is_ascii() {
input.len()
} else {
input.encode_utf16().count()
}
}
fn trim_common_utf8_affixes<'a>(mut source: &'a str, mut target: &'a str) -> (&'a str, &'a str) {
let mut prefix = source
.as_bytes()
.iter()
.zip(target.as_bytes())
.take_while(|(a, b)| a == b)
.count();
while prefix != 0 && (!source.is_char_boundary(prefix) || !target.is_char_boundary(prefix)) {
prefix -= 1;
}
source = &source[prefix..];
target = &target[prefix..];
let mut suffix = source
.as_bytes()
.iter()
.rev()
.zip(target.as_bytes().iter().rev())
.take_while(|(a, b)| a == b)
.count();
while suffix != 0
&& (!source.is_char_boundary(source.len() - suffix)
|| !target.is_char_boundary(target.len() - suffix))
{
suffix -= 1;
}
if suffix != 0 {
source = &source[..source.len() - suffix];
target = &target[..target.len() - suffix];
}
(source, target)
}
#[inline]
fn encode_utf16_into(input: &str, out: &mut [u16]) -> usize {
let mut len = 0usize;
for (slot, unit) in out.iter_mut().zip(input.encode_utf16()) {
*slot = unit;
len += 1;
}
len
}
fn distance_generic<T: BitPeq + DamerauScratch>(
source: &[T],
target: &[T],
opts: &Options,
damerau: bool,
) -> f64 {
match (damerau, opts.restricted) {
(true, false) => unrestricted_damerau(source, target, opts),
(true, true) => restricted_damerau(source, target, opts),
(false, _) => plain_levenshtein(source, target, opts),
}
}
fn unrestricted_damerau<T: BitPeq + DamerauScratch>(
source: &[T],
target: &[T],
opts: &Options,
) -> f64 {
if is_unit_cost(opts)
&& opts.transposition_cost == 1.0
&& source.len().saturating_add(target.len()) < u32::MAX as usize
{
return T::damerau_unit_dispatch(source, target);
}
full_matrix(source, target, opts, true, false).final_cost()
}
fn restricted_damerau<T: BitPeq>(source: &[T], target: &[T], opts: &Options) -> f64 {
if is_unit_cost(opts) && opts.transposition_cost == 1.0 {
let (shorter, longer) = if source.len() <= target.len() {
(source, target)
} else {
(target, source)
};
if (2..=64).contains(&shorter.len()) {
return osa_bit_vector(shorter, longer);
}
if shorter.len() > 64 {
return osa_bit_vector_blocks(shorter, longer);
}
}
restricted_rows(source, target, opts)
}
#[inline]
fn is_unit_cost(opts: &Options) -> bool {
opts.insertion_cost == 1.0 && opts.deletion_cost == 1.0 && opts.substitution_cost == 1.0
}
fn plain_levenshtein<T: BitPeq>(source: &[T], target: &[T], opts: &Options) -> f64 {
if is_unit_cost(opts) {
let (source, target) = if source.len().min(target.len()) > 16 {
trim_common_affixes(source, target)
} else {
(source, target)
};
if source.is_empty() {
return target.len() as f64;
}
if target.is_empty() {
return source.len() as f64;
}
let (shorter, longer) = if source.len() <= target.len() {
(source, target)
} else {
(target, source)
};
if (1..=4).contains(&shorter.len()) {
return bit_vector_distance_tiny(shorter, longer);
}
if (5..=64).contains(&shorter.len()) {
return bit_vector_distance(shorter, longer);
}
if shorter.len() > 64 {
return bit_vector_distance_blocks(shorter, longer);
}
}
plain_rows(source, target, opts)
}
fn trim_common_affixes<'a, T: Unit>(
mut source: &'a [T],
mut target: &'a [T],
) -> (&'a [T], &'a [T]) {
let shared = source.len().min(target.len());
let mut prefix = 0usize;
while prefix < shared && source[prefix] == target[prefix] {
prefix += 1;
}
source = &source[prefix..];
target = &target[prefix..];
let shared = source.len().min(target.len());
let mut suffix = 0usize;
while suffix < shared && source[source.len() - 1 - suffix] == target[target.len() - 1 - suffix]
{
suffix += 1;
}
if suffix != 0 {
source = &source[..source.len() - suffix];
target = &target[..target.len() - suffix];
}
(source, target)
}
fn bit_vector_distance_tiny<T: BitPeq>(shorter: &[T], longer: &[T]) -> f64 {
let m = shorter.len();
debug_assert!((1..=4).contains(&m));
if m == 1 {
return (longer.len() - usize::from(longer.contains(&shorter[0]))) as f64;
}
let last_bit = 1u64 << (m - 1);
let mut pv = (1u64 << m) - 1;
let mut mv = 0u64;
let mut score = m as i64;
for &c in longer {
let mut eq = u64::from(shorter[0] == c);
eq |= u64::from(shorter[1] == c) << 1;
if m > 2 {
eq |= u64::from(shorter[2] == c) << 2;
}
if m > 3 {
eq |= u64::from(shorter[3] == c) << 3;
}
let xv = eq | mv;
let xh = (((eq & pv).wrapping_add(pv)) ^ pv) | eq;
let mut ph = mv | !(xh | pv);
let mut mh = pv & xh;
score += i64::from(ph & last_bit != 0) - i64::from(mh & last_bit != 0);
ph = (ph << 1) | 1;
mh <<= 1;
pv = mh | !(xv | ph);
mv = ph & xv;
}
score as f64
}
fn bit_vector_distance<T: BitPeq>(shorter: &[T], longer: &[T]) -> f64 {
let m = shorter.len();
debug_assert!(m > 0 && m <= 64);
let peq = T::peq1(shorter);
let last_bit = 1u64 << (m - 1);
let mut pv: u64 = if m == 64 { u64::MAX } else { (1u64 << m) - 1 };
let mut mv: u64 = 0;
let mut score = m as i64;
for &c in longer {
let eq = T::peq1_get(&peq, c);
let xv = eq | mv;
let xh = (((eq & pv).wrapping_add(pv)) ^ pv) | eq;
let mut ph = mv | !(xh | pv);
let mut mh = pv & xh;
score += i64::from(ph & last_bit != 0) - i64::from(mh & last_bit != 0);
ph = (ph << 1) | 1;
mh <<= 1;
pv = mh | !(xv | ph);
mv = ph & xv;
}
score as f64
}
fn bit_vector_distance_blocks<T: BitPeq>(shorter: &[T], longer: &[T]) -> f64 {
if shorter.len() <= 256 {
bit_vector_distance_blocks_impl::<T, true>(shorter, longer)
} else {
bit_vector_distance_blocks_impl::<T, false>(shorter, longer)
}
}
fn bit_vector_distance_blocks_impl<T: BitPeq, const FINAL_POPCOUNT: bool>(
shorter: &[T],
longer: &[T],
) -> f64 {
const WORD: usize = 64;
let m = shorter.len();
debug_assert!(m > 0);
let blocks = m.div_ceil(WORD);
let last_block = blocks - 1;
let last_bit = 1u64 << ((m - 1) % WORD);
let peq = T::peqn(shorter, blocks);
let leading_absent = longer
.iter()
.position(|&unit| T::peqn_row(&peq, unit).is_some())
.unwrap_or(longer.len());
if leading_absent == longer.len() {
return m.max(longer.len()) as f64;
}
const STACK_BLOCKS: usize = 16;
let zeros_stack = [0u64; STACK_BLOCKS];
let zeros_heap;
let zeros: &[u64] = if blocks <= STACK_BLOCKS {
&zeros_stack[..blocks]
} else {
zeros_heap = vec![0u64; blocks];
&zeros_heap
};
let mut pv_stack = [u64::MAX; STACK_BLOCKS];
let mut pv_heap;
let pv: &mut [u64] = if blocks <= STACK_BLOCKS {
&mut pv_stack[..blocks]
} else {
pv_heap = vec![u64::MAX; blocks];
&mut pv_heap
};
for (b, word) in pv.iter_mut().enumerate() {
let skipped_here = leading_absent.saturating_sub(b * WORD).min(WORD);
*word = if skipped_here == WORD {
0
} else {
u64::MAX << skipped_here
};
}
let mut mv_stack = [0u64; STACK_BLOCKS];
let mut mv_heap;
let mv: &mut [u64] = if blocks <= STACK_BLOCKS {
&mut mv_stack[..blocks]
} else {
mv_heap = vec![0u64; blocks];
&mut mv_heap
};
let mut score = m.max(leading_absent) as i64;
for &c in &longer[leading_absent..] {
let row = T::peqn_row(&peq, c).unwrap_or(zeros);
let mut hp_carry_in = true;
let mut hn_carry_in = false;
for (b, &eq) in row.iter().enumerate() {
let x = eq | u64::from(hn_carry_in);
let d0 = ((x & pv[b]).wrapping_add(pv[b]) ^ pv[b]) | x | mv[b];
let mut hp = mv[b] | !(d0 | pv[b]);
let mut hn = d0 & pv[b];
let (hp_carry_out, hn_carry_out) = if b == last_block {
if FINAL_POPCOUNT {
(false, false)
} else {
(hp & last_bit != 0, hn & last_bit != 0)
}
} else {
(hp & (1u64 << 63) != 0, hn & (1u64 << 63) != 0)
};
if !FINAL_POPCOUNT && b == last_block {
score += i64::from(hp_carry_out) - i64::from(hn_carry_out);
}
hp = (hp << 1) | u64::from(hp_carry_in);
hn = (hn << 1) | u64::from(hn_carry_in);
pv[b] = hn | !(d0 | hp);
mv[b] = hp & d0;
hp_carry_in = hp_carry_out;
hn_carry_in = hn_carry_out;
}
}
if FINAL_POPCOUNT {
let mut distance = longer.len() as i64;
for b in 0..blocks {
let mask = if b == last_block {
if last_bit == 1u64 << 63 {
u64::MAX
} else {
last_bit | (last_bit - 1)
}
} else {
u64::MAX
};
distance += (pv[b] & mask).count_ones() as i64;
distance -= (mv[b] & mask).count_ones() as i64;
}
distance as f64
} else {
score as f64
}
}
fn osa_bit_vector<T: BitPeq>(shorter: &[T], longer: &[T]) -> f64 {
let m = shorter.len();
debug_assert!((1..=64).contains(&m));
let table = T::peq1(shorter);
let last_bit = 1u64 << (m - 1);
let mut pv: u64 = if m == 64 { u64::MAX } else { (1u64 << m) - 1 };
let mut mv: u64 = 0;
let mut prev_d0: u64 = 0;
let mut prev_pm: u64 = 0;
let mut score = m as i64;
for &c in longer {
let pm_j = T::peq1_get(&table, c);
let tr = (((!prev_d0) & pm_j) << 1) & prev_pm;
let d0 = (((pm_j & pv).wrapping_add(pv)) ^ pv) | pm_j | mv | tr;
let mut hp = mv | !(d0 | pv);
let mut hn = d0 & pv;
if hp & last_bit != 0 {
score += 1;
}
if hn & last_bit != 0 {
score -= 1;
}
hp = (hp << 1) | 1;
hn <<= 1;
pv = hn | !(d0 | hp);
mv = hp & d0;
prev_d0 = d0;
prev_pm = pm_j;
}
score as f64
}
struct OsaBlock {
pv: u64,
mv: u64,
d0: u64,
pm: u64,
}
fn osa_bit_vector_blocks<T: BitPeq>(shorter: &[T], longer: &[T]) -> f64 {
const WORD: usize = 64;
let m = shorter.len();
debug_assert!(m >= 1);
let blocks = m.div_ceil(WORD);
let table = T::peqn(shorter, blocks);
let zeros = vec![0u64; blocks];
let last_bit = 1u64 << ((m - 1) % WORD);
let mut state: Vec<OsaBlock> = (0..blocks)
.map(|_| OsaBlock {
pv: u64::MAX,
mv: 0,
d0: 0,
pm: 0,
})
.collect();
let mut score = m as i64;
for &c in longer {
let row = T::peqn_row(&table, c).unwrap_or(&zeros);
let mut hp_carry: u64 = 1;
let mut hn_carry: u64 = 0;
let mut below_prev_d0: u64 = 0;
let mut below_pm: u64 = 0;
for (b, blk) in state.iter_mut().enumerate() {
let pm_j = row[b];
let prev_d0 = blk.d0;
let prev_pm = blk.pm;
let tr = ((((!prev_d0) & pm_j) << 1) | (((!below_prev_d0) & below_pm) >> 63)) & prev_pm;
let x = pm_j | hn_carry;
let d0 = (((x & blk.pv).wrapping_add(blk.pv)) ^ blk.pv) | x | blk.mv | tr;
let mut hp = blk.mv | !(d0 | blk.pv);
let mut hn = d0 & blk.pv;
if b == blocks - 1 {
if hp & last_bit != 0 {
score += 1;
}
if hn & last_bit != 0 {
score -= 1;
}
}
let hp_out = hp >> 63;
let hn_out = hn >> 63;
hp = (hp << 1) | hp_carry;
hn = (hn << 1) | hn_carry;
blk.pv = hn | !(d0 | hp);
blk.mv = hp & d0;
blk.d0 = d0;
blk.pm = pm_j;
hp_carry = hp_out;
hn_carry = hn_out;
below_prev_d0 = prev_d0;
below_pm = pm_j;
}
}
score as f64
}
trait DamerauScratch: Unit {
type SymTable;
fn new_table() -> Self::SymTable;
fn get(table: &Self::SymTable, unit: Self) -> Option<(u32, u32)>;
fn set(table: &mut Self::SymTable, unit: Self, row: u32, next_slot: u32) -> u32;
fn damerau_unit_dispatch(source: &[Self], target: &[Self]) -> f64;
}
impl DamerauScratch for u8 {
type SymTable = ([u32; 256], [u32; 256]);
fn new_table() -> Self::SymTable {
([0u32; 256], [0u32; 256])
}
#[inline]
fn get(table: &Self::SymTable, unit: Self) -> Option<(u32, u32)> {
let row = table.0[unit as usize];
(row != 0).then(|| (row, table.1[unit as usize]))
}
#[inline]
fn set(table: &mut Self::SymTable, unit: Self, row: u32, next_slot: u32) -> u32 {
let u = unit as usize;
if table.0[u] == 0 {
table.1[u] = next_slot;
}
table.0[u] = row;
table.1[u]
}
fn damerau_unit_dispatch(source: &[Self], target: &[Self]) -> f64 {
let n = source.len();
let m = target.len();
if n <= 8 && m <= 8 {
return damerau_unit_small(source, target);
}
if n + m <= u16::MAX as usize {
if n <= 128 && m <= 128 {
return damerau_unit_mid(source, target);
}
return damerau_unit_large(source, target);
}
if n + m < u32::MAX as usize {
return damerau_unrestricted_unit::<u8, u32>(source, target);
}
f64::NAN }
}
impl DamerauScratch for u16 {
type SymTable = FxHashMap<u16, (u32, u32)>;
fn new_table() -> Self::SymTable {
FxHashMap::default()
}
#[inline]
fn get(table: &Self::SymTable, unit: Self) -> Option<(u32, u32)> {
table.get(&unit).copied()
}
#[inline]
fn set(table: &mut Self::SymTable, unit: Self, row: u32, next_slot: u32) -> u32 {
let entry = table.entry(unit).or_insert((0, next_slot));
entry.0 = row;
entry.1
}
fn damerau_unit_dispatch(source: &[Self], target: &[Self]) -> f64 {
let total = source.len().saturating_add(target.len());
if total <= u16::MAX as usize {
return damerau_unrestricted_unit::<u16, u16>(source, target);
}
if total < u32::MAX as usize {
return damerau_unrestricted_unit::<u16, u32>(source, target);
}
f64::NAN }
}
trait DamCell: Copy + Ord {
fn from_usize(v: usize) -> Self;
fn to_f64(self) -> f64;
fn plus(self, d: usize) -> Self;
}
impl DamCell for u16 {
#[inline]
fn from_usize(v: usize) -> Self {
v as u16
}
#[inline]
fn to_f64(self) -> f64 {
f64::from(self)
}
#[inline]
fn plus(self, d: usize) -> Self {
self + d as u16
}
}
impl DamCell for u32 {
#[inline]
fn from_usize(v: usize) -> Self {
v as u32
}
#[inline]
fn to_f64(self) -> f64 {
f64::from(self)
}
#[inline]
fn plus(self, d: usize) -> Self {
self + d as u32
}
}
fn damerau_unit_small(source: &[u8], target: &[u8]) -> f64 {
const CAP: usize = 9;
let n = source.len();
let m = target.len();
debug_assert!(n < CAP && m < CAP);
if n == 0 {
return m as f64;
}
if m == 0 {
return n as f64;
}
let w = m + 1;
let mut mat = [0u16; CAP * CAP];
for (c, cell) in mat[..=m].iter_mut().enumerate() {
*cell = c as u16;
}
for r in 1..=n {
mat[r * w] = r as u16;
}
for r in 1..=n {
let s = source[r - 1];
let base = r * w;
let pbase = base - w;
let mut lcm: usize = 0;
for c in 1..=m {
let t = target[c - 1];
let insert = mat[base + c - 1] + 1;
let delete = mat[pbase + c] + 1;
let sub = mat[pbase + c - 1] + u16::from(s != t);
let mut best = insert.min(delete).min(sub);
if r > 1 && c > 1 && lcm != 0 {
if let Some(p) = source[..r].iter().rposition(|&x| x == t) {
let lrm = p + 1;
let before = mat[(lrm - 1) * w + (lcm - 1)];
let gaps = r + c - lrm - lcm - 1;
let transpose = before + gaps as u16;
if transpose < best {
best = transpose;
}
}
}
mat[base + c] = best;
if s == t {
lcm = c;
}
}
}
f64::from(mat[n * w + m])
}
fn damerau_unit_mid(source: &[u8], target: &[u8]) -> f64 {
let n = source.len();
let m = target.len();
if n == 0 {
return m as f64;
}
if m == 0 {
return n as f64;
}
let w = m + 1;
let mut prev: Vec<u16> = (0..=m).map(|v| v as u16).collect();
let mut cur: Vec<u16> = vec![0u16; w];
let mut table = [0u32; 256]; let mut arena: Vec<u16> = Vec::new();
let mut next_slot: u32;
{
let s = source[0];
table[s as usize] = 1; arena.resize(w, 0);
next_slot = 1;
arena[..w].copy_from_slice(&prev);
cur[0] = 1;
let mut left: u16 = 1;
let mut diag = prev[0];
for c in 1..=m {
let t = target[c - 1];
let up = prev[c];
let insert = left + 1;
let delete = up + 1;
let sub = diag + u16::from(s != t);
let best = insert.min(delete).min(sub);
cur[c] = best;
left = best;
diag = up;
}
std::mem::swap(&mut prev, &mut cur);
}
for r in 2..=n {
let s = source[r - 1];
let su = s as usize;
let old = table[su];
let slot = if old == 0 {
let sl = next_slot;
next_slot += 1;
arena.resize(arena.len() + w, 0);
sl
} else {
old >> 16
};
table[su] = (slot << 16) | r as u32;
arena[slot as usize * w..][..w].copy_from_slice(&prev);
let rw = r as u16;
cur[0] = rw;
let t0 = target[0];
let mut diag;
let mut left;
{
let up = prev[1];
let insert = rw + 1;
let delete = up + 1;
let sub = prev[0] + u16::from(s != t0);
let best = insert.min(delete).min(sub);
cur[1] = best;
left = best;
diag = up;
}
let mut lcm: usize = if s == t0 { 1 } else { 0 };
let mut c = 2usize;
if lcm == 0 {
while c <= m {
let t = target[c - 1];
let up = prev[c];
let insert = left + 1;
let delete = up + 1;
let sub = diag + u16::from(s != t);
let best = insert.min(delete).min(sub);
cur[c] = best;
left = best;
diag = up;
if s == t {
lcm = c;
c += 1;
break;
}
c += 1;
}
}
while c <= m {
let t = target[c - 1];
let up = prev[c];
let insert = left + 1;
let delete = up + 1;
let sub = diag + u16::from(s != t);
let mut best = insert.min(delete).min(sub);
let e = table[t as usize];
let lrm = e & 0xFFFF;
if lrm != 0 {
let before = arena[(e >> 16) as usize * w + (lcm - 1)];
let gaps = r + c - lrm as usize - lcm - 1;
let transpose = before + gaps as u16;
if transpose < best {
best = transpose;
}
}
cur[c] = best;
left = best;
diag = up;
if s == t {
lcm = c;
}
c += 1;
}
std::mem::swap(&mut prev, &mut cur);
}
f64::from(prev[m])
}
fn damerau_unit_large(source: &[u8], target: &[u8]) -> f64 {
let n = source.len();
let m = target.len();
if n == 0 {
return m as f64;
}
if m == 0 {
return n as f64;
}
let w = m + 1;
let mut prev: Vec<u16> = (0..=m).map(|v| v as u16).collect();
let mut cur: Vec<u16> = vec![0u16; w];
let mut table = [0u32; 256];
let mut arena: Vec<u16> = Vec::new();
let mut next_slot: u32;
{
let s = source[0];
table[s as usize] = 1;
arena.resize(w, 0);
next_slot = 1;
arena[..w].copy_from_slice(&prev);
cur[0] = 1;
let mut diag = prev[0];
for c in 1..=m {
let t = target[c - 1];
let insert = cur[c - 1] + 1;
let delete = prev[c] + 1;
let sub = diag + u16::from(s != t);
cur[c] = insert.min(delete).min(sub);
diag = prev[c];
}
std::mem::swap(&mut prev, &mut cur);
}
for r in 2..=n {
let s = source[r - 1];
let su = s as usize;
let old = table[su];
let slot = if old == 0 {
let sl = next_slot;
next_slot += 1;
arena.resize(arena.len() + w, 0);
sl
} else {
old >> 16
};
table[su] = (slot << 16) | r as u32;
arena[slot as usize * w..][..w].copy_from_slice(&prev);
cur[0] = r as u16;
let t0 = target[0];
{
let insert = cur[0] + 1;
let delete = prev[1] + 1;
let sub = prev[0] + u16::from(s != t0);
cur[1] = insert.min(delete).min(sub);
}
let mut lcm: usize = if s == t0 { 1 } else { 0 };
let mut diag = prev[1];
let mut c = 2usize;
if lcm == 0 {
while c <= m {
let t = target[c - 1];
let insert = cur[c - 1] + 1;
let delete = prev[c] + 1;
let sub = diag + u16::from(s != t);
cur[c] = insert.min(delete).min(sub);
diag = prev[c];
if s == t {
lcm = c;
c += 1;
break;
}
c += 1;
}
}
while c <= m {
let t = target[c - 1];
let insert = cur[c - 1] + 1;
let delete = prev[c] + 1;
let sub = diag + u16::from(s != t);
let mut best = insert.min(delete).min(sub);
let e = table[t as usize];
let lrm = e & 0xFFFF;
if lrm != 0 {
let before = arena[(e >> 16) as usize * w + (lcm - 1)];
let gaps = r + c - lrm as usize - lcm - 1;
let transpose = before + gaps as u16;
if transpose < best {
best = transpose;
}
}
diag = prev[c];
cur[c] = best;
if s == t {
lcm = c;
}
c += 1;
}
std::mem::swap(&mut prev, &mut cur);
}
f64::from(prev[m])
}
fn damerau_unrestricted_unit<T: BitPeq + DamerauScratch, C: DamCell>(
source: &[T],
target: &[T],
) -> f64 {
let n = source.len();
let m = target.len();
if n == 0 {
return m as f64;
}
if m == 0 {
return n as f64;
}
let w = m + 1;
let mut prev: Vec<C> = (0..=m).map(C::from_usize).collect();
let mut cur: Vec<C> = vec![C::from_usize(0); w];
let mut table = T::new_table();
let mut arena: Vec<C> = Vec::new();
let mut next_slot: u32 = 0;
for r in 1..=n {
let s = source[r - 1];
let slot = T::set(&mut table, s, r as u32, next_slot);
if slot == next_slot {
arena.resize(arena.len() + w, C::from_usize(0));
next_slot += 1;
}
arena[slot as usize * w..][..w].copy_from_slice(&prev);
cur[0] = C::from_usize(r);
let mut lcm: usize = 0; let mut diag = prev[0];
for c in 1..=m {
let t = target[c - 1];
let insert = cur[c - 1].plus(1);
let delete = prev[c].plus(1);
let sub = diag.plus(usize::from(s != t));
let mut best = insert.min(delete).min(sub);
if r > 1 && c > 1 && lcm != 0 {
if let Some((lrm, tslot)) = <T as DamerauScratch>::get(&table, t) {
let before = arena[tslot as usize * w + (lcm - 1)];
let gaps = r + c - lrm as usize - lcm - 1;
let transpose = before.plus(gaps);
if transpose < best {
best = transpose;
}
}
}
diag = prev[c];
cur[c] = best;
if s == t {
lcm = c;
}
}
std::mem::swap(&mut prev, &mut cur);
}
prev[m].to_f64()
}
fn plain_rows<T: Unit>(source: &[T], target: &[T], opts: &Options) -> f64 {
let m = target.len();
let mut row: Vec<f64> = Vec::with_capacity(m + 1);
row.push(0.0);
for c in 1..=m {
row.push(row[c - 1] + opts.insertion_cost);
}
for &s in source {
let mut diag = row[0];
row[0] += opts.deletion_cost;
let mut left = row[0];
for c in 1..=m {
let up = row[c];
let insert = left + opts.insertion_cost;
let delete = up + opts.deletion_cost;
let mut sub = diag;
if s != target[c - 1] {
sub += opts.substitution_cost;
}
let best = min3(insert, delete, sub);
row[c] = best;
diag = up;
left = best;
}
}
row[m]
}
#[cfg(test)]
fn plain_rows_two_oracle<T: Unit>(source: &[T], target: &[T], opts: &Options) -> f64 {
let n = source.len();
let m = target.len();
let mut prev: Vec<f64> = Vec::with_capacity(m + 1);
prev.push(0.0);
for c in 1..=m {
prev.push(prev[c - 1] + opts.insertion_cost);
}
let mut cur = vec![0.0f64; m + 1];
for r in 1..=n {
cur[0] = prev[0] + opts.deletion_cost;
let s = source[r - 1];
for c in 1..=m {
let insert = cur[c - 1] + opts.insertion_cost;
let delete = prev[c] + opts.deletion_cost;
let mut sub = prev[c - 1];
if s != target[c - 1] {
sub += opts.substitution_cost;
}
cur[c] = min3(insert, delete, sub);
}
std::mem::swap(&mut prev, &mut cur);
}
prev[m]
}
fn restricted_rows<T: Unit>(source: &[T], target: &[T], opts: &Options) -> f64 {
let n = source.len();
let m = target.len();
let mut prev2: Vec<f64> = vec![0.0; m + 1];
let mut prev: Vec<f64> = Vec::with_capacity(m + 1);
prev.push(0.0);
for c in 1..=m {
prev.push(prev[c - 1] + opts.insertion_cost);
}
let mut cur = vec![0.0f64; m + 1];
for r in 1..=n {
cur[0] = prev[0] + opts.deletion_cost;
let s = source[r - 1];
for c in 1..=m {
let t = target[c - 1];
let insert = cur[c - 1] + opts.insertion_cost;
let delete = prev[c] + opts.deletion_cost;
let mut sub = prev[c - 1];
if s != t {
sub += opts.substitution_cost;
}
let mut best = min3(insert, delete, sub);
if r > 1 && c > 1 && s == target[c - 2] && source[r - 2] == t {
let transpose = prev2[c - 2] + opts.transposition_cost;
if transpose < best {
best = transpose;
}
}
cur[c] = best;
}
std::mem::swap(&mut prev2, &mut prev);
std::mem::swap(&mut prev, &mut cur);
}
prev[m]
}
#[inline]
fn min3(a: f64, b: f64, c: f64) -> f64 {
let mut best = a;
if b < best {
best = b;
}
if c < best {
best = c;
}
best
}
struct Matrix {
cols: usize,
cost: Vec<f64>,
parent: Vec<(u32, u32)>,
rows: usize,
}
impl Matrix {
#[inline]
fn idx(&self, r: usize, c: usize) -> usize {
r * self.cols + c
}
#[inline]
fn cost_at(&self, r: usize, c: usize) -> f64 {
self.cost[self.idx(r, c)]
}
fn final_cost(&self) -> f64 {
self.cost_at(self.rows - 1, self.cols - 1)
}
}
fn full_matrix<T: Unit>(
source: &[T],
target: &[T],
opts: &Options,
damerau: bool,
search: bool,
) -> Matrix {
let n = source.len();
let m = target.len();
let cols = m + 1;
let mut mat = Matrix {
cols,
cost: vec![0.0; (n + 1) * cols],
parent: vec![(0, 0); (n + 1) * cols],
rows: n + 1,
};
for r in 1..=n {
let i = mat.idx(r, 0);
mat.cost[i] = mat.cost[mat.idx(r - 1, 0)] + opts.deletion_cost;
mat.parent[i] = ((r - 1) as u32, 0);
}
for c in 1..=m {
let i = mat.idx(0, c);
if search {
mat.cost[i] = 0.0;
} else {
mat.cost[i] = mat.cost[mat.idx(0, c - 1)] + opts.insertion_cost;
mat.parent[i] = (0, (c - 1) as u32);
}
}
let unrestricted = damerau && !opts.restricted;
let restricted = damerau && opts.restricted;
let mut last_row_map = T::new_map();
let mut last_col_match: Option<usize> = None;
for r in 1..=n {
if unrestricted {
last_col_match = None;
}
let s = source[r - 1];
for c in 1..=m {
let t = target[c - 1];
let insert = mat.cost_at(r, c - 1) + opts.insertion_cost;
let delete = mat.cost_at(r - 1, c) + opts.deletion_cost;
let mut sub = mat.cost_at(r - 1, c - 1);
if s != t {
sub += opts.substitution_cost;
}
let mut best_cost = insert;
let mut best_parent = (r as u32, (c - 1) as u32);
if delete < best_cost {
best_cost = delete;
best_parent = ((r - 1) as u32, c as u32);
}
if sub < best_cost {
best_cost = sub;
best_parent = ((r - 1) as u32, (c - 1) as u32);
}
if unrestricted && r > 1 && c > 1 {
if let (Some(lcm), Some(lrm)) = (last_col_match, last_row_map.get(t)) {
let before = mat.cost_at(lrm - 1, lcm - 1);
let row_gap = r as isize - lrm as isize - 1;
let col_gap = c as isize - lcm as isize - 1;
let transpose = before
+ (row_gap as f64) * opts.deletion_cost
+ (col_gap as f64) * opts.insertion_cost
+ opts.transposition_cost;
if transpose < best_cost {
best_cost = transpose;
best_parent = ((lrm - 1) as u32, (lcm - 1) as u32);
}
}
}
if restricted && r > 1 && c > 1 && s == target[c - 2] && source[r - 2] == t {
let transpose = mat.cost_at(r - 2, c - 2) + opts.transposition_cost;
if transpose < best_cost {
best_cost = transpose;
best_parent = ((r - 2) as u32, (c - 2) as u32);
}
}
let i = mat.idx(r, c);
mat.cost[i] = best_cost;
mat.parent[i] = best_parent;
if unrestricted {
last_row_map.set(s, r);
if s == t {
last_col_match = Some(c);
}
}
}
}
mat
}
fn search_impl(source: &str, target: &str, opts: &Options, damerau: bool) -> SearchResult {
dispatch(source, target, |ops| match ops {
Operands::Bytes(s, t) => {
let (start, end, dist) = search_generic(s, t, opts, damerau);
SearchResult {
substring: String::from_utf8_lossy(slice_units(t, start, end)).into_owned(),
distance: dist,
offset: start,
}
}
Operands::Units(s, t) => {
let (start, end, dist) = search_generic(s, t, opts, damerau);
SearchResult {
substring: String::from_utf16_lossy(slice_units(t, start, end)),
distance: dist,
offset: start,
}
}
})
}
fn search_generic<T: BitPeq>(
source: &[T],
target: &[T],
opts: &Options,
damerau: bool,
) -> (isize, usize, f64) {
if !damerau && is_unit_cost(opts) && !source.is_empty() && !target.is_empty() {
return search_bits(source, target);
}
search_full_matrix(source, target, opts, damerau)
}
fn search_full_matrix<T: Unit>(
source: &[T],
target: &[T],
opts: &Options,
damerau: bool,
) -> (isize, usize, f64) {
let n = source.len();
let m = target.len();
let mat = full_matrix(source, target, opts, damerau, true);
let mut min_distance = (n + m) as f64;
let mut match_end = m;
for c in 0..=m {
let cost = mat.cost_at(n, c);
if min_distance > cost {
min_distance = cost;
match_end = c;
}
}
let match_start: isize = if match_end == 0 {
0
} else {
let mut row = n;
let mut col = match_end;
while row > 1 && col > 1 {
let (pr, pc) = mat.parent[mat.idx(row, col)];
row = pr as usize;
col = pc as usize;
}
col as isize - 1
};
(match_start, match_end, min_distance)
}
struct SearchColumns {
col_pv: Vec<u64>,
col_mv: Vec<u64>,
blocks: usize,
match_end: usize,
min_distance: i64,
}
fn search_bits<T: BitPeq>(source: &[T], target: &[T]) -> (isize, usize, f64) {
let n = source.len();
let m = target.len();
debug_assert!(n >= 1 && m >= 1);
let fw = if n <= 64 {
search_forward_word(source, target)
} else {
search_forward_blocks(source, target)
};
let match_end = fw.match_end;
let match_start: isize = if match_end == 0 {
0
} else {
let mut row = n;
let mut col = match_end;
while row > 1 && col > 1 {
let insert = search_cell_cost(&fw, row, col - 1) + 1;
let delete = search_cell_cost(&fw, row - 1, col) + 1;
let substitute = search_cell_cost(&fw, row - 1, col - 1)
+ i64::from(source[row - 1] != target[col - 1]);
let mut best = insert;
let mut parent = (row, col - 1);
if delete < best {
best = delete;
parent = (row - 1, col);
}
if substitute < best {
parent = (row - 1, col - 1);
}
(row, col) = parent;
}
col as isize - 1
};
(match_start, match_end, fw.min_distance as f64)
}
fn search_forward_word<T: BitPeq>(source: &[T], target: &[T]) -> SearchColumns {
let n = source.len();
let m = target.len();
debug_assert!((1..=64).contains(&n));
let peq = T::peq1(source);
let last_bit = 1u64 << (n - 1);
let mut pv: u64 = u64::MAX;
let mut mv: u64 = 0;
let mut score = n as i64;
let mut col_pv = vec![0u64; m];
let mut col_mv = vec![0u64; m];
let mut min_distance = (n + m) as i64;
let mut match_end = m;
if min_distance > n as i64 {
min_distance = n as i64;
match_end = 0;
}
for j in 1..=m {
let eq = T::peq1_get(&peq, target[j - 1]);
let xv = eq | mv;
let xh = (((eq & pv).wrapping_add(pv)) ^ pv) | eq;
let mut ph = mv | !(xh | pv);
let mut mh = pv & xh;
if ph & last_bit != 0 {
score += 1;
}
if mh & last_bit != 0 {
score -= 1;
}
ph <<= 1;
mh <<= 1;
pv = mh | !(xv | ph);
mv = ph & xv;
col_pv[j - 1] = pv;
col_mv[j - 1] = mv;
if min_distance > score {
min_distance = score;
match_end = j;
}
}
SearchColumns {
col_pv,
col_mv,
blocks: 1,
match_end,
min_distance,
}
}
fn search_forward_blocks<T: BitPeq>(source: &[T], target: &[T]) -> SearchColumns {
const WORD: usize = 64;
let n = source.len();
let m = target.len();
debug_assert!(n >= 1);
let blocks = n.div_ceil(WORD);
let last_block = blocks - 1;
let last_bit = 1u64 << ((n - 1) % WORD);
let peq = T::peqn(source, blocks);
let zeros = vec![0u64; blocks];
let mut pv = vec![u64::MAX; blocks];
let mut mv = vec![0u64; blocks];
let mut score = n as i64;
let mut col_pv = vec![0u64; m * blocks];
let mut col_mv = vec![0u64; m * blocks];
let mut min_distance = (n + m) as i64;
let mut match_end = m;
if min_distance > n as i64 {
min_distance = n as i64;
match_end = 0;
}
for j in 1..=m {
let row = T::peqn_row(&peq, target[j - 1]).unwrap_or(&zeros);
let mut hp_carry_in = false;
let mut hn_carry_in = false;
for (b, &eq) in row.iter().enumerate() {
let x = eq | u64::from(hn_carry_in);
let d0 = ((x & pv[b]).wrapping_add(pv[b]) ^ pv[b]) | x | mv[b];
let mut hp = mv[b] | !(d0 | pv[b]);
let mut hn = d0 & pv[b];
let (hp_carry_out, hn_carry_out) = if b == last_block {
(hp & last_bit != 0, hn & last_bit != 0)
} else {
(hp & (1u64 << 63) != 0, hn & (1u64 << 63) != 0)
};
if b == last_block {
score += i64::from(hp_carry_out) - i64::from(hn_carry_out);
}
hp = (hp << 1) | u64::from(hp_carry_in);
hn = (hn << 1) | u64::from(hn_carry_in);
pv[b] = hn | !(d0 | hp);
mv[b] = hp & d0;
hp_carry_in = hp_carry_out;
hn_carry_in = hn_carry_out;
}
let base = (j - 1) * blocks;
col_pv[base..base + blocks].copy_from_slice(&pv);
col_mv[base..base + blocks].copy_from_slice(&mv);
if min_distance > score {
min_distance = score;
match_end = j;
}
}
SearchColumns {
col_pv,
col_mv,
blocks,
match_end,
min_distance,
}
}
fn search_cell_cost(fw: &SearchColumns, r: usize, c: usize) -> i64 {
if r == 0 {
return 0;
}
if c == 0 {
return r as i64;
}
let base = (c - 1) * fw.blocks;
let pv = &fw.col_pv[base..base + fw.blocks];
let mv = &fw.col_mv[base..base + fw.blocks];
let full = r / 64;
let mut d = 0i64;
for (&p, &m_word) in pv[..full].iter().zip(&mv[..full]) {
d += i64::from(p.count_ones()) - i64::from(m_word.count_ones());
}
let rem = r % 64;
if rem > 0 {
let mask = (1u64 << rem) - 1;
d += i64::from((pv[full] & mask).count_ones()) - i64::from((mv[full] & mask).count_ones());
}
d
}
fn slice_units<T>(units: &[T], start: isize, end: usize) -> &[T] {
let len = units.len();
let s = if start < 0 {
(len as isize + start).max(0) as usize
} else {
(start as usize).min(len)
};
let e = end.min(len);
if s >= e { &[] } else { &units[s..e] }
}
#[cfg(test)]
mod tests {
use super::*;
fn lev(a: &str, b: &str) -> f64 {
levenshtein(a, b, &Options::default())
}
#[test]
fn classic_distances() {
assert_eq!(lev("kitten", "sitting"), 3.0);
assert_eq!(lev("saturday", "sunday"), 3.0);
assert_eq!(lev("", ""), 0.0);
assert_eq!(lev("abc", ""), 3.0);
assert_eq!(lev("", "abc"), 3.0);
assert_eq!(lev("same", "same"), 0.0);
}
#[test]
fn transposition_only_counts_for_damerau() {
let o = Options::default();
assert_eq!(levenshtein("ab", "ba", &o), 2.0);
assert_eq!(damerau_levenshtein("ab", "ba", &o), 1.0);
}
#[test]
fn restricted_and_unrestricted_damerau_differ() {
let unrestricted = Options {
restricted: false,
..Options::default()
};
let restricted = Options {
restricted: true,
..Options::default()
};
assert_eq!(damerau_levenshtein("ca", "abc", &unrestricted), 2.0);
assert_eq!(damerau_levenshtein("ca", "abc", &restricted), 3.0);
}
#[test]
fn asymmetric_costs_are_respected() {
let o = Options {
deletion_cost: 3.0,
..Options::default()
};
assert_eq!(levenshtein("abc", "ab", &o), 3.0);
assert_eq!(levenshtein("ab", "abc", &o), 1.0);
}
#[test]
fn fractional_and_zero_costs() {
let frac = Options {
insertion_cost: 0.5,
deletion_cost: 1.5,
substitution_cost: 0.75,
..Options::default()
};
assert_eq!(levenshtein("ab", "abc", &frac), 0.5);
let zero = Options {
insertion_cost: 0.0,
deletion_cost: 0.0,
substitution_cost: 0.0,
..Options::default()
};
assert_eq!(levenshtein("kitten", "sitting", &zero), 0.0);
}
#[test]
fn utf16_semantics_match_the_reference() {
assert_eq!(lev("a😀b", "ab"), 2.0);
assert_eq!(lev("😀", ""), 2.0);
assert_eq!(lev("😀", "😀"), 0.0);
}
#[test]
fn bmp_non_ascii_is_one_unit_per_char() {
assert_eq!(lev("café", "cafe"), 1.0);
assert_eq!(lev("Москва", "Москва"), 0.0);
}
#[test]
fn search_finds_best_substring() {
let r = levenshtein_search("ca", "abc", &Options::default());
assert_eq!(r.substring, "a");
assert_eq!(r.distance, 1.0);
assert_eq!(r.offset, 0);
}
#[test]
fn two_row_and_full_matrix_agree() {
let words = [
"kitten", "sitting", "flaw", "lawn", "", "a", "abcdef", "fedcba",
];
for a in words {
for b in words {
for restricted in [false, true] {
let o = Options {
restricted,
..Options::default()
};
let fast = distance_impl(a, b, &o, restricted);
let slow = dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => {
full_matrix(s, t, &o, restricted, false).final_cost()
}
Operands::Units(s, t) => {
full_matrix(s, t, &o, restricted, false).final_cost()
}
});
assert_eq!(fast, slow, "{a:?} vs {b:?} restricted={restricted}");
}
}
}
}
struct Xorshift64(u64);
impl Xorshift64 {
fn next_u64(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn next_range(&mut self, bound: usize) -> usize {
(self.next_u64() % bound as u64) as usize
}
}
fn random_string(rng: &mut Xorshift64, len: usize) -> String {
const ALPHABET: &[u8] = b"abcde";
(0..len)
.map(|_| ALPHABET[rng.next_range(ALPHABET.len())] as char)
.collect()
}
#[test]
fn bit_vector_agrees_with_plain_rows_on_random_pairs() {
let mut rng = Xorshift64(0x5EED_F00D_C0FF_EE42);
let lengths = [0usize, 1, 2, 5, 30, 63, 64, 65, 100, 200];
let opts = Options::default();
for &a_len in &lengths {
for &b_len in &lengths {
for _ in 0..20 {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
let via_fast_path = distance_impl(&a, &b, &opts, false);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(
via_fast_path, via_plain_rows,
"mismatch for {a:?} (len {a_len}) vs {b:?} (len {b_len})"
);
}
}
}
}
#[test]
fn bit_vector_agrees_on_utf16_input() {
let mut rng = Xorshift64(0x1234_5678_9ABC_DEF0);
let opts = Options::default();
let pairs = [
("café", "cafe"),
("Москва", "Масква"),
("😀😀😀", "😀"),
("a😀b😀c", "abc"),
];
for (a, b) in pairs {
let via_fast_path = levenshtein(a, b, &opts);
let via_plain_rows = dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(via_fast_path, via_plain_rows, "mismatch for {a:?} vs {b:?}");
}
const CYRILLIC: &[char] = &['а', 'б', 'в', 'г', 'д'];
for _ in 0..50 {
let a_len = rng.next_range(70);
let b_len = rng.next_range(70);
let a: String = (0..a_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let b: String = (0..b_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(via_fast_path, via_plain_rows, "mismatch for {a:?} vs {b:?}");
}
}
#[test]
fn utf8_affix_pretrim_matches_the_utf16_oracle() {
let opts = Options::default();
let mut cases = Vec::new();
let base = "аб😀中".repeat(100);
let mut changed: Vec<char> = base.chars().collect();
changed[200] = 'ж';
cases.push((base.clone(), changed.into_iter().collect::<String>()));
cases.push((base.clone(), base));
cases.push((
format!("{}é", "x".repeat(65)),
format!("{}©", "x".repeat(65)),
));
cases.push((
format!("{}😀", "д".repeat(40)),
format!("{}😁", "д".repeat(40)),
));
for (source, target) in cases {
let actual = levenshtein(&source, &target, &opts);
let expected = dispatch(&source, &target, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(actual, expected, "{source:?} -> {target:?}");
}
}
#[test]
fn bit_vector_fast_path_only_applies_to_unit_cost() {
let weighted = Options {
insertion_cost: 2.0,
..Options::default()
};
assert_eq!(levenshtein("abc", "ab", &weighted), 1.0); assert_eq!(levenshtein("ab", "abc", &weighted), 2.0); }
#[test]
fn one_row_weighted_matches_two_row_oracle_bit_for_bit() {
let options = [
Options::default(),
Options {
insertion_cost: 0.5,
deletion_cost: 1.5,
substitution_cost: 0.75,
..Options::default()
},
Options {
insertion_cost: 0.0,
deletion_cost: 0.0,
substitution_cost: 0.0,
..Options::default()
},
Options {
insertion_cost: -1.0,
deletion_cost: -0.5,
substitution_cost: 2.0,
..Options::default()
},
Options {
insertion_cost: f64::INFINITY,
deletion_cost: 1.0,
substitution_cost: 0.25,
..Options::default()
},
Options {
insertion_cost: f64::NAN,
deletion_cost: 1.0,
substitution_cost: 0.25,
..Options::default()
},
];
let pairs = [
("", ""),
("", "a😀b"),
("a😀b", ""),
("a", "abcdefghijklmnopqrstuvwxyz"),
("abcdefghijklmnopqrstuvwxyz", "a"),
("kitten", "sitting"),
("Москва", "Масква"),
("😀😃😄", "😃😄😁"),
];
for opts in options {
for (source, target) in pairs {
let (actual, expected) = dispatch(source, target, |ops| match ops {
Operands::Bytes(s, t) => {
(plain_rows(s, t, &opts), plain_rows_two_oracle(s, t, &opts))
}
Operands::Units(s, t) => {
(plain_rows(s, t, &opts), plain_rows_two_oracle(s, t, &opts))
}
});
if expected.is_nan() {
assert!(actual.is_nan(), "{source:?} -> {target:?}, {opts:?}");
} else {
assert_eq!(
actual.to_bits(),
expected.to_bits(),
"{source:?} -> {target:?}, {opts:?}"
);
}
}
}
}
#[test]
fn empty_plain_distance_matches_the_row_recurrence_bit_for_bit() {
for cost in [0.0, -0.0, 0.1, -0.5, f64::INFINITY, f64::NAN] {
let insert = Options {
insertion_cost: cost,
..Options::default()
};
let expected_insert = dispatch("", "a😀b", |ops| match ops {
Operands::Bytes(s, t) => plain_rows_two_oracle(s, t, &insert),
Operands::Units(s, t) => plain_rows_two_oracle(s, t, &insert),
});
let actual_insert = levenshtein("", "a😀b", &insert);
let delete = Options {
deletion_cost: cost,
..Options::default()
};
let expected_delete = dispatch("a😀b", "", |ops| match ops {
Operands::Bytes(s, t) => plain_rows_two_oracle(s, t, &delete),
Operands::Units(s, t) => plain_rows_two_oracle(s, t, &delete),
});
let actual_delete = levenshtein("a😀b", "", &delete);
for (actual, expected) in [
(actual_insert, expected_insert),
(actual_delete, expected_delete),
] {
if expected.is_nan() {
assert!(actual.is_nan());
} else {
assert_eq!(actual.to_bits(), expected.to_bits());
}
}
}
}
fn random_units(rng: &mut Xorshift64, len: usize) -> Vec<u8> {
const ALPHABET: &[u8] = b"abcde";
(0..len)
.map(|_| ALPHABET[rng.next_range(ALPHABET.len())])
.collect()
}
#[test]
fn bit_vector_blocks_agrees_with_plain_rows_on_random_pairs() {
let mut rng = Xorshift64(0xB10C_5EED_1234_5678);
let lengths = [
65usize, 127, 128, 129, 191, 192, 193, 255, 256, 257, 500, 1000,
];
let opts = Options::default();
for &a_len in &lengths {
for &b_len in &lengths {
for _ in 0..3 {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
let via_fast_path = distance_impl(&a, &b, &opts, false);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(
via_fast_path, via_plain_rows,
"mismatch for len {a_len} vs len {b_len}"
);
}
}
}
}
#[test]
fn bit_vector_blocks_skips_absent_prefix_without_changing_state() {
let opts = Options::default();
for pattern_len in [65usize, 129, 257] {
let shorter = vec![b'z'; pattern_len];
for prefix_len in [0usize, 1, 63, 64, 65, 127, 128, 300, 1_000] {
let mut longer = vec![b'a'; prefix_len];
longer.push(b'z');
longer.extend_from_slice(b"bbb");
assert_eq!(
bit_vector_distance_blocks(&shorter, &longer),
plain_rows(&shorter, &longer, &opts),
"pattern={pattern_len}, absent prefix={prefix_len}"
);
}
for target_len in [1usize, pattern_len, pattern_len * 2] {
let longer = vec![b'a'; target_len];
assert_eq!(
bit_vector_distance_blocks(&shorter, &longer),
plain_rows(&shorter, &longer, &opts),
"disjoint pattern={pattern_len}, target={target_len}"
);
}
}
}
#[test]
fn bit_vector_blocks_agrees_with_plain_rows_when_longer_is_much_bigger() {
let mut rng = Xorshift64(0xFEED_0BAD_C0FF_EE99);
let opts = Options::default();
for &shorter_len in &[65usize, 130, 260] {
for _ in 0..2 {
let a = random_string(&mut rng, shorter_len);
let b = random_string(&mut rng, 5000);
let via_fast_path = distance_impl(&a, &b, &opts, false);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(
via_fast_path, via_plain_rows,
"mismatch for shorter_len {shorter_len} vs longer_len 5000"
);
}
}
}
#[test]
fn bit_vector_blocks_agrees_with_bit_vector_distance_at_the_boundary() {
let mut rng = Xorshift64(0xC0DE_FACE_0BAD_F00D);
for shorter_len in 8usize..=64 {
for _ in 0..10 {
let longer_len = rng.next_range(300).max(1);
let shorter = random_units(&mut rng, shorter_len);
let longer = random_units(&mut rng, longer_len);
let via_single = bit_vector_distance(&shorter, &longer);
let via_blocks = bit_vector_distance_blocks(&shorter, &longer);
assert_eq!(
via_single, via_blocks,
"mismatch at shorter_len={shorter_len} longer_len={longer_len}"
);
}
}
}
#[test]
fn bit_vector_blocks_agrees_on_utf16_input() {
let mut rng = Xorshift64(0x9E37_79B9_7F4A_7C15);
let opts = Options::default();
const CYRILLIC: &[char] = &['а', 'б', 'в', 'г', 'д'];
let lengths = [65usize, 128, 200, 500];
for &a_len in &lengths {
for &b_len in &lengths {
let a: String = (0..a_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let b: String = (0..b_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(via_fast_path, via_plain_rows, "mismatch for {a:?} vs {b:?}");
}
}
let pairs = [
("😀".repeat(80), "😀".repeat(79)),
("a😀".repeat(70), "b😀".repeat(70)),
];
for (a, b) in pairs {
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = dispatch(&a, &b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, &opts),
Operands::Units(s, t) => plain_rows(s, t, &opts),
});
assert_eq!(via_fast_path, via_plain_rows, "mismatch for {a:?} vs {b:?}");
}
}
#[test]
fn bit_vector_blocks_matches_hand_computed_edge_cases() {
let opts = Options::default();
let a = "a".repeat(65);
let b = format!("{}b", "a".repeat(64));
assert_eq!(levenshtein(&a, &b, &opts), 1.0);
let c = "abcde".repeat(50); assert_eq!(levenshtein(&c, &c, &opts), 0.0);
let d = "x".repeat(200);
let e = "y".repeat(200);
assert_eq!(levenshtein(&d, &e, &opts), 200.0);
let f = "z".repeat(150);
assert_eq!(levenshtein(&f, "", &opts), 150.0);
}
fn oracle_plain_rows(a: &str, b: &str, opts: &Options) -> f64 {
dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => plain_rows(s, t, opts),
Operands::Units(s, t) => plain_rows(s, t, opts),
})
}
#[test]
fn bit_vector_blocks_full_vs_partial_last_block_boundaries() {
let mut rng = Xorshift64(0xB0BA_FACE_5EED_0001);
let opts = Options::default();
for &len in &[128usize, 129, 192, 193, 256, 257, 320, 321, 384, 385] {
let base = random_string(&mut rng, len);
let base_bytes = base.as_bytes();
assert_eq!(levenshtein(&base, &base, &opts), 0.0, "identical len {len}");
let mut positions: Vec<usize> = vec![0, len - 1];
for boundary in (63..len).step_by(64) {
positions.push(boundary);
if boundary + 1 < len {
positions.push(boundary + 1);
}
if boundary > 0 {
positions.push(boundary - 1);
}
}
positions.sort_unstable();
positions.dedup();
for pos in positions {
let mut mutated = base_bytes.to_vec();
mutated[pos] = b'a' + ((mutated[pos] - b'a' + 1) % 5);
let mutated_str = String::from_utf8(mutated).unwrap();
let via_fast_path = levenshtein(&base, &mutated_str, &opts);
let via_plain_rows = oracle_plain_rows(&base, &mutated_str, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"single substitution at pos {pos} of len {len} mismatch \
(fast={via_fast_path}, oracle={via_plain_rows})"
);
}
}
}
#[test]
fn bit_vector_blocks_all_one_character_repetition() {
let opts = Options::default();
for &(len_a, len_b) in &[
(128usize, 128usize),
(128, 129),
(129, 128),
(192, 64),
(64, 192),
(200, 400),
(400, 200),
(321, 321),
(500, 503),
] {
let a = "a".repeat(len_a);
let b = "a".repeat(len_b);
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"all-'a' mismatch len_a={len_a} len_b={len_b}"
);
}
for &(len_a, len_b) in &[(128usize, 128usize), (200, 200), (321, 321), (256, 300)] {
let a = "a".repeat(len_a);
let b = "b".repeat(len_b);
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"all-'a' vs all-'b' mismatch len_a={len_a} len_b={len_b}"
);
}
}
#[test]
fn bit_vector_blocks_alternating_two_characters() {
let opts = Options::default();
for &len in &[128usize, 129, 192, 200, 256, 257, 320, 400] {
let a: String = (0..len)
.map(|i| if i % 2 == 0 { 'a' } else { 'b' })
.collect();
let b: String = (0..len)
.map(|i| if i % 2 == 0 { 'b' } else { 'a' })
.collect();
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"alternating mismatch len {len}"
);
let b_longer = format!("{b}a");
let via_fast_path2 = levenshtein(&a, &b_longer, &opts);
let via_plain_rows2 = oracle_plain_rows(&a, &b_longer, &opts);
assert_eq!(
via_fast_path2, via_plain_rows2,
"alternating + 1 mismatch len {len}"
);
}
let a: String = (0..500).map(|i| ['a', 'b', 'c'][i % 3]).collect();
let b: String = (0..480).map(|i| ['b', 'a'][i % 2]).collect();
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(via_fast_path, via_plain_rows, "3-cycle vs 2-cycle mismatch");
}
#[test]
fn bit_vector_blocks_disjoint_and_near_identical_multiblock() {
let mut rng = Xorshift64(0xD15C_A5ED_9999_0001);
let opts = Options::default();
for &(len_a, len_b) in &[
(128usize, 128usize),
(200, 200),
(321, 321),
(256, 300),
(300, 256),
] {
let a = "x".repeat(len_a);
let b = "y".repeat(len_b);
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"disjoint mismatch {len_a} vs {len_b}"
);
}
for &len in &[128usize, 192, 256, 320, 400] {
let base = random_string(&mut rng, len);
let mut mutated = base.clone().into_bytes();
for _ in 0..5 {
let pos = rng.next_range(len);
let delta = 1 + rng.next_range(4) as u8;
mutated[pos] = b'a' + ((mutated[pos] - b'a' + delta) % 5);
}
let mutated_str = String::from_utf8(mutated).unwrap();
let via_fast_path = levenshtein(&base, &mutated_str, &opts);
let via_plain_rows = oracle_plain_rows(&base, &mutated_str, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"near-identical mismatch len {len}"
);
}
}
#[test]
fn bit_vector_blocks_boundary_pattern_lengths_against_huge_targets() {
let mut rng = Xorshift64(0x8000_0001_DEAD_10CC);
let opts = Options::default();
for &shorter_len in &[65usize, 129, 193, 257, 321, 385] {
for &longer_len in &[4000usize, 10_007] {
let a = random_string(&mut rng, shorter_len);
let b = random_string(&mut rng, longer_len);
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"boundary pattern len {shorter_len} vs huge target len {longer_len}"
);
}
}
let pattern = "q".repeat(193);
let mut target = "q".repeat(9001).into_bytes();
for i in (100..9000).step_by(777) {
target[i] = b'r';
}
let target_str = String::from_utf8(target).unwrap();
let via_fast_path = levenshtein(&pattern, &target_str, &opts);
let via_plain_rows = oracle_plain_rows(&pattern, &target_str, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"degenerate huge-target mismatch"
);
}
#[test]
fn bit_vector_blocks_direct_call_empty_longer() {
let mut rng = Xorshift64(0xE3E4_1234_0000_AAAA);
for &len in &[65usize, 128, 129, 300] {
let shorter = random_units(&mut rng, len);
let longer: Vec<u8> = vec![];
let via_blocks = bit_vector_distance_blocks(&shorter, &longer);
let via_plain_rows = plain_rows(&shorter, &longer, &Options::default());
assert_eq!(
via_blocks, via_plain_rows,
"empty-longer mismatch len {len}"
);
assert_eq!(via_blocks, len as f64);
}
assert_eq!(
levenshtein("", &"m".repeat(12_000), &Options::default()),
12_000.0
);
}
#[test]
fn bit_vector_blocks_edit_distance_at_block_boundary() {
let opts = Options::default();
for &blocks_to_flip in &[1usize, 2, 3] {
let total_len = 400;
let mut longer = "a".repeat(total_len).into_bytes();
for b in 0..blocks_to_flip {
let start = 64 * b;
let end = (start + 64).min(total_len);
for byte in &mut longer[start..end] {
*byte = b'z';
}
}
let shorter = "a".repeat(total_len);
let longer_str = String::from_utf8(longer).unwrap();
let via_fast_path = levenshtein(&shorter, &longer_str, &opts);
let via_plain_rows = oracle_plain_rows(&shorter, &longer_str, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"block-boundary distance mismatch, blocks_to_flip={blocks_to_flip}"
);
assert_eq!(via_plain_rows, (blocks_to_flip * 64) as f64);
}
}
struct SplitMix64(u64);
impl SplitMix64 {
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_range(&mut self, bound: usize) -> usize {
(self.next_u64() % bound as u64) as usize
}
}
fn random_ascii_wide(rng: &mut SplitMix64, len: usize) -> String {
const ALPHABET: &[u8] = b"abcdefghijklmnopqrstuvwxyz0123456789";
(0..len)
.map(|_| ALPHABET[rng.next_range(ALPHABET.len())] as char)
.collect()
}
fn random_unicode_wide(rng: &mut SplitMix64, units: usize) -> String {
const BMP: &[char] = &['а', 'б', 'в', 'ñ', 'ü', '中', '字'];
const ASTRAL: &[char] = &['😀', '𝔘', '𝕏', '🎉'];
let mut s = String::new();
let mut remaining = units;
while remaining > 0 {
if remaining >= 2 && rng.next_range(3) == 0 {
s.push(ASTRAL[rng.next_range(ASTRAL.len())]);
remaining -= 2;
} else {
s.push(BMP[rng.next_range(BMP.len())]);
remaining -= 1;
}
}
s
}
#[test]
fn bit_vector_blocks_large_scale_differential_ascii_and_utf16() {
let mut rng = SplitMix64(0x243F_6A88_85A3_08D3);
let opts = Options::default();
let length_pairs = [
(65usize, 70usize),
(66, 4096),
(127, 130),
(200, 4096),
(383, 500),
(512, 10_001),
(1000, 1050),
(2049, 2100),
(65, 12_345),
(300, 10_007),
];
for &(shorter_len, longer_len) in &length_pairs {
let a = random_ascii_wide(&mut rng, shorter_len);
let b = random_ascii_wide(&mut rng, longer_len);
let via_fast_path = levenshtein(&a, &b, &opts);
let via_plain_rows = oracle_plain_rows(&a, &b, &opts);
assert_eq!(
via_fast_path, via_plain_rows,
"ascii mismatch shorter_len={shorter_len} longer_len={longer_len}"
);
let ua = random_unicode_wide(&mut rng, shorter_len);
let ub = random_unicode_wide(&mut rng, longer_len);
let via_fast_path_u = levenshtein(&ua, &ub, &opts);
let via_plain_rows_u = oracle_plain_rows(&ua, &ub, &opts);
assert_eq!(
via_fast_path_u, via_plain_rows_u,
"utf16 mismatch shorter_len={shorter_len} longer_len={longer_len}"
);
}
}
fn osa_opts() -> Options {
Options {
restricted: true,
..Options::default()
}
}
fn oracle_osa(a: &str, b: &str) -> f64 {
let opts = osa_opts();
dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => restricted_rows(s, t, &opts),
Operands::Units(s, t) => restricted_rows(s, t, &opts),
})
}
#[test]
fn osa_bit_vector_agrees_with_restricted_rows_on_random_pairs() {
let mut rng = Xorshift64(0x05A0_5A05_A05A);
let opts = osa_opts();
let lengths = [
0usize, 1, 2, 3, 7, 8, 9, 63, 64, 65, 127, 128, 129, 191, 192, 193, 256, 500,
];
for &a_len in &lengths {
for &b_len in &lengths {
for _ in 0..4 {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
let expected = oracle_osa(&a, &b);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
expected,
"mismatch for len {a_len} vs len {b_len}"
);
assert_eq!(
damerau_levenshtein(&b, &a, &opts),
expected,
"symmetry mismatch for len {b_len} vs len {a_len}"
);
}
}
}
}
#[test]
fn osa_transposition_heavy_inputs_agree() {
let mut rng = Xorshift64(0x7A57_A57A_57A5);
let opts = osa_opts();
for &len in &[10usize, 30, 64, 65, 100, 200, 300] {
for _ in 0..10 {
let a = random_string(&mut rng, len);
let mut b: Vec<char> = a.chars().collect();
let swaps = 1 + rng.next_range(len / 2);
for _ in 0..swaps {
let i = rng.next_range(len - 1);
b.swap(i, i + 1);
}
let b: String = b.into_iter().collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_osa(&a, &b),
"mismatch for len {len} after {swaps} adjacent swaps"
);
}
}
}
#[test]
fn osa_block_boundary_transpositions_agree() {
let mut rng = Xorshift64(0xB0B0_B0B0_B0B0);
let opts = osa_opts();
for &len in &[130usize, 200, 260] {
for &boundary in &[63usize, 127, 191] {
if boundary + 1 >= len {
continue;
}
for _ in 0..5 {
let a = random_string(&mut rng, len);
let mut b: Vec<char> = a.chars().collect();
b.swap(boundary, boundary + 1);
let b: String = b.into_iter().collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_osa(&a, &b),
"mismatch for len {len}, swap at ({boundary},{})",
boundary + 1
);
}
}
}
}
#[test]
fn osa_alternating_and_degenerate_inputs_agree() {
let opts = osa_opts();
for &len in &[63usize, 64, 65, 128, 129, 200] {
let ab: String = "ab".chars().cycle().take(len).collect();
let ba: String = "ba".chars().cycle().take(len).collect();
assert_eq!(damerau_levenshtein(&ab, &ba, &opts), oracle_osa(&ab, &ba));
let aa = "a".repeat(len);
let bb = "b".repeat(len);
assert_eq!(damerau_levenshtein(&aa, &bb, &opts), oracle_osa(&aa, &bb));
assert_eq!(damerau_levenshtein(&aa, &ab, &opts), oracle_osa(&aa, &ab));
}
assert_eq!(damerau_levenshtein("", "abc", &opts), 3.0);
assert_eq!(damerau_levenshtein("a", "abc", &opts), 2.0);
}
#[test]
fn osa_classic_fixtures() {
let opts = osa_opts();
assert_eq!(damerau_levenshtein("CA", "ABC", &opts), 3.0);
assert_eq!(damerau_levenshtein("CA", "AC", &opts), 1.0);
assert_eq!(damerau_levenshtein("ab", "ba", &opts), 1.0);
assert_eq!(damerau_levenshtein("kitten", "sitting", &opts), 3.0);
let filler = "a".repeat(64);
let s1 = format!("a{filler}CA{filler}a");
let s2 = format!("b{filler}AC{filler}b");
assert_eq!(damerau_levenshtein(&s1, &s2, &opts), 3.0);
assert_eq!(oracle_osa(&s1, &s2), 3.0);
}
#[test]
fn osa_single_word_and_blocks_agree_on_the_shared_domain() {
let mut rng = Xorshift64(0xC0DE_0511_0511);
for shorter_len in 2usize..=64 {
for _ in 0..6 {
let longer_len = rng.next_range(300).max(1);
let shorter = random_units(&mut rng, shorter_len);
let longer = random_units(&mut rng, longer_len);
assert_eq!(
osa_bit_vector(&shorter, &longer),
osa_bit_vector_blocks(&shorter, &longer),
"mismatch at shorter_len={shorter_len} longer_len={longer_len}"
);
}
}
}
#[test]
fn osa_utf16_and_astral_inputs_agree() {
let mut rng = Xorshift64(0x0111_0111_0111);
let opts = osa_opts();
const CYRILLIC: &[char] = &['\u{430}', '\u{431}', '\u{432}', '\u{433}', '\u{434}'];
for &(a_len, b_len) in &[
(10usize, 12usize),
(40, 40),
(64, 70),
(65, 130),
(200, 210),
] {
let a: String = (0..a_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let b: String = (0..b_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_osa(&a, &b),
"cyrillic mismatch {a_len}x{b_len}"
);
}
assert_eq!(
damerau_levenshtein(
"\u{418}\u{432}\u{430}\u{43d}\u{43a}\u{43e}",
"\u{41f}\u{435}\u{442}\u{440}\u{443}\u{43d}\u{43a}\u{43e}",
&opts
),
5.0
);
let a = "\u{1F600}".repeat(40);
let b = format!("a{}", "\u{1F600}".repeat(39));
assert_eq!(damerau_levenshtein(&a, &b, &opts), oracle_osa(&a, &b));
}
#[test]
fn osa_weighted_costs_never_take_the_fast_path() {
for transposition_cost in [0.5, 2.0, 0.0, f64::NAN] {
let opts = Options {
restricted: true,
transposition_cost,
..Options::default()
};
let got = damerau_levenshtein("abcd", "abdc", &opts);
let want = dispatch("abcd", "abdc", |ops| match ops {
Operands::Bytes(s, t) => restricted_rows(s, t, &opts),
Operands::Units(s, t) => restricted_rows(s, t, &opts),
});
assert_eq!(got.to_bits(), want.to_bits());
}
let weighted = Options {
restricted: true,
insertion_cost: 2.0,
..Options::default()
};
assert_eq!(damerau_levenshtein("ab", "abc", &weighted), 2.0);
assert_eq!(damerau_levenshtein("abc", "ab", &weighted), 1.0);
}
fn oracle_unrestricted(a: &str, b: &str) -> f64 {
let opts = Options::default();
dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => full_matrix(s, t, &opts, true, false).final_cost(),
Operands::Units(s, t) => full_matrix(s, t, &opts, true, false).final_cost(),
})
}
#[test]
fn damerau_unit_fast_path_matches_the_pinned_quirk_fixtures() {
let opts = Options::default();
assert_eq!(damerau_levenshtein("bb", "abbb", &opts), 1.0);
assert_eq!(damerau_levenshtein("abbb", "bb", &opts), 2.0);
assert_eq!(damerau_levenshtein("dfcb", "bdffc", &opts), 2.0);
assert_eq!(damerau_levenshtein("aabcbbb", "cabbccaab", &opts), 3.0);
assert_eq!(damerau_levenshtein("ca", "abc", &opts), 2.0);
for (a, b) in [
("bb", "abbb"),
("abbb", "bb"),
("dfcb", "bdffc"),
("aabcbbb", "cabbccaab"),
("ca", "abc"),
] {
assert_eq!(damerau_levenshtein(a, b, &opts), oracle_unrestricted(a, b));
}
}
#[test]
fn damerau_unit_fast_path_agrees_with_full_matrix_on_random_pairs() {
let mut rng = Xorshift64(0xDA3E_DA3E_DA3E);
let opts = Options::default();
let lengths = [0usize, 1, 2, 3, 5, 8, 13, 21, 34, 55, 80];
for &a_len in &lengths {
for &b_len in &lengths {
for _ in 0..6 {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"mismatch for {a:?} vs {b:?}"
);
}
}
}
for &(a_len, b_len) in &[(200usize, 210usize), (300, 40), (129, 500)] {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"mismatch at {a_len}x{b_len}"
);
}
}
#[test]
fn damerau_unit_fast_path_agrees_on_utf16_input() {
let mut rng = Xorshift64(0xDA3E_0016_0016);
let opts = Options::default();
const CYRILLIC: &[char] = &['\u{430}', '\u{431}', '\u{432}'];
for &(a_len, b_len) in &[(5usize, 7usize), (20, 20), (40, 60), (80, 30)] {
let a: String = (0..a_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let b: String = (0..b_len)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"cyrillic mismatch {a_len}x{b_len}"
);
}
let a = "\u{1F600}\u{1F601}\u{1F600}\u{1F601}";
let b = "\u{1F601}\u{1F600}\u{1F601}";
assert_eq!(damerau_levenshtein(a, b, &opts), oracle_unrestricted(a, b));
}
#[test]
fn damerau_weighted_costs_never_take_the_unit_fast_path() {
for opts in [
Options {
transposition_cost: 0.5,
..Options::default()
},
Options {
insertion_cost: 2.0,
..Options::default()
},
Options {
transposition_cost: f64::NAN,
..Options::default()
},
] {
let got = damerau_levenshtein("ca", "abc", &opts);
let want = dispatch("ca", "abc", |ops| match ops {
Operands::Bytes(s, t) => full_matrix(s, t, &opts, true, false).final_cost(),
Operands::Units(s, t) => full_matrix(s, t, &opts, true, false).final_cost(),
});
assert_eq!(got.to_bits(), want.to_bits());
}
let half = Options {
transposition_cost: 0.5,
..Options::default()
};
assert_eq!(damerau_levenshtein("ab", "ba", &half), 0.5);
}
fn sm_units(rng: &mut SplitMix64, len: usize) -> Vec<u8> {
const ALPHABET: &[u8] = b"abcde";
(0..len)
.map(|_| ALPHABET[rng.next_range(ALPHABET.len())])
.collect()
}
fn ascii_string(units: &[u8]) -> String {
String::from_utf8(units.to_vec()).expect("ascii")
}
#[test]
fn osa_offset_straddle_transpositions_agree() {
let mut rng = SplitMix64(0x05A0_0FF5_E7B0_0001);
let opts = osa_opts();
for &len in &[130usize, 200, 300] {
for &p in &[0usize, 1, 62, 63, 64, 65, 126, 127, 128, 129] {
if p + 1 >= len {
continue;
}
for prefix_len in 0usize..=2 {
let a = sm_units(&mut rng, len);
let mut b = sm_units(&mut rng, prefix_len);
b.extend_from_slice(&a);
b.swap(prefix_len + p, prefix_len + p + 1);
assert_eq!(
osa_bit_vector_blocks(&a, &b),
restricted_rows(&a, &b, &opts),
"direct blocks mismatch len={len} p={p} prefix={prefix_len}"
);
let sa = ascii_string(&a);
let sb = ascii_string(&b);
let expected = oracle_osa(&sa, &sb);
assert_eq!(
damerau_levenshtein(&sa, &sb, &opts),
expected,
"public mismatch len={len} p={p} prefix={prefix_len}"
);
assert_eq!(
damerau_levenshtein(&sb, &sa, &opts),
expected,
"public reversed mismatch len={len} p={p} prefix={prefix_len}"
);
}
}
let a = sm_units(&mut rng, len);
let mut b = a.clone();
b.swap(len - 2, len - 1);
assert_eq!(
osa_bit_vector_blocks(&a, &b),
restricted_rows(&a, &b, &opts),
"tail swap mismatch len={len}"
);
}
}
#[test]
fn osa_single_symbol_seas_with_boundary_swaps() {
let opts = osa_opts();
for &len in &[65usize, 129, 193, 260] {
for &p in &[0usize, 62, 63, 64, 127, 128, 191, 192] {
if p + 1 >= len {
continue;
}
let mut s1 = vec![b'a'; len];
let mut s2 = vec![b'a'; len];
s1[p] = b'b';
s2[p + 1] = b'b';
assert_eq!(
osa_bit_vector_blocks(&s1, &s2),
restricted_rows(&s1, &s2, &opts),
"sea swap mismatch len={len} p={p}"
);
let mut s3 = vec![b'a'; len];
let mut s4 = vec![b'a'; len];
s3[p] = b'b';
s3[p + 1] = b'c';
s4[p] = b'c';
s4[p + 1] = b'b';
let sa = ascii_string(&s3);
let sb = ascii_string(&s4);
let expected = oracle_osa(&sa, &sb);
assert_eq!(
damerau_levenshtein(&sa, &sb, &opts),
expected,
"bc-sea mismatch len={len} p={p}"
);
assert_eq!(expected, 1.0, "a bc<->cb swap must cost exactly 1");
}
let mut s1 = vec![b'a'; len];
let mut s2 = vec![b'a'; len];
s1[len - 2] = b'b';
s2[len - 1] = b'b';
assert_eq!(
osa_bit_vector_blocks(&s1, &s2),
restricted_rows(&s1, &s2, &opts),
"sea tail swap mismatch len={len}"
);
}
}
#[test]
fn osa_tiny_and_empty_operands_all_entries() {
let opts = osa_opts();
let tiny = ["", "a", "b", "ab", "ba", "aa", "abc", "cba", "aab"];
for a in tiny {
for b in tiny {
assert_eq!(
damerau_levenshtein(a, b, &opts),
oracle_osa(a, b),
"tiny mismatch {a:?} vs {b:?}"
);
}
}
let long = vec![b'q'; 65];
assert_eq!(osa_bit_vector_blocks(&long, b""), 65.0);
assert_eq!(
osa_bit_vector_blocks(&long, b"q"),
restricted_rows(&long, b"q", &opts)
);
assert_eq!(
osa_bit_vector_blocks(b"x", b"xy"),
restricted_rows(b"x", b"xy", &opts)
);
assert_eq!(
osa_bit_vector(b"xy", b"yx"),
restricted_rows(b"xy", b"yx", &opts)
);
}
#[test]
fn osa_large_randomized_differential_splitmix() {
let mut rng = SplitMix64(0x05A0_2026_0816_AAAA);
let opts = osa_opts();
for round in 0..120 {
let len = 65 + rng.next_range(350);
let a = sm_units(&mut rng, len);
let mut b = a.clone();
let swaps = 1 + rng.next_range(10);
for _ in 0..swaps {
let i = rng.next_range(b.len() - 1);
b.swap(i, i + 1);
}
for _ in 0..rng.next_range(5) {
let i = rng.next_range(b.len());
b[i] = b"abcde"[rng.next_range(5)];
}
if rng.next_range(3) == 0 {
let cut = 1 + rng.next_range(4);
b.drain(..cut);
}
let sa = ascii_string(&a);
let sb = ascii_string(&b);
let expected = oracle_osa(&sa, &sb);
assert_eq!(
damerau_levenshtein(&sa, &sb, &opts),
expected,
"round {round} ({len})"
);
assert_eq!(
damerau_levenshtein(&sb, &sa, &opts),
expected,
"round {round} reversed ({len})"
);
}
const BMP: &[char] = &['\u{430}', '\u{431}', '\u{432}', '\u{4E2D}'];
for &p in &[62usize, 63, 64, 65, 127, 128] {
let chars: Vec<char> = (0..160).map(|_| BMP[rng.next_range(BMP.len())]).collect();
let mut swapped = chars.clone();
swapped.swap(p, p + 1);
let a: String = chars.into_iter().collect();
let b: String = swapped.into_iter().collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_osa(&a, &b),
"utf16 swap at {p}"
);
}
let a = "\u{1F600}\u{1F601}".repeat(40); let b = format!("\u{1F601}\u{1F600}{}", "\u{1F600}\u{1F601}".repeat(39));
assert_eq!(damerau_levenshtein(&a, &b, &opts), oracle_osa(&a, &b));
}
#[test]
fn bitpeq_full_byte_alphabet_kernels_direct() {
let mut rng = SplitMix64(0xB17E_0000_FFFF_0001);
let opts = Options::default();
for &m in &[8usize, 64, 65, 200, 256, 300] {
for _ in 0..3 {
let shorter: Vec<u8> = (0..m)
.map(|i| {
if i % 2 == 0 {
(i % 256) as u8
} else {
(rng.next_u64() % 256) as u8
}
})
.collect();
let longer: Vec<u8> = (0..600).map(|_| (rng.next_u64() % 256) as u8).collect();
let expected = plain_rows(&shorter, &longer, &opts);
if (8..=64).contains(&m) {
assert_eq!(
bit_vector_distance(&shorter, &longer),
expected,
"single-word full-alphabet mismatch m={m}"
);
}
assert_eq!(
bit_vector_distance_blocks(&shorter, &longer),
expected,
"blocks full-alphabet mismatch m={m}"
);
}
}
}
#[test]
fn bitpeq_u16_wide_alphabet_kernels_direct() {
let mut rng = SplitMix64(0x0016_31DE_A1FA_0001);
let opts = Options::default();
for &m in &[65usize, 200, 400] {
let shorter: Vec<u16> = (0..m)
.map(|i| {
if i % 3 == 0 {
(i % 1000) as u16
} else {
(rng.next_u64() % 1000) as u16
}
})
.collect();
let longer: Vec<u16> = (0..700).map(|_| (rng.next_u64() % 1000) as u16).collect();
assert_eq!(
bit_vector_distance_blocks(&shorter, &longer),
plain_rows(&shorter, &longer, &opts),
"u16 wide-alphabet blocks mismatch m={m}"
);
}
let shorter: Vec<u16> = (0..60).map(|_| (rng.next_u64() % 1000) as u16).collect();
let longer: Vec<u16> = (0..500).map(|_| (rng.next_u64() % 1000) as u16).collect();
assert_eq!(
bit_vector_distance(&shorter, &longer),
plain_rows(&shorter, &longer, &opts)
);
}
#[test]
fn bitpeq_randomized_differential_splitmix() {
let mut rng = SplitMix64(0xB17E_2026_0816_BBBB);
let opts = Options::default();
let lengths = [8usize, 63, 64, 65, 66, 127, 128, 129, 130, 192, 250];
for &m in &lengths {
for _ in 0..4 {
let narrow = rng.next_range(2) == 0;
let gen_byte = |rng: &mut SplitMix64| -> u8 {
if narrow {
b'a' + (rng.next_u64() % 3) as u8
} else {
(rng.next_u64() % 256) as u8
}
};
let shorter: Vec<u8> = (0..m).map(|_| gen_byte(&mut rng)).collect();
let longer_len = 1 + rng.next_range(400);
let longer: Vec<u8> = (0..longer_len).map(|_| gen_byte(&mut rng)).collect();
let expected = plain_rows(&shorter, &longer, &opts);
if (8..=64).contains(&m) {
assert_eq!(
bit_vector_distance(&shorter, &longer),
expected,
"word mismatch m={m} n={longer_len} narrow={narrow}"
);
}
assert_eq!(
bit_vector_distance_blocks(&shorter, &longer),
expected,
"blocks mismatch m={m} n={longer_len} narrow={narrow}"
);
}
}
}
#[test]
fn damerau_unit_snapshot_overwrite_stress() {
let opts = Options::default();
for &k in &[2usize, 3, 5, 8, 40, 100] {
let pairs = [
("ab".repeat(k), "ba".repeat(k)),
("aab".repeat(k), "aba".repeat(k)),
("abc".repeat(k), "cab".repeat(k)),
("ab".repeat(k), format!("b{}", "ab".repeat(k))),
("aabb".repeat(k), "bbaa".repeat(k)),
];
for (a, b) in &pairs {
let expected = oracle_unrestricted(a, b);
assert_eq!(
damerau_levenshtein(a, b, &opts),
expected,
"structured mismatch k={k} {a:?} vs {b:?}"
);
let expected_rev = oracle_unrestricted(b, a);
assert_eq!(
damerau_levenshtein(b, a, &opts),
expected_rev,
"structured reversed mismatch k={k}"
);
}
}
let mut rng = SplitMix64(0xDA3E_2026_0816_CCCC);
for round in 0..200 {
let len_a = 1 + rng.next_range(120);
let len_b = 1 + rng.next_range(120);
let a: String = (0..len_a)
.map(|_| if rng.next_range(2) == 0 { 'a' } else { 'b' })
.collect();
let b: String = (0..len_b)
.map(|_| if rng.next_range(2) == 0 { 'a' } else { 'b' })
.collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"ab-random mismatch round {round} ({len_a}x{len_b})"
);
}
}
#[test]
fn damerau_unit_degenerate_tiny_and_nul() {
let opts = Options::default();
let tiny = ["", "a", "b", "ab", "ba", "aa", "aba", "bab"];
for a in tiny {
for b in tiny {
assert_eq!(
damerau_levenshtein(a, b, &opts),
oracle_unrestricted(a, b),
"tiny mismatch {a:?} vs {b:?}"
);
}
}
for (a, b) in [
("a".repeat(500), "a".repeat(497)),
("a".repeat(300), "b".repeat(300)),
("ab".repeat(150), "ba".repeat(150)),
("a".repeat(400), format!("{}b", "a".repeat(399))),
] {
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"degenerate mismatch {}x{}",
a.len(),
b.len()
);
}
for (a, b) in [
("\0ab\0", "b\0a"),
("\0\0", "\0"),
("a\0b", "ab\0"),
("\0a", "a\0"),
] {
assert_eq!(
damerau_levenshtein(a, b, &opts),
oracle_unrestricted(a, b),
"nul mismatch {a:?} vs {b:?}"
);
}
}
#[test]
fn damerau_unit_many_distinct_symbols() {
let opts = Options::default();
let mut rng = SplitMix64(0xDA3E_A1FA_BE7A_0001);
let source: Vec<u8> = (0..300).map(|i| (i % 256) as u8).collect();
let target: Vec<u8> = (0..310).map(|_| (rng.next_u64() % 256) as u8).collect();
assert_eq!(
damerau_unrestricted_unit::<_, u16>(&source, &target),
full_matrix(&source, &target, &opts, true, false).final_cost()
);
assert_eq!(
damerau_unrestricted_unit::<_, u16>(&target, &source),
full_matrix(&target, &source, &opts, true, false).final_cost()
);
let wide_char = |i: usize| char::from_u32(0x400 + (i % 400) as u32).unwrap();
let a: String = (0..350).map(wide_char).collect();
let b: String = (0..350).map(|i| wide_char(i + 7)).collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"u16 wide-alphabet mismatch"
);
let a = "\u{1F600}\u{1F601}\u{1F602}".repeat(30);
let b = format!(
"\u{1F601}\u{1F600}{}",
"\u{1F602}\u{1F601}\u{1F600}".repeat(29)
);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"astral mismatch"
);
}
#[test]
fn damerau_unit_large_randomized_differential_splitmix() {
let mut rng = SplitMix64(0xDA3E_2026_0816_DDDD);
let opts = Options::default();
const ALPHABETS: [&[u8]; 2] = [b"ab", b"abcde"];
for round in 0..80 {
let alphabet = ALPHABETS[round % 2];
let len_a = 1 + rng.next_range(250);
let len_b = 1 + rng.next_range(250);
let a: String = (0..len_a)
.map(|_| alphabet[rng.next_range(alphabet.len())] as char)
.collect();
let b: String = (0..len_b)
.map(|_| alphabet[rng.next_range(alphabet.len())] as char)
.collect();
let fwd = oracle_unrestricted(&a, &b);
let rev = oracle_unrestricted(&b, &a);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
fwd,
"round {round} ({len_a}x{len_b})"
);
assert_eq!(
damerau_levenshtein(&b, &a, &opts),
rev,
"round {round} reversed ({len_b}x{len_a})"
);
}
const CYR: &[char] = &['\u{430}', '\u{431}', '\u{432}'];
for round in 0..15 {
let len_a = 1 + rng.next_range(150);
let len_b = 1 + rng.next_range(150);
let a: String = (0..len_a).map(|_| CYR[rng.next_range(CYR.len())]).collect();
let b: String = (0..len_b).map(|_| CYR[rng.next_range(CYR.len())]).collect();
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
oracle_unrestricted(&a, &b),
"u16 round {round}"
);
}
}
#[test]
fn damerau_unit_u16_and_u32_cells_agree() {
let mut rng = Xorshift64(0xCE11_CE11_CE11);
for &(a_len, b_len) in &[(0usize, 5usize), (7, 7), (40, 60), (128, 200), (300, 41)] {
for _ in 0..8 {
let a = random_units(&mut rng, a_len);
let b = random_units(&mut rng, b_len);
assert_eq!(
damerau_unrestricted_unit::<u8, u16>(&a, &b),
damerau_unrestricted_unit::<u8, u32>(&a, &b),
"cell-width mismatch at {a_len}x{b_len}"
);
}
}
}
#[test]
fn damerau_byte_tiers_agree_with_full_matrix_at_boundaries() {
let mut rng = Xorshift64(0x71E5_71E5_71E5);
let opts = Options::default();
let sizes = [
(7usize, 7usize),
(8, 8),
(8, 9),
(9, 8),
(9, 9),
(8, 200),
(127, 128),
(128, 128),
(128, 129),
(129, 129),
(129, 40),
(160, 160),
(200, 130),
];
for &(a_len, b_len) in &sizes {
for _ in 0..6 {
let a = random_string(&mut rng, a_len);
let b = random_string(&mut rng, b_len);
let expected = oracle_unrestricted(&a, &b);
assert_eq!(
damerau_levenshtein(&a, &b, &opts),
expected,
"mismatch at {a_len}x{b_len}"
);
assert_eq!(
damerau_levenshtein(&b, &a, &opts),
oracle_unrestricted(&b, &a),
"reverse mismatch at {b_len}x{a_len}"
);
}
}
}
#[test]
fn damerau_byte_tiers_agree_with_each_other_on_shared_domains() {
let mut rng = Xorshift64(0x3B1D_3B1D_3B1D);
for _ in 0..300 {
let a_len = 1 + rng.next_range(8);
let b_len = 1 + rng.next_range(8);
let a = random_units(&mut rng, a_len);
let b = random_units(&mut rng, b_len);
let small = damerau_unit_small(&a, &b);
let mid = damerau_unit_mid(&a, &b);
let large = damerau_unit_large(&a, &b);
let generic = damerau_unrestricted_unit::<u8, u16>(&a, &b);
assert_eq!(small, mid, "small/mid at {a_len}x{b_len}");
assert_eq!(mid, large, "mid/large at {a_len}x{b_len}");
assert_eq!(large, generic, "large/generic at {a_len}x{b_len}");
}
for _ in 0..100 {
let a_len = 9 + rng.next_range(120);
let b_len = 9 + rng.next_range(120);
let a = random_units(&mut rng, a_len);
let b = random_units(&mut rng, b_len);
assert_eq!(
damerau_unit_mid(&a, &b),
damerau_unit_large(&a, &b),
"mid/large at {a_len}x{b_len}"
);
assert_eq!(
damerau_unit_large(&a, &b),
damerau_unrestricted_unit::<u8, u16>(&a, &b),
"large/generic at {a_len}x{b_len}"
);
}
}
#[test]
fn damerau_byte_tiers_handle_quirks_and_degenerate_shapes() {
for (a, b, want) in [
("bb", "abbb", 1.0),
("abbb", "bb", 2.0),
("dfcb", "bdffc", 2.0),
("aabcbbb", "cabbccaab", 3.0),
("ca", "abc", 2.0),
] {
let ab = a.as_bytes();
let bb = b.as_bytes();
if ab.len() <= 8 && bb.len() <= 8 {
assert_eq!(damerau_unit_small(ab, bb), want, "small {a:?}");
}
assert_eq!(damerau_unit_mid(ab, bb), want, "mid {a:?}");
assert_eq!(damerau_unit_large(ab, bb), want, "large {a:?}");
}
let opts = Options::default();
for len in [8usize, 9, 60, 129, 200] {
let aa = "a".repeat(len);
let ab: String = "ab".chars().cycle().take(len).collect();
let zz = "z".repeat(len + 3);
assert_eq!(
damerau_levenshtein(&aa, &ab, &opts),
oracle_unrestricted(&aa, &ab)
);
assert_eq!(
damerau_levenshtein(&aa, &zz, &opts),
oracle_unrestricted(&aa, &zz)
);
}
}
fn oracle_search(a: &str, b: &str, opts: &Options, damerau: bool) -> SearchResult {
dispatch(a, b, |ops| match ops {
Operands::Bytes(s, t) => {
let (start, end, dist) = search_full_matrix(s, t, opts, damerau);
SearchResult {
substring: String::from_utf8_lossy(slice_units(t, start, end)).into_owned(),
distance: dist,
offset: start,
}
}
Operands::Units(s, t) => {
let (start, end, dist) = search_full_matrix(s, t, opts, damerau);
SearchResult {
substring: String::from_utf16_lossy(slice_units(t, start, end)),
distance: dist,
offset: start,
}
}
})
}
fn search_rand(rng: &mut SplitMix64, len: usize, alphabet: usize) -> String {
(0..len)
.map(|_| (b'a' + rng.next_range(alphabet) as u8) as char)
.collect()
}
fn embed_near_match(
rng: &mut SplitMix64,
needle: &str,
haystack: &mut String,
alphabet: usize,
) {
let n = needle.len();
let m = haystack.len();
if m <= n {
return;
}
let pos = rng.next_range(m - n);
let mut copy = needle.to_owned().into_bytes();
for _ in 0..rng.next_range(3) {
let i = rng.next_range(copy.len());
copy[i] = b'a' + rng.next_range(alphabet) as u8;
}
haystack.replace_range(pos..pos + n, std::str::from_utf8(©).unwrap());
}
#[test]
fn search_bits_agrees_with_full_matrix_on_random_ascii() {
let mut rng = SplitMix64(0x5EA2_C4B1_D00D_0001);
let opts = Options::default();
for case in 0..3000usize {
let alphabet = [2usize, 3, 4, 26][rng.next_range(4)];
let n = 1 + rng.next_range(if case % 5 == 0 { 200 } else { 90 });
let m = 1 + rng.next_range(220);
let s = search_rand(&mut rng, n, alphabet);
let mut t = search_rand(&mut rng, m, alphabet);
if rng.next_range(2) == 0 {
embed_near_match(&mut rng, &s, &mut t, alphabet);
}
let got = levenshtein_search(&s, &t, &opts);
let want = oracle_search(&s, &t, &opts, false);
assert_eq!(got, want, "search mismatch: s={s:?} t={t:?}");
}
}
#[test]
fn search_bits_boundary_needle_lengths_agree() {
let mut rng = SplitMix64(0x5EA2_C4B1_D00D_0002);
let opts = Options::default();
for &n in &[1usize, 2, 63, 64, 65, 66, 127, 128, 129, 130] {
for &m in &[1usize, 64, 65, 129, 200] {
for _ in 0..6 {
let s = search_rand(&mut rng, n, 3);
let mut t = search_rand(&mut rng, m, 3);
embed_near_match(&mut rng, &s, &mut t, 3);
let got = levenshtein_search(&s, &t, &opts);
let want = oracle_search(&s, &t, &opts, false);
assert_eq!(got, want, "boundary mismatch n={n} m={m}");
}
}
}
}
#[test]
fn search_bits_agrees_on_utf16_input() {
let mut rng = SplitMix64(0x5EA2_C4B1_D00D_0003);
let opts = Options::default();
const CYRILLIC: &[char] = &['а', 'б', 'в', 'г'];
for _ in 0..600 {
let n = 1 + rng.next_range(120);
let m = 1 + rng.next_range(140);
let s: String = (0..n)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let t: String = (0..m)
.map(|_| CYRILLIC[rng.next_range(CYRILLIC.len())])
.collect();
let got = levenshtein_search(&s, &t, &opts);
let want = oracle_search(&s, &t, &opts, false);
assert_eq!(got, want, "u16 search mismatch: s={s:?} t={t:?}");
}
for _ in 0..200 {
let s_units = 2 + rng.next_range(100);
let t_units = 2 + rng.next_range(140);
let s = random_unicode_wide(&mut rng, s_units);
let t = random_unicode_wide(&mut rng, t_units);
let got = levenshtein_search(&s, &t, &opts);
let want = oracle_search(&s, &t, &opts, false);
assert_eq!(got, want, "astral search mismatch: s={s:?} t={t:?}");
}
}
#[test]
fn search_cell_costs_match_the_full_matrix() {
let mut rng = SplitMix64(0x5EA2_C4B1_D00D_0004);
let opts = Options::default();
for _ in 0..150 {
let n = 1 + rng.next_range(150);
let m = 1 + rng.next_range(150);
let s = search_rand(&mut rng, n, 3).into_bytes();
let t = search_rand(&mut rng, m, 3).into_bytes();
let mat = full_matrix(&s, &t, &opts, false, true);
let fw = if n <= 64 {
search_forward_word(&s, &t)
} else {
search_forward_blocks(&s, &t)
};
if n <= 64 {
for r in 0..=n {
for c in 0..=m {
assert_eq!(
search_cell_cost(&fw, r, c) as f64,
mat.cost_at(r, c),
"cell ({r},{c}) n={n} m={m}"
);
}
}
} else {
for _ in 0..60 {
let r = rng.next_range(n + 1);
let c = rng.next_range(m + 1);
assert_eq!(
search_cell_cost(&fw, r, c) as f64,
mat.cost_at(r, c),
"cell ({r},{c}) n={n} m={m}"
);
}
for r in [64usize, 65, n] {
for c in [1usize, m / 2, m] {
assert_eq!(
search_cell_cost(&fw, r, c) as f64,
mat.cost_at(r, c),
"word-boundary cell ({r},{c}) n={n} m={m}"
);
}
}
}
}
}
#[test]
fn search_word_and_blocks_agree_on_the_shared_domain() {
let mut rng = SplitMix64(0x5EA2_C4B1_D00D_0005);
for &n in &[1usize, 7, 32, 63, 64] {
for _ in 0..8 {
let m = 1 + rng.next_range(120);
let s = search_rand(&mut rng, n, 3).into_bytes();
let t = search_rand(&mut rng, m, 3).into_bytes();
let word = search_forward_word(&s, &t);
let blocks = search_forward_blocks(&s, &t);
assert_eq!(word.match_end, blocks.match_end, "match_end n={n} m={m}");
assert_eq!(
word.min_distance, blocks.min_distance,
"min_distance n={n} m={m}"
);
for r in 0..=n {
for c in 0..=m {
assert_eq!(
search_cell_cost(&word, r, c),
search_cell_cost(&blocks, r, c),
"cell ({r},{c}) n={n} m={m}"
);
}
}
}
}
}
#[test]
fn search_tie_breaking_pinned_examples() {
let opts = Options::default();
for (s, t) in [
("aaa", "aaaaaa"),
("aa", "aa"),
("ab", "ababab"),
("aba", "bab"),
("ca", "abc"),
("b", "aaa"),
] {
let got = levenshtein_search(s, t, &opts);
let want = oracle_search(s, t, &opts, false);
assert_eq!(got, want, "tie-break mismatch for {s:?} in {t:?}");
}
let r = levenshtein_search("aaa", "aaaaaa", &opts);
assert_eq!(
(r.substring.as_str(), r.distance, r.offset),
("aaa", 0.0, 0)
);
}
#[test]
fn search_weighted_damerau_and_empty_operands_keep_the_matrix_path() {
let weighted = Options {
substitution_cost: 0.5,
..Options::default()
};
let got = levenshtein_search("kitten", "sitting", &weighted);
let want = oracle_search("kitten", "sitting", &weighted, false);
assert_eq!(got, want);
let opts = Options::default();
for (s, t) in [("ca", "abc"), ("ab", "xxbaxx"), ("abcd", "acbd")] {
let got = damerau_levenshtein_search(s, t, &opts);
let want = oracle_search(s, t, &opts, true);
assert_eq!(got, want, "damerau search mismatch for {s:?} in {t:?}");
}
for (s, t) in [("", "abc"), ("abc", ""), ("", "")] {
let got = levenshtein_search(s, t, &opts);
let want = oracle_search(s, t, &opts, false);
assert_eq!(got, want, "empty-operand mismatch for {s:?} in {t:?}");
}
assert_eq!(levenshtein_search("", "abc", &opts).distance, 0.0);
assert_eq!(levenshtein_search("abc", "", &opts).distance, 3.0);
}
#[test]
fn search_bench_corpus_pairs_agree() {
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.ancestors()
.nth(2)
.unwrap()
.join("benches/data/distance-pairs.json");
let Ok(body) = std::fs::read_to_string(&path) else {
eprintln!("skipping: {} not generated", path.display());
return;
};
let json: serde_json::Value = serde_json::from_str(&body).expect("valid bench data");
let opts = Options::default();
for key in ["ascii", "cyrillic"] {
for (size, pair) in json["pairs"][key].as_object().expect("pair map") {
let a = pair[0].as_str().unwrap();
let b = pair[1].as_str().unwrap();
let got = levenshtein_search(a, b, &opts);
let want = oracle_search(a, b, &opts, false);
assert_eq!(got, want, "bench pair {key}/{size}");
}
}
}
}