pub mod metrics;
use std::borrow::Borrow;
use std::collections::{hash_map, HashMap};
use std::fmt::{self, Debug, Formatter};
use std::iter::Extend;
pub trait Metric<K: ?Sized> {
fn distance(&self, a: &K, b: &K) -> u64;
}
struct BKNode<K> {
key: K,
children: HashMap<u64, BKNode<K>>,
}
impl<K> BKNode<K>
{
pub fn new(key: K) -> BKNode<K>
{
BKNode {
key: key,
children: HashMap::new(),
}
}
pub fn add_child(&mut self, distance: u64, key: K) {
self.children.insert(distance, BKNode::new(key));
}
}
impl<K> Debug for BKNode<K> where K: Debug
{
fn fmt(&self, f: &mut Formatter) -> fmt::Result {
f.debug_map().entry(&self.key, &self.children).finish()
}
}
#[derive(Debug)]
pub struct BKTree<K, M = metrics::Levenshtein>
{
root: Option<BKNode<K>>,
metric: M,
}
impl<K, M> BKTree<K, M>
where M: Metric<K>
{
pub fn new(metric: M) -> BKTree<K, M>
{
BKTree {
root: None,
metric: metric,
}
}
pub fn add(&mut self, key: K) {
match self.root {
Some(ref mut root) => {
let mut cur_node = root;
let mut cur_dist = self.metric.distance(&cur_node.key, &key);
while cur_node.children.contains_key(&cur_dist) && cur_dist > 0 {
let current = cur_node;
let next_node = current.children.get_mut(&cur_dist).unwrap();
cur_node = next_node;
cur_dist = self.metric.distance(&cur_node.key, &key);
}
cur_node.add_child(cur_dist, key);
}
None => {
self.root = Some(BKNode::new(key));
}
}
}
pub fn find<'a, 'q, Q: ?Sized>(&'a self, key: &'q Q, tolerance: u64) -> Find<'a, 'q, K, Q, M>
where K: Borrow<Q>, M: Metric<Q>
{
Find {
root: self.root.as_ref(),
stack: Vec::new(),
tolerance: tolerance,
metric: &self.metric,
key: key,
}
}
pub fn find_exact<Q: ?Sized>(&self, key: &Q) -> Option<&K>
where K: Borrow<Q>, M: Metric<Q>
{
self.find(key, 0).next().map(|(_, found_key)| found_key)
}
}
impl<K, M: Metric<K>> Extend<K> for BKTree<K, M> {
fn extend<I: IntoIterator<Item = K>>(&mut self, keys: I) {
for key in keys {
self.add(key);
}
}
}
impl<K: AsRef<str>> Default for BKTree<K> {
fn default() -> BKTree<K> {
BKTree::new(metrics::Levenshtein)
}
}
pub struct Find<'a, 'q, K: 'a, Q: 'q + ?Sized, M: 'a>
{
root: Option<&'a BKNode<K>>,
stack: Vec<StackItem<'a, K>>,
tolerance: u64,
metric: &'a M,
key: &'q Q,
}
struct StackItem<'a, K: 'a> {
cur_dist: u64,
children_iter: hash_map::Iter<'a, u64, BKNode<K>>,
}
enum StackAction<'a, K: 'a>
{
Push(&'a BKNode<K>),
Pop,
}
impl<'a, 'q, K, Q: ?Sized, M> Iterator for Find<'a, 'q, K, Q, M>
where K: Borrow<Q>, M: Metric<Q>
{
type Item = (u64, &'a K);
fn next(&mut self) -> Option<(u64, &'a K)> {
if let Some(root) = self.root.take() {
let cur_dist = self.metric.distance(self.key, root.key.borrow() as &Q);
self.stack.push(StackItem {
cur_dist: cur_dist,
children_iter: root.children.iter(),
});
if cur_dist <= self.tolerance {
return Some((cur_dist, &root.key));
}
}
loop {
let action = match self.stack.last_mut() {
Some(stack_top) => {
let min_dist = stack_top.cur_dist.saturating_sub(self.tolerance);
let max_dist = stack_top.cur_dist.saturating_add(self.tolerance);
let mut action = StackAction::Pop;
for (dist, child_node) in &mut stack_top.children_iter {
if min_dist <= *dist && *dist <= max_dist {
action = StackAction::Push(child_node);
break;
}
}
action
},
None => return None,
};
match action {
StackAction::Push(child_node) => {
let cur_dist = self.metric.distance(self.key, child_node.key.borrow() as &Q);
self.stack.push(StackItem {
cur_dist: cur_dist,
children_iter: child_node.children.iter(),
});
if cur_dist <= self.tolerance {
return Some((cur_dist, &child_node.key));
}
},
StackAction::Pop => {
self.stack.pop();
},
}
}
}
}
#[cfg(test)]
mod tests {
use std::fmt::Debug;
use {BKNode, BKTree};
fn assert_eq_sorted<'t, T: 't, I>(left: I, right: &[(u64, T)])
where T: Ord + Debug, I: Iterator<Item=(u64, &'t T)>
{
let mut left_mut: Vec<_> = left.collect();
let mut right_mut: Vec<_> = right.iter().map(|&(dist, ref key)| (dist, key)).collect();
left_mut.sort();
right_mut.sort();
assert_eq!(left_mut, right_mut);
}
#[test]
fn node_construct() {
let node: BKNode<&str> = BKNode::new("foo");
assert_eq!(node.key, "foo");
assert!(node.children.is_empty());
}
#[test]
fn tree_construct() {
let tree: BKTree<&str> = Default::default();
assert!(tree.root.is_none());
}
#[test]
fn tree_add() {
let mut tree: BKTree<&str> = Default::default();
tree.add("foo");
match tree.root {
Some(ref root) => {
assert_eq!(root.key, "foo");
},
None => { assert!(false); }
}
tree.add("fop");
tree.add("f\u{e9}\u{e9}");
match tree.root {
Some(ref root) => {
assert_eq!(root.children.get(&1).unwrap().key, "fop");
assert_eq!(root.children.get(&2).unwrap().key, "f\u{e9}\u{e9}");
},
None => { assert!(false); }
}
}
#[test]
fn tree_extend() {
let mut tree: BKTree<&str> = Default::default();
tree.extend(vec!["foo", "fop"]);
match tree.root {
Some(ref root) => {
assert_eq!(root.key, "foo");
},
None => { assert!(false); }
}
assert_eq!(tree.root.unwrap().children.get(&1).unwrap().key, "fop");
}
#[test]
fn tree_find() {
let mut tree: BKTree<&str> = Default::default();
tree.add("book");
tree.add("books");
tree.add("cake");
tree.add("boo");
tree.add("cape");
tree.add("boon");
tree.add("cook");
tree.add("cart");
assert_eq_sorted(tree.find("caqe", 1), &[(1, "cake"), (1, "cape")]);
assert_eq_sorted(tree.find("cape", 1), &[(1, "cake"), (0, "cape")]);
assert_eq_sorted(tree.find("book", 1), &[(0, "book"), (1, "books"), (1, "boo"), (1, "boon"), (1, "cook")]);
assert_eq_sorted(tree.find("book", 0), &[(0, "book")]);
assert!(tree.find("foobar", 1).next().is_none());
}
#[test]
fn tree_find_exact() {
let mut tree: BKTree<&str> = Default::default();
tree.add("book");
tree.add("books");
tree.add("cake");
tree.add("boo");
tree.add("cape");
tree.add("boon");
tree.add("cook");
tree.add("cart");
assert_eq!(tree.find_exact("caqe"), None);
assert_eq!(tree.find_exact("cape"), Some(&"cape"));
assert_eq!(tree.find_exact("book"), Some(&"book"));
}
}