use crate::error::TrieError;
use crate::trie_node::TrieNode;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use std::clone::Clone;
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
#[derive(Debug, Clone)]
pub struct Trie<K: Eq + Ord + Clone, V> {
root: TrieNode<K, V>,
}
impl<K: Eq + Ord + Clone, V> Trie<K, V> {
pub fn new() -> Self {
Trie {
root: TrieNode::default(),
}
}
pub fn contains_key<I: Iterator<Item = K>>(&self, key: I) -> bool {
if self.is_empty() {
return false;
}
match self.find_node(key) {
Some(node) => node.may_be_leaf(),
None => false,
}
}
pub fn get<I: Iterator<Item = K>>(&self, key: I) -> Option<&V> {
self.find_node(key).and_then(|node| node.get_value())
}
pub fn get_mut<I: Iterator<Item = K>>(&mut self, key: I) -> Option<&mut V> {
self.find_node_mut(key)
.and_then(|node| Some(node.value.as_mut().unwrap()))
}
pub fn set_value<I: Iterator<Item = K>>(&mut self, key: I, value: V) -> Result<(), TrieError> {
self.find_node_mut(key)
.ok_or_else(|| TrieError::NotFound("Key not found".to_string()))
.map(|node| node.set_value(value))
}
pub fn find_prefixes<I: Iterator<Item = K>>(&self, key: I) -> Vec<(usize, &V)> {
let mut node = &self.root;
let mut prefixes = Vec::new();
for (i, k) in key.enumerate() {
if let Some((nk, next)) = node
.children
.binary_search_by_key(&&k, |(k, n)| k)
.ok()
.and_then(|ix| Some(&node.children[ix]))
{
if let Some(value) = &next.value {
prefixes.push((i, value));
}
node = next;
} else {
break;
}
}
prefixes
}
pub fn iter_prefixes<I: Iterator<Item = K>>(
&mut self,
key: I,
mut cb: impl FnMut(usize, &mut TrieNode<K, V>),
) {
let mut node = &mut self.root;
for (i, k) in key.enumerate() {
if let Ok(ix) = node.children.binary_search_by_key(&&k, |(k, n)| k) {
let (nk, next) = &mut node.children[ix];
if let Some(_) = &mut next.value {
cb(i, next);
}
node = next;
} else {
cb(i, node);
break;
}
}
}
pub fn find_longest_prefix<I: Iterator<Item = K>>(&self, key: I) -> Option<&V> {
{
let mut current = &self.root;
let mut last_value: Option<&V> = None.as_ref();
for k in key {
if let Some((_, next_node)) = current.children.iter().find(|(key, _)| key == &k) {
if next_node.value.is_some() {
last_value = next_node.value.as_ref();
}
current = next_node;
} else {
break;
}
}
last_value
}
}
pub fn find_postfixes<I: Iterator<Item = K>>(&self, prefix: I) -> Vec<&V> {
let mut postfixes = Vec::new();
if let Some(node) = self.find_node(prefix) {
self.collect_values(node, &mut postfixes);
}
postfixes
}
#[allow(clippy::only_used_in_recursion)]
fn collect_values<'a>(&self, node: &'a TrieNode<K, V>, values: &mut Vec<&'a V>) {
if let Some(ref value) = node.value {
values.push(value);
}
for (_, child) in &node.children {
self.collect_values(child, values);
}
}
pub fn is_empty(&self) -> bool {
self.root.children.is_empty()
}
pub fn clear(&mut self) {
self.root = TrieNode::default();
}
pub fn insert<I: Iterator<Item = K>>(
&mut self,
key: I,
value_cb: impl FnMut(&mut TrieNode<K, V>, Option<usize>),
) -> Option<&mut V> {
self.root.insert(key.enumerate(), value_cb, None)
}
pub fn remove_subtree<I: Iterator<Item = K>>(&mut self, key: I) {
self.root.remove_subtree(key.peekable())
}
fn find_node<I: Iterator<Item = K>>(&self, key: I) -> Option<&TrieNode<K, V>> {
self.root.find_node(key)
}
fn find_node_mut<I: Iterator<Item = K>>(&mut self, key: I) -> Option<&mut TrieNode<K, V>> {
self.root.find_node_mut(key)
}
pub fn iter(&self) -> TrieIterator<K, V> {
TrieIterator::new(&self)
}
}
impl<T: Eq + Ord + Clone, U> Default for Trie<T, U> {
fn default() -> Self {
Self::new()
}
}
pub struct TrieIterator<'a, K: Eq + Ord + Clone, V> {
stack: Vec<(&'a TrieNode<K, V>, Vec<K>)>,
}
impl<'a, K: Eq + Ord + Clone, V> TrieIterator<'a, K, V> {
fn new(trie: &'a Trie<K, V>) -> Self {
TrieIterator {
stack: vec![(&trie.root, Vec::new())],
}
}
}
impl<'a, K: Eq + Ord + Clone, V> Iterator for TrieIterator<'a, K, V> {
type Item = (Vec<K>, &'a V);
fn next(&mut self) -> Option<Self::Item> {
while let Some((node, path)) = self.stack.pop() {
for (key_part, child) in &node.children {
let mut new_path = path.clone();
new_path.push(key_part.clone());
self.stack.push((child, new_path));
}
if let Some(ref value) = node.value {
return Some((path, value));
}
}
None
}
}