use std::cmp::Ordering;
#[derive(Debug, Clone)]
pub struct TrieNode<Key, Val> {
key: Vec<Key>,
value: Option<Val>,
children: Vec<Box<TrieNode<Key, Val>>>,
}
impl<Key, Val> TrieNode<Key, Val>
where
Key: Clone + Ord + PartialEq,
{
pub fn new(key: Vec<Key>) -> Self {
Self {
key,
value: None,
children: Vec::new(),
}
}
pub fn new_root() -> Self {
Self {
key: Vec::new(),
value: None,
children: Vec::new(),
}
}
pub fn key(&self) -> &[Key] {
&self.key
}
pub fn value(&self) -> Option<&Val> {
self.value.as_ref()
}
pub fn set_value(&mut self, value: Val) {
self.value = Some(value);
}
pub fn remove_value(&mut self) -> Option<Val> {
self.value.take()
}
fn common_prefix_len(a: &[Key], b: &[Key]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
fn find_child_index(&self, first_key: &Key) -> Result<usize, usize> {
self.children.binary_search_by(|child| {
match child.key.first() {
Some(k) => k.cmp(first_key),
None => Ordering::Less, }
})
}
pub fn insert(&mut self, key: &[Key], value: Val) {
if key.is_empty() {
self.value = Some(value);
return;
}
let first_key = &key[0];
match self.find_child_index(first_key) {
Ok(index) => {
let child = &mut self.children[index];
let common_len = Self::common_prefix_len(&child.key, key);
if common_len == child.key.len() {
child.insert(&key[common_len..], value);
} else if common_len == key.len() {
let old_key = child.key.clone();
let old_value = child.value.take();
let old_children = std::mem::take(&mut child.children);
child.key = key.to_vec();
child.value = Some(value);
child.children.clear();
let new_child = Box::new(TrieNode {
key: old_key[common_len..].to_vec(),
value: old_value,
children: old_children,
});
child.children.push(new_child);
} else {
let old_key = child.key.clone();
let old_value = child.value.take();
let old_children = std::mem::take(&mut child.children);
child.key = key[..common_len].to_vec();
child.value = None;
child.children.clear();
let old_child = Box::new(TrieNode {
key: old_key[common_len..].to_vec(),
value: old_value,
children: old_children,
});
let new_child = Box::new(TrieNode {
key: key[common_len..].to_vec(),
value: Some(value),
children: Vec::new(),
});
if old_child.key.first() <= new_child.key.first() {
child.children.push(old_child);
child.children.push(new_child);
} else {
child.children.push(new_child);
child.children.push(old_child);
}
}
}
Err(index) => {
let new_child = Box::new(TrieNode {
key: key.to_vec(),
value: Some(value),
children: Vec::new(),
});
self.children.insert(index, new_child);
}
}
}
pub fn get(&self, key: &[Key]) -> Option<&Val> {
if key.is_empty() {
return self.value.as_ref();
}
let first_key = &key[0];
if let Ok(index) = self.find_child_index(first_key) {
let child = &self.children[index];
let common_len = Self::common_prefix_len(&child.key, key);
if common_len == child.key.len() && common_len <= key.len() {
child.get(&key[common_len..])
} else {
None
}
} else {
None
}
}
pub fn contains_key(&self, key: &[Key]) -> bool {
self.get(key).is_some()
}
pub fn remove(&mut self, key: &[Key]) -> (Option<Val>, bool) {
if key.is_empty() {
let old_value = self.value.take();
let should_remove = old_value.is_some() && self.children.is_empty();
return (old_value, should_remove);
}
let first_key = &key[0];
if let Ok(index) = self.find_child_index(first_key) {
let common_len = Self::common_prefix_len(&self.children[index].key, key);
if common_len == self.children[index].key.len() && common_len <= key.len() {
let (removed_value, should_remove_child) =
self.children[index].remove(&key[common_len..]);
if should_remove_child {
self.children.remove(index);
}
let should_remove_self = self.value.is_none() && self.children.is_empty();
(removed_value, should_remove_self)
} else {
(None, false)
}
} else {
(None, false)
}
}
pub fn iter(&self) -> Vec<(Vec<Key>, &Val)> {
let mut result = Vec::new();
self.collect_all(&mut Vec::new(), &mut result);
result
}
fn collect_all<'a>(
&'a self,
current_path: &mut Vec<Key>,
result: &mut Vec<(Vec<Key>, &'a Val)>,
) {
let mut full_path = current_path.clone();
full_path.extend_from_slice(&self.key);
if let Some(value) = &self.value {
result.push((full_path.clone(), value));
}
for child in &self.children {
child.collect_all(&mut full_path, result);
}
}
pub fn keys_with_prefix(&self, prefix: &[Key]) -> Vec<Vec<Key>> {
if let Some((node, actual_prefix)) = self.find_node_with_prefix_and_path(prefix) {
let mut result = Vec::new();
let mut current_path = actual_prefix;
node.collect_keys(&mut current_path, &mut result);
result
} else {
Vec::new()
}
}
fn find_node_with_prefix_and_path(&self, prefix: &[Key]) -> Option<(&Self, Vec<Key>)> {
self.find_node_with_prefix_and_path_helper(prefix, Vec::new())
}
fn find_node_with_prefix_and_path_helper(
&self,
prefix: &[Key],
mut current_path: Vec<Key>,
) -> Option<(&Self, Vec<Key>)> {
current_path.extend_from_slice(&self.key);
if prefix.is_empty() {
return Some((self, current_path));
}
if current_path.len() >= prefix.len() {
if current_path[..prefix.len()] == prefix[..] {
return Some((self, current_path));
} else {
return None;
}
}
let remaining_prefix = &prefix[current_path.len()..];
if remaining_prefix.is_empty() {
return Some((self, current_path));
}
let first_key = &remaining_prefix[0];
if let Ok(index) = self.find_child_index(first_key) {
let child = &self.children[index];
child.find_node_with_prefix_and_path_helper(prefix, current_path)
} else {
None
}
}
fn collect_keys(&self, current_path: &mut Vec<Key>, result: &mut Vec<Vec<Key>>) {
if self.value.is_some() {
result.push(current_path.clone());
}
for child in &self.children {
let old_len = current_path.len();
current_path.extend_from_slice(&child.key);
child.collect_keys(current_path, result);
current_path.truncate(old_len);
}
}
fn common_prefix_length_helper(&self, key: &[Key], current_path: &mut Vec<Key>) -> usize {
let path_before = current_path.len();
current_path.extend_from_slice(&self.key);
let current_common = Self::common_prefix_len(current_path, key);
let max_possible = current_path.len().min(key.len());
if current_common < max_possible && current_common < current_path.len() {
current_path.truncate(path_before);
return current_common;
}
if key.len() <= current_path.len() {
current_path.truncate(path_before);
return key.len();
}
let remaining_key = &key[current_path.len()..];
if remaining_key.is_empty() {
current_path.truncate(path_before);
return current_path.len();
}
let first_remaining = &remaining_key[0];
let mut best_common = current_path.len();
if let Ok(index) = self.find_child_index(first_remaining) {
let child = &self.children[index];
let child_common = child.common_prefix_length_helper(key, current_path);
best_common = best_common.max(child_common);
}
current_path.truncate(path_before);
best_common
}
}
#[derive(Debug, Clone)]
pub struct Trie<Key, Val> {
root: Box<TrieNode<Key, Val>>,
}
impl<Key, Val> Trie<Key, Val>
where
Key: Clone + Ord + PartialEq,
{
pub fn new() -> Self {
Self {
root: Box::new(TrieNode::new_root()),
}
}
pub fn insert(&mut self, key: &[Key], value: Val) {
self.root.insert(key, value);
}
pub fn get(&self, key: &[Key]) -> Option<&Val> {
self.root.get(key)
}
pub fn contains_key(&self, key: &[Key]) -> bool {
self.root.contains_key(key)
}
pub fn remove(&mut self, key: &[Key]) -> Option<Val> {
let (removed_value, _) = self.root.remove(key);
removed_value
}
pub fn iter(&self) -> Vec<(Vec<Key>, &Val)> {
self.root.iter()
}
pub fn keys(&self) -> Vec<Vec<Key>> {
self.root.keys_with_prefix(&[])
}
pub fn keys_with_prefix(&self, prefix: &[Key]) -> Vec<Vec<Key>> {
self.root.keys_with_prefix(prefix)
}
pub fn is_empty(&self) -> bool {
self.root.value.is_none() && self.root.children.is_empty()
}
pub fn len(&self) -> usize {
self.iter().len()
}
pub fn common_prefix_length(&self, key: &[Key]) -> usize {
let mut path = Vec::new();
self.root.common_prefix_length_helper(key, &mut path)
}
}
impl<Key, Val> Default for Trie<Key, Val>
where
Key: Clone + Ord + PartialEq,
{
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_operations() {
let mut trie = Trie::new();
assert!(trie.is_empty());
assert_eq!(trie.len(), 0);
trie.insert(&['h', 'e', 'l', 'l', 'o'], "world");
trie.insert(&['h', 'e', 'l', 'p'], "me");
trie.insert(&['h', 'i'], "there");
assert!(!trie.is_empty());
assert_eq!(trie.len(), 3);
assert_eq!(trie.get(&['h', 'e', 'l', 'l', 'o']), Some(&"world"));
assert_eq!(trie.get(&['h', 'e', 'l', 'p']), Some(&"me"));
assert_eq!(trie.get(&['h', 'i']), Some(&"there"));
assert_eq!(trie.get(&['h', 'e']), None);
assert!(trie.contains_key(&['h', 'e', 'l', 'l', 'o']));
assert!(trie.contains_key(&['h', 'i']));
assert!(!trie.contains_key(&['h', 'e']));
let hello_keys = trie.keys_with_prefix(&['h', 'e']);
assert_eq!(hello_keys.len(), 2);
assert!(hello_keys.contains(&vec!['h', 'e', 'l', 'l', 'o']));
assert!(hello_keys.contains(&vec!['h', 'e', 'l', 'p']));
}
#[test]
fn test_removal() {
let mut trie = Trie::new();
trie.insert(&['a', 'b', 'c'], 1);
trie.insert(&['a', 'b', 'd'], 2);
trie.insert(&['a', 'c'], 3);
assert_eq!(trie.len(), 3);
assert_eq!(trie.remove(&['a', 'b', 'c']), Some(1));
assert_eq!(trie.len(), 2);
assert!(!trie.contains_key(&['a', 'b', 'c']));
assert!(trie.contains_key(&['a', 'b', 'd']));
assert_eq!(trie.remove(&['x', 'y', 'z']), None);
assert_eq!(trie.len(), 2);
}
#[test]
fn test_radix_compression() {
let mut trie = Trie::new();
trie.insert(&['a', 'b', 'c', 'd', 'e', 'f'], "value");
assert_eq!(trie.get(&['a', 'b', 'c', 'd', 'e', 'f']), Some(&"value"));
assert_eq!(trie.get(&['a', 'b', 'c']), None);
trie.insert(&['a', 'b', 'c', 'x', 'y'], "other");
assert_eq!(trie.get(&['a', 'b', 'c', 'd', 'e', 'f']), Some(&"value"));
assert_eq!(trie.get(&['a', 'b', 'c', 'x', 'y']), Some(&"other"));
assert_eq!(trie.get(&['a', 'b', 'c']), None);
}
#[test]
fn test_prefix_search() {
let mut trie = Trie::new();
trie.insert(&['c', 'a', 't'], "feline");
trie.insert(&['c', 'a', 'r'], "vehicle");
trie.insert(&['c', 'a', 'r', 'd'], "paper");
trie.insert(&['d', 'o', 'g'], "canine");
let ca_keys = trie.keys_with_prefix(&['c', 'a']);
assert_eq!(ca_keys.len(), 3);
let car_keys = trie.keys_with_prefix(&['c', 'a', 'r']);
assert_eq!(car_keys.len(), 2);
assert!(car_keys.contains(&vec!['c', 'a', 'r']));
assert!(car_keys.contains(&vec!['c', 'a', 'r', 'd']));
let empty_prefix = trie.keys_with_prefix(&[]);
assert_eq!(empty_prefix.len(), 4);
}
#[test]
fn test_root_key_is_empty() {
let trie = Trie::<char, i32>::new();
assert_eq!(trie.root.key().len(), 0);
}
#[test]
fn test_usage_example() {
let mut trie = Trie::new();
trie.insert(&['c', 'a', 't'], 1);
trie.insert(&['c', 'a', 'r'], 2);
trie.insert(&['c', 'a', 'r', 'd'], 3);
trie.insert(&['d', 'o', 'g'], 4);
trie.insert(&['d', 'o', 'o', 'r'], 5);
assert_eq!(trie.get(&['c', 'a', 't']), Some(&1));
assert_eq!(trie.get(&['c', 'a', 'r']), Some(&2));
assert_eq!(trie.get(&['d', 'o', 'g']), Some(&4));
assert_eq!(trie.get(&['n', 'o', 't']), None);
let ca_keys = trie.keys_with_prefix(&['c', 'a']);
assert_eq!(ca_keys.len(), 3);
assert!(ca_keys.contains(&vec!['c', 'a', 't']));
assert!(ca_keys.contains(&vec!['c', 'a', 'r']));
assert!(ca_keys.contains(&vec!['c', 'a', 'r', 'd']));
assert_eq!(trie.remove(&['c', 'a', 't']), Some(1));
assert!(!trie.contains_key(&['c', 'a', 't']));
let all_pairs = trie.iter();
assert_eq!(all_pairs.len(), 4);
}
#[test]
fn test_common_prefix_length() {
let mut trie = Trie::new();
trie.insert(&['h', 'e', 'l', 'l', 'o'], "world");
trie.insert(&['h', 'e', 'l', 'p'], "me");
trie.insert(&['h', 'i'], "there");
trie.insert(&['c', 'a', 't'], "animal");
assert_eq!(trie.common_prefix_length(&['h', 'e', 'l', 'l', 'o']), 5);
assert_eq!(trie.common_prefix_length(&['h', 'i']), 2);
assert_eq!(trie.common_prefix_length(&['h', 'e', 'l']), 3); assert_eq!(
trie.common_prefix_length(&['h', 'e', 'l', 'l', 'o', 'w']),
5
); assert_eq!(trie.common_prefix_length(&['h', 'e', 'x']), 2);
assert_eq!(trie.common_prefix_length(&['x', 'y', 'z']), 0);
assert_eq!(trie.common_prefix_length(&[]), 0);
assert_eq!(trie.common_prefix_length(&['h']), 1); assert_eq!(trie.common_prefix_length(&['c']), 1);
let mut empty_trie = Trie::new();
assert_eq!(empty_trie.common_prefix_length(&['a', 'b', 'c']), 0);
empty_trie.insert(&[], "root");
assert_eq!(empty_trie.common_prefix_length(&['a']), 0);
let mut short_trie = Trie::new();
short_trie.insert(&['a'], "short");
assert_eq!(short_trie.common_prefix_length(&['a', 'b', 'c', 'd']), 1);
let mut branch_trie = Trie::new();
branch_trie.insert(&['p', 'r', 'e', 'f', 'i', 'x', '1'], "first");
branch_trie.insert(&['p', 'r', 'e', 'f', 'i', 'x', '2'], "second");
branch_trie.insert(&['p', 'r', 'e', 'f'], "prefix");
assert_eq!(
branch_trie.common_prefix_length(&['p', 'r', 'e', 'f', 'i', 'x', '3']),
6
); assert_eq!(
branch_trie.common_prefix_length(&['p', 'r', 'e', 'f', 'i', 'x', '1']),
7
); assert_eq!(
branch_trie.common_prefix_length(&['p', 'r', 'e', 'f', 'y']),
4
); }
#[test]
fn test_children_ordering_after_operations() {
let mut trie = Trie::new();
fn assert_children_sorted<Key: Clone + Ord + PartialEq + std::fmt::Debug, Val>(
node: &TrieNode<Key, Val>,
) {
for i in 0..node.children.len().saturating_sub(1) {
let current_first = node.children[i].key.first();
let next_first = node.children[i + 1].key.first();
assert!(
current_first <= next_first,
"Children not sorted: {:?} should come before {:?}",
current_first,
next_first
);
}
for child in &node.children {
assert_children_sorted(child);
}
}
trie.insert(&['z', 'e', 'b', 'r', 'a'], 1);
assert_children_sorted(&trie.root);
trie.insert(&['a', 'p', 'p', 'l', 'e'], 2);
assert_children_sorted(&trie.root);
trie.insert(&['m', 'a', 'n', 'g', 'o'], 3);
assert_children_sorted(&trie.root);
trie.insert(&['b', 'a', 'n', 'a', 'n', 'a'], 4);
assert_children_sorted(&trie.root);
trie.insert(&['a', 'p', 'r', 'i', 'c', 'o', 't'], 5); assert_children_sorted(&trie.root);
trie.insert(&['a', 'p', 'p'], 6); assert_children_sorted(&trie.root);
trie.insert(&['a', 'p', 'p', 'l', 'e', 's'], 7); assert_children_sorted(&trie.root);
trie.insert(&['m', 'a', 'n'], 8); assert_children_sorted(&trie.root);
trie.insert(&['m', 'a', 'n', 'd', 'a', 'r', 'i', 'n'], 9); assert_children_sorted(&trie.root);
trie.remove(&['a', 'p', 'p']);
assert_children_sorted(&trie.root);
trie.remove(&['a', 'p', 'p', 'l', 'e']);
assert_children_sorted(&trie.root);
trie.remove(&['m', 'a', 'n']);
assert_children_sorted(&trie.root);
trie.insert(&['a', 'b', 'c'], 10);
assert_children_sorted(&trie.root);
trie.insert(&['a', 'x', 'y'], 11);
assert_children_sorted(&trie.root);
let mut prefix_trie = Trie::new();
prefix_trie.insert(&['p', 'r', 'e'], 1);
assert_children_sorted(&prefix_trie.root);
prefix_trie.insert(&['p', 'r', 'e', 'f', 'i', 'x'], 2);
assert_children_sorted(&prefix_trie.root);
prefix_trie.insert(&['p', 'r'], 3);
assert_children_sorted(&prefix_trie.root);
prefix_trie.insert(&['p'], 4);
assert_children_sorted(&prefix_trie.root);
let mut multi_trie = Trie::new();
multi_trie.insert(&['c', 'a', 't'], 1);
multi_trie.insert(&['c', 'a', 'r'], 2);
multi_trie.insert(&['c', 'a', 'n'], 3);
multi_trie.insert(&['c', 'a', 'p'], 4);
multi_trie.insert(&['c', 'a', 'b'], 5);
assert_children_sorted(&multi_trie.root);
multi_trie.remove(&['c', 'a', 'n']);
assert_children_sorted(&multi_trie.root);
multi_trie.remove(&['c', 'a', 'r']);
assert_children_sorted(&multi_trie.root);
assert!(multi_trie.contains_key(&['c', 'a', 't']));
assert!(multi_trie.contains_key(&['c', 'a', 'p']));
assert!(multi_trie.contains_key(&['c', 'a', 'b']));
assert!(!multi_trie.contains_key(&['c', 'a', 'n']));
assert!(!multi_trie.contains_key(&['c', 'a', 'r']));
}
#[test]
fn test_binary_search_efficiency() {
let mut trie = Trie::new();
let keys = vec![
['a', 'a'],
['a', 'b'],
['a', 'c'],
['a', 'd'],
['a', 'e'],
['b', 'a'],
['b', 'b'],
['b', 'c'],
['b', 'd'],
['b', 'e'],
['c', 'a'],
['c', 'b'],
['c', 'c'],
['c', 'd'],
['c', 'e'],
['d', 'a'],
['d', 'b'],
['d', 'c'],
['d', 'd'],
['d', 'e'],
['e', 'a'],
['e', 'b'],
['e', 'c'],
['e', 'd'],
['e', 'e'],
];
for (i, key) in keys.iter().enumerate() {
trie.insert(key, i);
}
for (i, key) in keys.iter().enumerate() {
assert_eq!(trie.get(key), Some(&i));
}
assert_eq!(trie.get(&['f', 'f']), None);
assert_eq!(trie.get(&['a', 'f']), None);
assert_eq!(trie.get(&['z', 'z']), None);
fn verify_binary_search_property<Key: Clone + Ord + PartialEq + std::fmt::Debug, Val>(
node: &TrieNode<Key, Val>,
) {
for i in 0..node.children.len() {
for j in i + 1..node.children.len() {
let left_first = node.children[i].key.first();
let right_first = node.children[j].key.first();
assert!(
left_first < right_first,
"Binary search property violated: {:?} >= {:?}",
left_first,
right_first
);
}
}
for child in &node.children {
verify_binary_search_property(child);
}
}
verify_binary_search_property(&trie.root);
}
#[test]
fn test_children_ordering_extreme_cases() {
let mut trie = Trie::new();
trie.insert(&['a', 'a', 'a', 'a'], 1);
trie.insert(&['a', 'a', 'a', 'b'], 2);
trie.insert(&['a', 'a', 'b', 'a'], 3);
trie.insert(&['a', 'b', 'a', 'a'], 4);
trie.insert(&['b', 'a', 'a', 'a'], 5);
fn check_ordering<Key: Clone + Ord + PartialEq + std::fmt::Debug, Val>(
node: &TrieNode<Key, Val>,
path: &mut Vec<Key>,
) {
for i in 0..node.children.len().saturating_sub(1) {
let left = &node.children[i];
let right = &node.children[i + 1];
match (left.key.first(), right.key.first()) {
(Some(l), Some(r)) => assert!(
l <= r,
"Children not ordered at path {:?}: {:?} should come before {:?}",
path,
l,
r
),
(None, Some(_)) => {} (Some(_), None) => panic!("Empty key should come first"),
(None, None) => {} }
}
for child in &node.children {
path.extend_from_slice(&child.key);
check_ordering(child, path);
path.truncate(path.len() - child.key.len());
}
}
let mut path = Vec::new();
check_ordering(&trie.root, &mut path);
let mut reverse_trie = Trie::new();
let reverse_keys = vec![
['z', 'z', 'z'],
['y', 'y', 'y'],
['x', 'x', 'x'],
['w', 'w', 'w'],
['v', 'v', 'v'],
['u', 'u', 'u'],
['t', 't', 't'],
['s', 's', 's'],
['r', 'r', 'r'],
['q', 'q', 'q'],
];
for (i, key) in reverse_keys.iter().enumerate() {
reverse_trie.insert(key, i);
let mut path = Vec::new();
check_ordering(&reverse_trie.root, &mut path);
}
let mut interleaved_trie = Trie::new();
interleaved_trie.insert(&['m'], 1);
interleaved_trie.insert(&['a'], 2);
interleaved_trie.insert(&['z'], 3);
let mut path = Vec::new();
check_ordering(&interleaved_trie.root, &mut path);
interleaved_trie.remove(&['m']);
path.clear();
check_ordering(&interleaved_trie.root, &mut path);
interleaved_trie.insert(&['k'], 4);
interleaved_trie.insert(&['p'], 5);
path.clear();
check_ordering(&interleaved_trie.root, &mut path);
let mut split_trie = Trie::new();
split_trie.insert(
&[
'c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'f', 'i', 'x', '1',
],
1,
);
split_trie.insert(
&[
'c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'f', 'i', 'x', '2',
],
2,
);
split_trie.insert(
&[
'c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'f', 'i', 'x', '3',
],
3,
);
split_trie.insert(
&['c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'f', 'i', 'y'],
4,
);
split_trie.insert(&['c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'g'], 5);
split_trie.insert(&['c', 'o', 'm', 'm', 'o', 'n', 'q'], 6);
split_trie.insert(&['c', 'o', 'm', 'p'], 7);
split_trie.insert(&['c', 'o', 'n'], 8);
split_trie.insert(&['d'], 9);
path.clear();
check_ordering(&split_trie.root, &mut path);
assert_eq!(
split_trie.get(&[
'c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'f', 'i', 'x', '2'
]),
Some(&2)
);
assert_eq!(split_trie.get(&['c', 'o', 'm', 'p']), Some(&7));
assert_eq!(split_trie.get(&['d']), Some(&9));
assert_eq!(
split_trie.get(&['c', 'o', 'm', 'm', 'o', 'n', 'p', 'r', 'e', 'g']),
Some(&5)
);
}
#[test]
fn test_random_operations_with_int_keys() {
use std::collections::HashMap;
let mut trie = Trie::new();
let mut reference = HashMap::new();
let mut rng_state = 12345u64;
fn next_random(state: &mut u64) -> u64 {
*state = state.wrapping_mul(1103515245).wrapping_add(12345);
*state
}
fn generate_random_key(state: &mut u64, len: usize) -> Vec<i32> {
(0..len)
.map(|_| (next_random(state) % 100) as i32)
.collect()
}
fn verify_consistency(trie: &Trie<i32, String>, reference: &HashMap<Vec<i32>, String>) {
for (key, expected_value) in reference {
let trie_value = trie.get(key);
assert_eq!(
trie_value,
Some(expected_value),
"Key {:?} mismatch: trie={:?}, reference={:?}",
key,
trie_value,
Some(expected_value)
);
}
let trie_pairs = trie.iter();
assert_eq!(
trie_pairs.len(),
reference.len(),
"Trie has {} items but reference has {}",
trie_pairs.len(),
reference.len()
);
fn check_ordering(node: &TrieNode<i32, String>) {
for i in 0..node.children.len().saturating_sub(1) {
let left_first = node.children[i].key.first();
let right_first = node.children[i + 1].key.first();
assert!(
left_first <= right_first,
"Children not sorted: {:?} should come before {:?}",
left_first,
right_first
);
}
for child in &node.children {
check_ordering(child);
}
}
check_ordering(&trie.root);
}
println!("Testing random insertions...");
for i in 0..100 {
let key_len = (next_random(&mut rng_state) % 5) + 1; let key = generate_random_key(&mut rng_state, key_len as usize);
let value = format!("value_{}", i);
trie.insert(&key, value.clone());
reference.insert(key, value);
if i % 10 == 9 {
verify_consistency(&trie, &reference);
}
}
println!("Testing random queries...");
for _ in 0..50 {
let key_len = (next_random(&mut rng_state) % 5) + 1;
let key = generate_random_key(&mut rng_state, key_len as usize);
let trie_result = trie.get(&key);
let reference_result = reference.get(&key);
assert_eq!(
trie_result, reference_result,
"Query mismatch for key {:?}",
key
);
}
println!("Testing random deletions...");
let keys_to_delete: Vec<_> = reference.keys().cloned().collect();
let mut deleted_count = 0;
for (_i, key) in keys_to_delete.iter().enumerate() {
if next_random(&mut rng_state) % 3 == 0 {
let trie_removed = trie.remove(key);
let reference_removed = reference.remove(key);
assert_eq!(
trie_removed, reference_removed,
"Delete mismatch for key {:?}",
key
);
deleted_count += 1;
if deleted_count % 5 == 0 {
verify_consistency(&trie, &reference);
}
}
}
println!("Testing mixed operations...");
for i in 0..100 {
let operation = next_random(&mut rng_state) % 3;
match operation {
0 => {
let key_len = (next_random(&mut rng_state) % 4) + 1;
let key = generate_random_key(&mut rng_state, key_len as usize);
let value = format!("mixed_value_{}", i);
trie.insert(&key, value.clone());
reference.insert(key, value);
}
1 => {
let key_len = (next_random(&mut rng_state) % 4) + 1;
let key = generate_random_key(&mut rng_state, key_len as usize);
let trie_result = trie.get(&key);
let reference_result = reference.get(&key);
assert_eq!(trie_result, reference_result);
}
2 => {
if !reference.is_empty() {
let keys: Vec<_> = reference.keys().cloned().collect();
let key_index = (next_random(&mut rng_state) as usize) % keys.len();
let key = &keys[key_index];
let trie_removed = trie.remove(key);
let reference_removed = reference.remove(key);
assert_eq!(trie_removed, reference_removed);
}
}
_ => unreachable!(),
}
if i % 20 == 19 {
verify_consistency(&trie, &reference);
}
}
verify_consistency(&trie, &reference);
println!(
"Random operations test completed successfully! Final state: {} keys",
reference.len()
);
}
#[test]
fn test_edge_cases_with_int_keys() {
let mut trie = Trie::new();
trie.insert(&[], 42);
assert_eq!(trie.get(&[]), Some(&42));
assert_eq!(trie.remove(&[]), Some(42));
assert_eq!(trie.get(&[]), None);
for i in 0..10 {
trie.insert(&[i], i * 10);
}
for i in 0..10 {
assert_eq!(trie.get(&[i]), Some(&(i * 10)));
}
trie.insert(&[1, 2, 3], 123);
trie.insert(&[1, 2, 3, 4], 1234);
trie.insert(&[1, 2, 3, 4, 5], 12345);
trie.insert(&[1, 2], 12);
trie.insert(&[1], 1);
assert_eq!(trie.get(&[1]), Some(&1));
assert_eq!(trie.get(&[1, 2]), Some(&12));
assert_eq!(trie.get(&[1, 2, 3]), Some(&123));
assert_eq!(trie.get(&[1, 2, 3, 4]), Some(&1234));
assert_eq!(trie.get(&[1, 2, 3, 4, 5]), Some(&12345));
trie.insert(&[-1, -2, -3], -123);
trie.insert(&[-5, 0, 5], 505);
assert_eq!(trie.get(&[-1, -2, -3]), Some(&-123));
assert_eq!(trie.get(&[-5, 0, 5]), Some(&505));
trie.insert(&[i32::MAX, i32::MIN], 999);
assert_eq!(trie.get(&[i32::MAX, i32::MIN]), Some(&999));
assert_eq!(trie.common_prefix_length(&[1, 2, 3, 4, 5, 6]), 5);
assert_eq!(trie.common_prefix_length(&[1, 2, 3, 4]), 4);
assert_eq!(trie.common_prefix_length(&[1, 2, 3]), 3);
assert_eq!(trie.common_prefix_length(&[1, 2]), 2);
assert_eq!(trie.common_prefix_length(&[1]), 1);
assert_eq!(trie.common_prefix_length(&[2]), 1); }
#[test]
fn test_stress_with_large_int_sequences() {
use std::collections::HashMap;
let mut trie = Trie::new();
let mut reference = HashMap::new();
for length in 1..=20 {
for start in 0..5 {
let key: Vec<i32> = (start..start + length).collect();
let value = format!("seq_{}_{}", length, start);
trie.insert(&key, value.clone());
reference.insert(key, value);
}
}
for (key, expected_value) in &reference {
assert_eq!(trie.get(key), Some(expected_value));
}
let prefix_1 = trie.keys_with_prefix(&[0]);
assert!(!prefix_1.is_empty());
let prefix_0_1 = trie.keys_with_prefix(&[0, 1]);
assert!(!prefix_0_1.is_empty());
assert_eq!(trie.common_prefix_length(&[0, 1, 2, 3, 4, 5]), 6);
assert_eq!(trie.common_prefix_length(&[0, 1, 2, 999]), 3);
let keys_to_delete: Vec<_> = reference
.keys()
.filter(|k| k.len() % 3 == 0)
.cloned()
.collect();
for key in keys_to_delete {
let removed = trie.remove(&key);
let expected = reference.remove(&key);
assert_eq!(removed, expected);
}
for (key, expected_value) in &reference {
assert_eq!(trie.get(key), Some(expected_value));
}
println!(
"Stress test completed with {} remaining keys",
reference.len()
);
}
#[test]
fn test_intensive_random_operations() {
use std::collections::HashMap;
let mut trie = Trie::new();
let mut reference = HashMap::new();
let mut rng_state = 98765u64;
fn next_random(state: &mut u64) -> u64 {
*state = state.wrapping_mul(1103515245).wrapping_add(12345);
*state
}
fn generate_diverse_key(state: &mut u64, pattern: usize) -> Vec<i32> {
match pattern % 5 {
0 => {
let len = (next_random(state) % 3) + 1;
(0..len).map(|_| (next_random(state) % 20) as i32).collect()
}
1 => {
let len = (next_random(state) % 5) + 4;
(0..len).map(|_| (next_random(state) % 10) as i32).collect()
}
2 => {
let start = (next_random(state) % 10) as i32;
let len = (next_random(state) % 6) + 1;
(start..start + len as i32).collect()
}
3 => {
let len = (next_random(state) % 4) + 1;
(0..len)
.map(|_| -((next_random(state) % 50) as i32))
.collect()
}
4 => {
let len = (next_random(state) % 5) + 1;
(0..len)
.map(|_| ((next_random(state) % 100) as i32) - 50)
.collect()
}
_ => unreachable!(),
}
}
fn verify_complete_consistency(
trie: &Trie<i32, String>,
reference: &HashMap<Vec<i32>, String>,
) {
for (key, expected_value) in reference {
match trie.get(key) {
Some(actual_value) => {
assert_eq!(
actual_value, expected_value,
"Value mismatch for key {:?}: expected {:?}, got {:?}",
key, expected_value, actual_value
);
}
None => panic!("Key {:?} exists in reference but not in trie", key),
}
}
let trie_pairs = trie.iter();
for (trie_key, trie_value) in &trie_pairs {
match reference.get(trie_key) {
Some(ref_value) => {
assert_eq!(
trie_value.as_str(),
ref_value.as_str(),
"Trie has key {:?} with value {:?} but reference has {:?}",
trie_key,
trie_value,
ref_value
);
}
None => panic!(
"Trie has key {:?} that doesn't exist in reference",
trie_key
),
}
}
assert_eq!(
trie_pairs.len(),
reference.len(),
"Size mismatch: trie has {} keys, reference has {}",
trie_pairs.len(),
reference.len()
);
fn verify_ordering_recursive(node: &TrieNode<i32, String>) {
for i in 0..node.children.len().saturating_sub(1) {
let left_first = node.children[i].key.first();
let right_first = node.children[i + 1].key.first();
assert!(
left_first <= right_first,
"Children ordering violation: {:?} should come before {:?}",
left_first,
right_first
);
}
for child in &node.children {
verify_ordering_recursive(child);
}
}
verify_ordering_recursive(&trie.root);
}
println!("Starting intensive random operations test...");
for round in 0..10 {
for i in 0..50 {
let pattern = (round * 50 + i) % 5;
let key = generate_diverse_key(&mut rng_state, pattern);
let value = format!("intensive_{}_{}", round, i);
trie.insert(&key, value.clone());
reference.insert(key, value);
}
if round % 3 == 2 {
verify_complete_consistency(&trie, &reference);
}
}
println!("Phase 1 completed: {} keys inserted", reference.len());
for round in 0..20 {
for op in 0..25 {
let operation = next_random(&mut rng_state) % 4;
match operation {
0 => {
let key = generate_diverse_key(&mut rng_state, round + op);
let value = format!("mixed_{}_{}", round, op);
trie.insert(&key, value.clone());
reference.insert(key, value);
}
1 => {
let key = generate_diverse_key(&mut rng_state, round + op + 100);
let trie_result = trie.get(&key);
let ref_result = reference.get(&key);
assert_eq!(trie_result, ref_result, "Query mismatch for {:?}", key);
}
2 => {
if !reference.is_empty() {
let keys: Vec<_> = reference.keys().cloned().collect();
let key_idx = (next_random(&mut rng_state) as usize) % keys.len();
let key = &keys[key_idx];
let trie_removed = trie.remove(key);
let ref_removed = reference.remove(key);
assert_eq!(trie_removed, ref_removed, "Remove mismatch for {:?}", key);
}
}
3 => {
if !reference.is_empty() {
let keys: Vec<_> = reference.keys().collect();
let sample_key =
&keys[(next_random(&mut rng_state) as usize) % keys.len()];
if !sample_key.is_empty() {
let prefix_len =
((next_random(&mut rng_state) as usize) % sample_key.len()) + 1;
let prefix = &sample_key[..prefix_len];
let trie_prefix_keys = trie.keys_with_prefix(prefix);
let expected_prefix_keys: Vec<_> = reference
.keys()
.filter(|k| {
k.len() >= prefix.len() && &k[..prefix.len()] == prefix
})
.cloned()
.collect();
assert_eq!(
trie_prefix_keys.len(),
expected_prefix_keys.len(),
"Prefix key count mismatch for prefix {:?}",
prefix
);
for expected_key in &expected_prefix_keys {
assert!(
trie_prefix_keys.contains(expected_key),
"Expected key {:?} not found in prefix results for {:?}",
expected_key,
prefix
);
}
}
}
}
_ => unreachable!(),
}
}
if round % 5 == 4 {
verify_complete_consistency(&trie, &reference);
println!(
"Round {} verification passed: {} keys",
round,
reference.len()
);
}
}
println!("Phase 3: Testing edge cases...");
trie.insert(&[], "empty_key".to_string());
reference.insert(vec![], "empty_key".to_string());
let long_key: Vec<i32> = (0..100).collect();
trie.insert(&long_key, "long_key".to_string());
reference.insert(long_key, "long_key".to_string());
let extreme_key = vec![i32::MIN, 0, i32::MAX];
trie.insert(&extreme_key, "extreme".to_string());
reference.insert(extreme_key, "extreme".to_string());
for i in 0..20 {
let similar_key = vec![1, 2, 3, 4, 5, i];
trie.insert(&similar_key, format!("similar_{}", i));
reference.insert(similar_key, format!("similar_{}", i));
}
verify_complete_consistency(&trie, &reference);
for _ in 0..50 {
let random_val = next_random(&mut rng_state);
let test_key = generate_diverse_key(&mut rng_state, random_val as usize);
let prefix_len = trie.common_prefix_length(&test_key);
if prefix_len > 0 {
let prefix = &test_key[..prefix_len.min(test_key.len())];
let matching_keys = trie.keys_with_prefix(prefix);
assert!(
!matching_keys.is_empty(),
"common_prefix_length returned {} for {:?} but no keys match prefix {:?}",
prefix_len,
test_key,
prefix
);
}
}
println!("Intensive random operations test completed successfully!");
println!("Final state: {} keys in trie", reference.len());
println!("Tree structure verified for complete correctness");
}
}
impl<Key, Val> std::fmt::Display for Trie<Key, Val>
where
Key: Clone + Ord + PartialEq + std::fmt::Debug,
Val: std::fmt::Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(f, "Trie {{")?;
let pairs = self.iter();
for (key, value) in pairs {
writeln!(f, " {:?} => {:?}", key, value)?;
}
write!(f, "}}")
}
}