use crate::normalize;
use std::collections::HashMap;
pub fn trigram_similarity(a: &str, b: &str) -> f64 {
let a = normalize(a);
let b = normalize(b);
trigram_similarity_raw(&a, &b)
}
pub(crate) fn trigram_similarity_raw(a: &str, b: &str) -> f64 {
let a_tri = trigrams(a);
let b_tri = trigrams(b);
if a_tri.is_empty() && b_tri.is_empty() {
return 1.0;
}
if a_tri.is_empty() || b_tri.is_empty() {
return 0.0;
}
let total = a_tri.len() + b_tri.len();
let mut counts: HashMap<String, i32> = HashMap::new();
for t in &a_tri {
*counts.entry(t.clone()).or_insert(0) += 1;
}
let mut shared = 0;
for t in &b_tri {
if let Some(c) = counts.get_mut(t) {
if *c > 0 {
shared += 1;
*c -= 1;
}
}
}
2.0 * shared as f64 / total as f64
}
fn trigrams(s: &str) -> Vec<String> {
let chars: Vec<char> = s.chars().collect();
if chars.len() < 3 {
return vec![];
}
chars.windows(3).map(|w| w.iter().collect()).collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f64, b: f64) -> bool {
(a - b).abs() < 0.001
}
#[test]
fn test_identical() {
assert!(approx(trigram_similarity("hello", "hello"), 1.0));
}
#[test]
fn test_both_empty() {
assert!(approx(trigram_similarity("", ""), 1.0));
}
#[test]
fn test_no_overlap() {
assert!(approx(trigram_similarity("abc", "xyz"), 0.0));
}
#[test]
fn test_partial() {
let score = trigram_similarity("hello", "helo");
assert!(score > 0.0 && score < 1.0);
}
#[test]
fn test_case_insensitive() {
assert!(approx(trigram_similarity("Hello", "hello"), 1.0));
}
#[test]
fn test_raw_skips_normalization() {
assert!(trigram_similarity_raw("Hello", "hello") < 1.0);
assert!(approx(trigram_similarity_raw("hello", "hello"), 1.0));
}
#[test]
fn test_raw_matches_public_on_normalized_input() {
assert!(approx(
trigram_similarity_raw("hello", "helo"),
trigram_similarity("hello", "helo"),
));
}
}