#[cfg(feature = "serde")]
extern crate serde;
pub mod metrics;
use std::{
borrow::Borrow,
collections::VecDeque,
fmt::{Debug, Formatter, Result as FmtResult},
iter::Extend,
};
#[cfg(feature = "enable-fnv")]
extern crate fnv;
#[cfg(feature = "enable-fnv")]
use fnv::FnvHashMap;
#[cfg(not(feature = "enable-fnv"))]
use std::collections::HashMap;
pub trait Metric<K: ?Sized> {
fn distance(&self, a: &K, b: &K) -> u32;
fn threshold_distance(&self, a: &K, b: &K, threshold: u32) -> Option<u32>;
}
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
struct BKNode<K> {
key: K,
#[cfg(feature = "enable-fnv")]
children: FnvHashMap<u32, BKNode<K>>,
#[cfg(not(feature = "enable-fnv"))]
children: HashMap<u32, BKNode<K>>,
max_child_distance: Option<u32>,
}
impl<K> BKNode<K> {
pub fn new(key: K) -> BKNode<K> {
BKNode {
key,
#[cfg(feature = "enable-fnv")]
children: fnv::FnvHashMap::default(),
#[cfg(not(feature = "enable-fnv"))]
children: HashMap::default(),
max_child_distance: None,
}
}
pub fn add_child(&mut self, distance: u32, key: K) {
self.children.insert(distance, BKNode::new(key));
self.max_child_distance = self.max_child_distance.max(Some(distance));
}
}
impl<K> Debug for BKNode<K>
where
K: Debug,
{
fn fmt(&self, f: &mut Formatter) -> FmtResult {
f.debug_map().entry(&self.key, &self.children).finish()
}
}
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
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 }
}
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);
}
if cur_dist > 0 {
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: u32) -> Find<'a, 'q, K, Q, M>
where
K: Borrow<Q>,
M: Metric<Q>,
{
let candidates = if let Some(root) = &self.root {
VecDeque::from(vec![root])
} else {
VecDeque::new()
};
Find {
candidates,
tolerance,
metric: &self.metric,
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> {
candidates: VecDeque<&'a BKNode<K>>,
tolerance: u32,
metric: &'a M,
key: &'q Q,
}
impl<'a, 'q, K, Q: ?Sized, M> Iterator for Find<'a, 'q, K, Q, M>
where
K: Borrow<Q>,
M: Metric<Q>,
{
type Item = (u32, &'a K);
fn next(&mut self) -> Option<(u32, &'a K)> {
while let Some(current) = self.candidates.pop_front() {
let BKNode {
key,
children,
max_child_distance,
} = current;
let distance_cutoff = max_child_distance.unwrap_or(0) + self.tolerance;
let cur_dist = self.metric.threshold_distance(
self.key,
current.key.borrow() as &Q,
distance_cutoff,
);
if let Some(dist) = cur_dist {
let min_dist = dist.saturating_sub(self.tolerance);
let max_dist = dist.saturating_add(self.tolerance);
for (dist, child_node) in &mut children.iter() {
if min_dist <= *dist && *dist <= max_dist {
self.candidates.push_back(child_node);
}
}
if dist <= self.tolerance {
return Some((dist, &key));
}
}
}
None
}
}
#[cfg(test)]
mod tests {
extern crate bincode;
use std::fmt::Debug;
use {BKNode, BKTree};
fn assert_eq_sorted<'t, T: 't, I>(left: I, right: &[(u32, T)])
where
T: Ord + Debug,
I: Iterator<Item = (u32, &'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(&4).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"));
}
#[test]
fn one_node_tree() {
let mut tree: BKTree<&str> = Default::default();
tree.add("book");
tree.add("book");
assert_eq!(tree.root.unwrap().children.len(), 0);
}
#[cfg(feature = "serde")]
#[test]
fn test_serialization() {
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("book", 0), &[(0, "book")]);
assert_eq_sorted(tree.find("books", 0), &[(0, "books")]);
assert_eq_sorted(tree.find("cake", 0), &[(0, "cake")]);
assert_eq_sorted(tree.find("boo", 0), &[(0, "boo")]);
assert_eq_sorted(tree.find("cape", 0), &[(0, "cape")]);
assert_eq_sorted(tree.find("boon", 0), &[(0, "boon")]);
assert_eq_sorted(tree.find("cook", 0), &[(0, "cook")]);
assert_eq_sorted(tree.find("cart", 0), &[(0, "cart")]);
assert_eq_sorted(
tree.find("book", 1),
&[
(0, "book"),
(1, "books"),
(1, "boo"),
(1, "boon"),
(1, "cook"),
],
);
assert_eq!(None, tree.find_exact("This &str hasn't been added"));
let encoded_tree: Vec<u8> = bincode::serialize(&tree).unwrap();
let decoded_tree: BKTree<&str> = bincode::deserialize(&encoded_tree[..]).unwrap();
assert_eq_sorted(decoded_tree.find("book", 0), &[(0, "book")]);
assert_eq_sorted(decoded_tree.find("books", 0), &[(0, "books")]);
assert_eq_sorted(decoded_tree.find("cake", 0), &[(0, "cake")]);
assert_eq_sorted(decoded_tree.find("boo", 0), &[(0, "boo")]);
assert_eq_sorted(decoded_tree.find("cape", 0), &[(0, "cape")]);
assert_eq_sorted(decoded_tree.find("boon", 0), &[(0, "boon")]);
assert_eq_sorted(decoded_tree.find("cook", 0), &[(0, "cook")]);
assert_eq_sorted(decoded_tree.find("cart", 0), &[(0, "cart")]);
assert_eq_sorted(
decoded_tree.find("book", 1),
&[
(0, "book"),
(1, "books"),
(1, "boo"),
(1, "boon"),
(1, "cook"),
],
);
assert_eq!(None, decoded_tree.find_exact("This &str hasn't been added"));
}
}