#![feature(slicing_syntax)]
#![allow(unstable)]
extern crate "test" as test_crate;
use std::hash::Hash;
use std::collections::hash_map::{self, HashMap, Hasher, Entry};
use std::fmt::{self, Formatter, Debug};
pub struct SequenceTrie<K, V> {
pub value: Option<V>,
pub children: HashMap<K, SequenceTrie<K, V>>
}
impl<K, V> SequenceTrie<K, V> where K: PartialEq + Eq + Hash<Hasher> + Clone {
pub fn new() -> SequenceTrie<K, V> {
SequenceTrie {
value: None,
children: HashMap::new()
}
}
pub fn insert(&mut self, key: &[K], value: V) -> bool {
let key_node = key.iter().fold(self, |current_node, fragment| {
match current_node.children.entry(fragment.clone()) {
Entry::Vacant(slot) => slot.insert(SequenceTrie::new()),
Entry::Occupied(slot) => slot.into_mut()
}
});
let is_new_value = match key_node.value {
Some(_) => false,
None => true
};
key_node.value = Some(value);
is_new_value
}
pub fn get(&self, key: &[K]) -> Option<&V> {
self.get_node(key).and_then(|node| node.value.as_ref())
}
pub fn get_node(&self, key: &[K]) -> Option<&SequenceTrie<K, V>> {
let fragment = match key.first() {
Some(head) => head,
None => return Some(self)
};
match self.children.get(fragment) {
Some(node) => node.get_node(&key[1..]),
None => None
}
}
pub fn get_mut(&mut self, key: &[K]) -> Option<&mut V> {
self.get_mut_node(key).and_then(|node| node.value.as_mut())
}
pub fn get_mut_node(&mut self, key: &[K]) -> Option<&mut SequenceTrie<K, V>> {
let fragment = match key.first() {
Some(head) => head,
None => return Some(self)
};
match self.children.get_mut(fragment) {
Some(node) => node.get_mut_node(&key[1..]),
None => None
}
}
pub fn get_prefix_nodes(&self, key: &[K]) -> Vec<&SequenceTrie<K, V>> {
let mut node_path = vec![self];
for fragment in key.iter() {
match node_path.last().unwrap().children.get(fragment) {
Some(node) => node_path.push(node),
None => break
}
}
node_path
}
pub fn get_ancestor(&self, key: &[K]) -> Option<&V> {
self.get_ancestor_node(key).and_then(|node| node.value.as_ref())
}
pub fn get_ancestor_node(&self, key: &[K]) -> Option<&SequenceTrie<K, V>> {
let node_path = self.get_prefix_nodes(key);
for node in node_path.iter().rev() {
if node.value.is_some() {
return Some(*node);
}
}
None
}
pub fn remove(&mut self, key: &[K]) {
self.remove_recursive(key);
}
fn remove_recursive(&mut self, key: &[K]) -> bool {
match key.first() {
None => { self.value = None; },
Some(fragment) => {
if let Entry::Occupied(mut entry) = self.children.entry(fragment.clone()) {
let delete_child = entry.get_mut().remove_recursive(&key[1..]);
if delete_child {
entry.remove();
}
}
}
}
if self.children.is_empty() && self.value.is_none() {
true
} else {
false
}
}
pub fn keys<'a>(&'a self) -> Keys<'a, K, V> {
Keys {
stack: vec![IterItem { node: self, key: None, child_iter: self.children.keys() }]
}
}
}
pub struct Keys<'a, K: 'a, V: 'a> {
stack: Vec<IterItem<'a, K, V,>>
}
struct IterItem<'a, K: 'a, V: 'a> {
node: &'a SequenceTrie<K, V>,
key: Option<&'a K>,
child_iter: hash_map::Keys<'a, K, SequenceTrie<K, V>>
}
impl<'a, K, V> Iterator for Keys<'a, K, V>
where
K: PartialEq + Eq + Hash<Hasher> + Clone {
type Item = Vec<&'a K>;
fn next(&mut self) -> Option<Vec<&'a K>> {
loop {
match self.stack.last().map(|x| x.node.children.is_empty()) {
Some(true) => {
let result: Vec<&'a K> = self.stack.iter()
.skip(1)
.map(|x| x.key.unwrap())
.collect();
self.stack.pop();
return Some(result);
},
Some(false) => {
match self.stack.last_mut().unwrap().child_iter.next() {
Some(child_key) => {
let child = self.stack.last().unwrap().node.children.get(child_key).unwrap();
self.stack.push(IterItem {
node: child,
key: Some(child_key),
child_iter: child.children.keys()
});
}
None => { self.stack.pop(); }
}
},
None => return None
}
}
}
}
impl<K, V> Debug for SequenceTrie<K, V>
where
K: PartialEq + Eq + Hash<Hasher> + Clone + Debug,
V: Debug {
fn fmt(&self, fmt: &mut Formatter) -> Result<(), fmt::Error> {
try!("Trie { value: ".fmt(fmt));
try!(self.value.fmt(fmt));
try!(", children: ".fmt(fmt));
try!(self.children.fmt(fmt));
" }".fmt(fmt)
}
}
impl<K, V> Clone for SequenceTrie<K, V> where K: Clone, V: Clone {
fn clone(&self) -> SequenceTrie<K, V> {
SequenceTrie {
value: self.value.clone(),
children: self.children.clone()
}
}
}
#[cfg(test)]
mod test {
use super::SequenceTrie;
use std::collections::HashSet;
fn make_trie() -> SequenceTrie<char, u32> {
let mut trie = SequenceTrie::new();
trie.insert(&[], 0u32);
trie.insert(&['a'], 1u32);
trie.insert(&['a', 'b', 'c', 'd'], 4u32);
trie.insert(&['a', 'b', 'x', 'y'], 25u32);
trie
}
#[test]
fn get() {
let trie = make_trie();
let data = [
(vec![], Some(0u32)),
(vec!['a'], Some(1u32)),
(vec!['a', 'b'], None),
(vec!['a', 'b', 'c'], None),
(vec!['a', 'b', 'x'], None),
(vec!['a', 'b', 'c', 'd'], Some(4u32)),
(vec!['a', 'b', 'x', 'y'], Some(25u32)),
(vec!['b', 'x', 'y'], None)
];
for &(ref key, value) in data.iter() {
assert_eq!(trie.get(key.as_slice()), value.as_ref());
}
}
#[test]
fn get_mut() {
let mut trie = make_trie();
let key = ['a', 'b', 'c', 'd'];
*trie.get_mut(&key).unwrap() = 77u32;
assert_eq!(*trie.get(&key).unwrap(), 77u32);
}
#[test]
fn get_ancestor() {
let trie = make_trie();
let data = [
(vec![], 0u32),
(vec!['a'], 1u32),
(vec!['a', 'b'], 1u32),
(vec!['a', 'b', 'c'], 1u32),
(vec!['a', 'b', 'c', 'd'], 4u32),
(vec!['a', 'b', 'x'], 1u32),
(vec!['a', 'b', 'x', 'y'], 25u32),
(vec!['p', 'q'], 0u32),
(vec!['a', 'p', 'q'], 1u32)
];
for &(ref key, value) in data.iter() {
assert_eq!(*trie.get_ancestor(key.as_slice()).unwrap(), value);
}
}
#[test]
fn get_prefix_nodes() {
let trie = make_trie();
let prefix_nodes = trie.get_prefix_nodes(&['a', 'b', 'z']);
assert_eq!(prefix_nodes.len(), 3);
let values = [Some(0u32), Some(1u32), None];
for (node, value) in prefix_nodes.iter().zip(values.iter()) {
assert_eq!(node.value, *value);
}
}
#[test]
fn remove() {
let mut trie= make_trie();
println!("Remove ['a']");
println!("Before: {:?}", trie);
trie.remove(&['a']);
println!("After: {:?}", trie);
assert_eq!(trie.get_node(&['a']).unwrap().value, None);
println!("Remove []");
println!("Before: {:?}", trie);
trie.remove(&[]);
println!("After: {:?}", trie);
assert_eq!(trie.get_node(&[]).unwrap().value, None);
assert_eq!(trie.get(&['a', 'b', 'c', 'd']), Some(&4u32));
println!("Remove ['a', 'b', 'c', 'd']");
println!("Before: {:?}", trie);
trie.remove(&['a', 'b', 'c', 'd']);
println!("After: {:?}", trie);
assert!(trie.get_node(&['a', 'b', 'c', 'd']).is_none());
assert!(trie.get_node(&['a', 'b', 'c']).is_none());
assert!(trie.get_node(&['a', 'b']).is_some());
println!("Remove ['a', 'b', 'x', 'y']");
println!("Before: {:?}", trie);
trie.remove(&['a', 'b', 'x', 'y']);
println!("After: {:?}", trie);
assert!(trie.get_node(&['a', 'b', 'x', 'y']).is_none());
assert!(trie.get_node(&['a', 'b', 'x']).is_none());
assert!(trie.get_node(&['a', 'b']).is_none());
assert!(trie.get_node(&['a']).is_none());
assert!(trie.value.is_none());
assert!(trie.children.is_empty());
}
#[test]
fn key_iter() {
let trie = make_trie();
let obs_keys: HashSet<Vec<char>> = trie.keys().map(|v| -> Vec<char> {
v.iter().map(|&&x| x).collect()
}).collect();
let mut exp_keys: HashSet<Vec<char>> = HashSet::new();
exp_keys.insert(vec!['a', 'b', 'c', 'd']);
exp_keys.insert(vec!['a', 'b', 'x', 'y']);
assert_eq!(exp_keys, obs_keys);
}
#[derive(PartialEq, Eq, Hash, Clone)]
struct Key {
field: usize
}
#[test]
fn struct_key() {
SequenceTrie::<Key, usize>::new();
}
}
#[cfg(test)]
mod benchmark {
use super::SequenceTrie;
use std::collections::HashMap;
use test_crate::Bencher;
macro_rules! u32_benchmark {
($map_constructor: expr, $test_id: ident, $num_keys: expr, $key_length: expr) => (
#[bench]
fn $test_id(b: &mut Bencher) {
let mut map = $map_constructor;
let mut test_data = Vec::<Vec<u32>>::with_capacity($num_keys);
for i in range(0, $num_keys) {
let mut key = Vec::<u32>::with_capacity($key_length);
for j in range(0, $key_length) {
key[j] = (i * j) as u32;
}
test_data[i] = key;
}
b.iter(|| {
for key in test_data.iter() {
map.insert(key.as_slice(), 7u32);
}
});
}
)
}
u32_benchmark! { HashMap::new(), hashmap_k1024_l16, 1024, 16 }
u32_benchmark! { SequenceTrie::new(), trie_k1024_l16, 1024, 16 }
u32_benchmark! { HashMap::new(), hashmap_k64_l128, 64, 128 }
u32_benchmark! { SequenceTrie::new(), trie_k64_l128, 64, 128 }
}