#![allow(dead_code)]
use std::cell::RefCell; use std::cmp::min; use std::collections::HashMap; use std::hash::Hash; use std::rc::{Rc, Weak};
#[derive(Default, Debug)]
struct TrieNode<T: Default + PartialEq> {
children: HashMap<char, Rc<RefCell<TrieNode<T>>>>,
word: Option<String>,
data: Vec<T>,
is_end: bool,
parent: Option<Weak<RefCell<TrieNode<T>>>>,
node_char: char,
}
pub struct Trie<T: Clone + Default + PartialEq + Eq + Hash> {
root: Rc<RefCell<TrieNode<T>>>,
data_map: HashMap<T, Vec<Weak<RefCell<TrieNode<T>>>>>,
}
#[derive(Debug)]
pub struct SearchResult<T> {
pub word: String,
pub data: Vec<T>,
}
#[derive(Debug)]
pub struct SearchResultWithScore<T> {
pub word: String,
pub data: Vec<T>,
pub score: f32,
}
impl<T: PartialEq> PartialEq for SearchResult<T> {
fn eq(&self, other: &Self) -> bool {
self.word == other.word && self.data == other.data
}
}
impl<T: Clone + Default + PartialEq + Eq + Hash> Trie<T> {
pub fn new() -> Self {
Trie {
root: Rc::new(RefCell::new(TrieNode {
node_char: '$', ..Default::default()
})),
data_map: HashMap::new(),
}
}
pub fn insert(&mut self, word: &str, data: T) {
let mut current = Rc::clone(&self.root);
let augmented_word = format!("${}", word);
for c in augmented_word.chars() {
let next = {
let mut current_ref = current.borrow_mut();
current_ref
.children
.entry(c)
.or_insert_with(|| {
Rc::new(RefCell::new(TrieNode {
parent: Some(Rc::downgrade(¤t)),
node_char: c,
..Default::default()
}))
})
.clone()
};
current = next;
}
let mut current_ref = current.borrow_mut();
if current_ref.word.is_none() {
current_ref.word = Some(word.to_string());
}
current_ref.data.push(data.clone());
current_ref.is_end = true;
self.data_map
.entry(data)
.or_default()
.push(Rc::downgrade(¤t));
}
pub fn search_within_distance(&self, word: &str, max_distance: usize) -> Vec<SearchResult<T>> {
let augmented_word = format!("${}", word);
let augmented_word_length = augmented_word.len();
let mut rows = vec![vec![0; augmented_word_length + 1]];
for i in 0..=augmented_word_length {
rows[0][i] = i;
}
let mut results = Vec::new();
self.search_impl(
&self.root.borrow(),
'$',
&mut rows,
&augmented_word,
max_distance,
&mut results,
true,
);
results
}
pub fn search_within_distance_scored(
&self,
word: &str,
max_distance: usize,
) -> Vec<SearchResultWithScore<T>> {
self.search_within_distance(word, max_distance)
.into_iter()
.map(|result| {
let score = self.calculate_jaro_winkler_score(word, &result.word);
SearchResultWithScore {
word: result.word.clone(), data: result.data,
score,
}
})
.collect()
}
fn search_impl(
&self,
node: &TrieNode<T>,
ch: char,
rows: &mut Vec<Vec<usize>>,
word: &str,
max_distance: usize,
results: &mut Vec<SearchResult<T>>,
is_root: bool,
) {
let row_length = word.len() + 1;
let mut current_row = vec![0; row_length];
current_row[0] = if is_root {
0
} else {
rows.last().unwrap()[0] + 1
};
for i in 1..row_length {
let insert_or_del = min(current_row[i - 1] + 1, rows.last().unwrap()[i] + 1);
let replace = if word.chars().nth(i - 1) == Some(ch) {
rows.last().unwrap()[i - 1] } else {
rows.last().unwrap()[i - 1] + 1 };
current_row[i] = min(insert_or_del, replace);
}
let should_search_childs = *current_row.iter().min().unwrap() <= max_distance;
if node.word.is_some() {
if current_row[row_length - 1] <= max_distance {
collect_all_words_from_this_node(node, results);
return;
}
}
rows.push(current_row);
if should_search_childs {
for (next_ch, child) in &node.children {
self.search_impl(
&child.borrow(),
*next_ch,
rows,
word,
max_distance,
results,
false,
);
}
}
else if rows.len() > max_distance && rows.len() - max_distance >= word.len() + 1 {
if rows.len() > word.len() {
for i in word.len() - max_distance + 1..word.len() + 2 {
if rows[i].last().unwrap() <= &max_distance {
collect_all_words_from_this_node(node, results);
rows.pop();
return;
}
}
for i in word.len() + 2..word.len() + 2 + max_distance + 1 {
if rows.len() > i {
if rows[i].last().unwrap() <= &max_distance {
collect_all_words_from_this_node(node, results);
rows.pop();
return;
}
}
}
}
}
rows.pop();
}
pub fn remove_all(&mut self, data: &T) {
if let Some(nodes) = self.data_map.get_mut(data) {
let mut empty_nodes = Vec::new();
nodes.retain(|node_weak| {
if let Some(node) = node_weak.upgrade() {
let mut node_ref = node.borrow_mut();
node_ref.data.retain(|d| d != data);
if node_ref.data.is_empty() && node_ref.word.is_some() {
node_ref.word = None;
node_ref.is_end = false;
}
if node_ref.data.is_empty() {
empty_nodes.push(Rc::clone(&node));
}
!node_ref.data.is_empty()
} else {
false }
});
for node in empty_nodes {
self.remove_node(node);
}
}
self.data_map.remove(data);
}
fn remove_node(&mut self, node: Rc<RefCell<TrieNode<T>>>) {
let mut current = node;
loop {
let parent = {
let current_ref = current.borrow();
if !current_ref.children.is_empty()
|| current_ref.word.is_some()
|| !current_ref.data.is_empty()
{
break;
}
current_ref.parent.as_ref().and_then(Weak::upgrade)
};
if let Some(parent_node) = parent {
let node_char = current.borrow().node_char;
parent_node.borrow_mut().children.remove(&node_char);
current = parent_node;
} else {
break;
}
}
}
}
fn collect_all_words_from_this_node<T: Clone + Default + PartialEq>(
node: &TrieNode<T>,
results: &mut Vec<SearchResult<T>>,
) {
if let Some(ref node_word) = node.word {
results.push(SearchResult {
word: node_word.clone(),
data: node.data.clone(),
});
}
for (_, child) in &node.children {
collect_all_words_from_this_node(&child.borrow(), results);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_insert_and_search() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("app", 2);
trie.insert("application", 3);
let results = trie.search_within_distance("app", 0);
assert_eq!(results.len(), 3);
assert!(results.contains(&SearchResult {
word: "app".to_string(),
data: vec![2]
}));
assert!(results.contains(&SearchResult {
word: "apple".to_string(),
data: vec![1]
}));
assert!(results.contains(&SearchResult {
word: "application".to_string(),
data: vec![3]
}));
}
#[test]
fn test_search_with_distance() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("appl", 2);
trie.insert("aple", 3);
trie.insert("applet", 4);
let results = trie.search_within_distance("apple", 1);
assert_eq!(results.len(), 4);
}
#[test]
fn test_multiple_data_per_word() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("apple", 2);
trie.insert("apple", 3);
let results = trie.search_within_distance("apple", 0);
assert_eq!(results.len(), 1);
assert_eq!(results[0].data, vec![1, 2, 3]);
}
#[test]
fn test_remove_all() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("app", 2);
trie.insert("application", 2);
trie.insert("apple", 2);
trie.remove_all(&2);
let results = trie.search_within_distance("app", 0);
assert_eq!(results.len(), 1);
assert_eq!(
results[0],
SearchResult {
word: "apple".to_string(),
data: vec![1]
}
);
}
#[test]
fn test_empty_string() {
let mut trie = Trie::new();
trie.insert("", 1);
trie.insert("a", 2);
let results = trie.search_within_distance("", 0);
assert_eq!(results.len(), 2);
assert_eq!(
results[0],
SearchResult {
word: "".to_string(),
data: vec![1]
}
);
assert_eq!(
results[1],
SearchResult {
word: "a".to_string(),
data: vec![2]
}
);
}
#[test]
fn test_long_words() {
let mut trie = Trie::new();
let long_word = "supercalifragilisticexpialidocious";
trie.insert(long_word, 1);
let results = trie.search_within_distance(long_word, 0);
assert_eq!(results.len(), 1);
assert_eq!(
results[0],
SearchResult {
word: long_word.to_string(),
data: vec![1]
}
);
}
#[test]
fn test_prefix_search() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("application", 2);
trie.insert("appreciate", 3);
let results = trie.search_within_distance("app", 0);
assert_eq!(results.len(), 3);
}
#[test]
fn test_case_sensitivity() {
let mut trie = Trie::new();
trie.insert("Apple", 1);
trie.insert("apple", 2);
let results = trie.search_within_distance("Apple", 0);
assert_eq!(results.len(), 1);
assert_eq!(
results[0],
SearchResult {
word: "Apple".to_string(),
data: vec![1]
}
);
}
#[test]
fn test_remove_and_reinsert() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.remove_all(&1);
trie.insert("apple", 2);
let results = trie.search_within_distance("apple", 0);
assert_eq!(results.len(), 1);
assert_eq!(
results[0],
SearchResult {
word: "apple".to_string(),
data: vec![2]
}
);
}
#[test]
fn test_large_distance_search() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("banana", 2);
trie.insert("cherry", 3);
let results = trie.search_within_distance("grape", 5);
println!("{:?}", results);
assert_eq!(results.len(), 2);
assert!(results.contains(&SearchResult {
word: "apple".to_string(),
data: vec![1]
}));
assert!(results.contains(&SearchResult {
word: "banana".to_string(),
data: vec![2]
}));
}
#[test]
fn test_prefix_additions_with_distance() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("app", 2);
trie.insert("application", 3);
trie.insert("shouldnotbefound", 4);
let results = trie.search_within_distance("capp", 1);
assert_eq!(results.len(), 3);
}
#[test]
fn test_prefix_deletions_with_distance() {
let mut trie = Trie::new();
trie.insert("apple", 1);
trie.insert("application", 3);
trie.insert("app", 2);
trie.insert("shouldnotbefound", 4);
let results = trie.search_within_distance("ppl", 1);
assert_eq!(results.len(), 2);
}
#[test]
fn test_prefix_deletions_with_distance_2() {
let mut trie = Trie::new();
trie.insert("appleton", 1);
trie.insert("apple", 2);
trie.insert("matrix", 3);
trie.insert("appleton", 2);
trie.insert("applitix", 3);
trie.insert("applutux", 4);
trie.insert("applet", 5);
trie.insert("capplenex", 6);
trie.insert("capplunix", 7);
let results = trie.search_within_distance("applu", 1);
assert_eq!(results.len(), 6);
assert!(results.contains(&SearchResult {
word: "capplunix".to_string(),
data: vec![7]
}));
assert!(results.contains(&SearchResult {
word: "applitix".to_string(),
data: vec![3]
}));
}
}