1#![cfg_attr(not(doctest), doc = include_str!("../README.md"))]
16#![deny(
17 missing_docs,
18 missing_debug_implementations,
19 unreachable_pub,
20 rustdoc::broken_intra_doc_links,
21 unsafe_code
22)]
23#![warn(rust_2018_idioms)]
24#![no_std]
25
26extern crate alloc;
27
28#[cfg(test)]
29extern crate std;
30
31mod u256;
32
33use alloc::{
34 collections::{BTreeMap, BTreeSet, VecDeque},
35 vec,
36 vec::Vec,
37};
38use core::ops::RangeFrom;
39
40use self::u256::U256;
41
42#[derive(Debug, Clone)]
43#[repr(transparent)]
44struct MasksByByteSized<I>([I; 256]);
45
46impl<I> Default for MasksByByteSized<I>
47where
48 I: Default + Copy,
49{
50 fn default() -> Self {
51 Self([I::default(); 256])
52 }
53}
54
55#[allow(clippy::large_enum_variant)]
56enum MasksByByte {
57 U8(MasksByByteSized<u8>),
58 U16(MasksByByteSized<u16>),
59 U32(MasksByByteSized<u32>),
60 U64(MasksByByteSized<u64>),
61 U128(MasksByByteSized<u128>),
62 U256(MasksByByteSized<U256>),
63}
64
65impl MasksByByte {
66 fn new(used_bytes: BTreeSet<u8>) -> Self {
67 match used_bytes.len() {
68 ..=8 => MasksByByte::U8(MasksByByteSized::<u8>::new(used_bytes)),
69 9..=16 => {
70 MasksByByte::U16(MasksByByteSized::<u16>::new(used_bytes))
71 }
72 17..=32 => {
73 MasksByByte::U32(MasksByByteSized::<u32>::new(used_bytes))
74 }
75 33..=64 => {
76 MasksByByte::U64(MasksByByteSized::<u64>::new(used_bytes))
77 }
78 65..=128 => {
79 MasksByByte::U128(MasksByByteSized::<u128>::new(used_bytes))
80 }
81 129..=256 => {
82 MasksByByte::U256(MasksByByteSized::<U256>::new(used_bytes))
83 }
84 _ => unreachable!("There are only 256 possible u8s"),
85 }
86 }
87}
88
89#[derive(Debug, Clone)]
92pub struct TrieHardSized<'a, T, I> {
93 masks: MasksByByteSized<I>,
94 nodes: Vec<TrieState<'a, T, I>>,
95}
96
97impl<'a, T, I> Default for TrieHardSized<'a, T, I>
98where
99 I: Default + Copy,
100{
101 fn default() -> Self {
102 Self {
103 masks: MasksByByteSized::default(),
104 nodes: Default::default(),
105 }
106 }
107}
108
109#[derive(PartialEq, Eq, PartialOrd, Ord)]
110struct StateSpec<'a> {
111 prefix: &'a [u8],
112 index: usize,
113}
114
115#[derive(Debug, Clone)]
116struct SearchNode<I> {
117 mask: I,
118 edge_start: usize,
119}
120
121#[derive(Debug, Clone)]
122enum TrieState<'a, T, I> {
123 Leaf(&'a [u8], T),
124 Search(SearchNode<I>),
125 SearchOrLeaf(&'a [u8], T, SearchNode<I>),
126}
127
128#[allow(clippy::large_enum_variant)]
150#[derive(Debug, Clone)]
151pub enum TrieHard<'a, T> {
152 U8(TrieHardSized<'a, T, u8>),
154 U16(TrieHardSized<'a, T, u16>),
156 U32(TrieHardSized<'a, T, u32>),
158 U64(TrieHardSized<'a, T, u64>),
160 U128(TrieHardSized<'a, T, u128>),
162 U256(TrieHardSized<'a, T, U256>),
164}
165
166impl<'a, T> Default for TrieHard<'a, T> {
167 fn default() -> Self {
168 TrieHard::U8(TrieHardSized::default())
169 }
170}
171
172impl<'a, T> TrieHard<'a, T>
173where
174 T: 'a + Copy,
175{
176 pub fn new(values: Vec<(&'a [u8], T)>) -> Self {
198 if values.is_empty() {
199 return Self::default();
200 }
201
202 let used_bytes = values
203 .iter()
204 .flat_map(|(k, _)| k.iter())
205 .cloned()
206 .collect::<BTreeSet<_>>();
207
208 let masks = MasksByByte::new(used_bytes);
209
210 match masks {
211 MasksByByte::U8(masks) => {
212 TrieHard::U8(TrieHardSized::<'_, _, u8>::new(masks, values))
213 }
214 MasksByByte::U16(masks) => {
215 TrieHard::U16(TrieHardSized::<'_, _, u16>::new(masks, values))
216 }
217 MasksByByte::U32(masks) => {
218 TrieHard::U32(TrieHardSized::<'_, _, u32>::new(masks, values))
219 }
220 MasksByByte::U64(masks) => {
221 TrieHard::U64(TrieHardSized::<'_, _, u64>::new(masks, values))
222 }
223 MasksByByte::U128(masks) => {
224 TrieHard::U128(TrieHardSized::<'_, _, u128>::new(masks, values))
225 }
226 MasksByByte::U256(masks) => {
227 TrieHard::U256(TrieHardSized::<'_, _, U256>::new(masks, values))
228 }
229 }
230 }
231
232 pub fn get<K: AsRef<[u8]>>(&self, raw_key: K) -> Option<T> {
246 match self {
247 TrieHard::U8(trie) => trie.get(raw_key),
248 TrieHard::U16(trie) => trie.get(raw_key),
249 TrieHard::U32(trie) => trie.get(raw_key),
250 TrieHard::U64(trie) => trie.get(raw_key),
251 TrieHard::U128(trie) => trie.get(raw_key),
252 TrieHard::U256(trie) => trie.get(raw_key),
253 }
254 }
255
256 pub fn get_from_bytes(&self, key: &[u8]) -> Option<T> {
268 match self {
269 TrieHard::U8(trie) => trie.get_from_bytes(key),
270 TrieHard::U16(trie) => trie.get_from_bytes(key),
271 TrieHard::U32(trie) => trie.get_from_bytes(key),
272 TrieHard::U64(trie) => trie.get_from_bytes(key),
273 TrieHard::U128(trie) => trie.get_from_bytes(key),
274 TrieHard::U256(trie) => trie.get_from_bytes(key),
275 }
276 }
277
278 pub fn iter(&self) -> TrieIter<'_, 'a, T> {
293 match self {
294 TrieHard::U8(trie) => TrieIter::U8(trie.iter()),
295 TrieHard::U16(trie) => TrieIter::U16(trie.iter()),
296 TrieHard::U32(trie) => TrieIter::U32(trie.iter()),
297 TrieHard::U64(trie) => TrieIter::U64(trie.iter()),
298 TrieHard::U128(trie) => TrieIter::U128(trie.iter()),
299 TrieHard::U256(trie) => TrieIter::U256(trie.iter()),
300 }
301 }
302
303 pub fn prefix_search<K: AsRef<[u8]>>(
318 &self,
319 prefix: K,
320 ) -> TrieIter<'_, 'a, T> {
321 match self {
322 TrieHard::U8(trie) => TrieIter::U8(trie.prefix_search(prefix)),
323 TrieHard::U16(trie) => TrieIter::U16(trie.prefix_search(prefix)),
324 TrieHard::U32(trie) => TrieIter::U32(trie.prefix_search(prefix)),
325 TrieHard::U64(trie) => TrieIter::U64(trie.prefix_search(prefix)),
326 TrieHard::U128(trie) => TrieIter::U128(trie.prefix_search(prefix)),
327 TrieHard::U256(trie) => TrieIter::U256(trie.prefix_search(prefix)),
328 }
329 }
330
331 pub fn ancestor<K: AsRef<[u8]>>(&self, key: K) -> Option<(&[u8], T)> {
350 match self {
351 TrieHard::U8(trie) => trie.ancestor(key),
352 TrieHard::U16(trie) => trie.ancestor(key),
353 TrieHard::U32(trie) => trie.ancestor(key),
354 TrieHard::U64(trie) => trie.ancestor(key),
355 TrieHard::U128(trie) => trie.ancestor(key),
356 TrieHard::U256(trie) => trie.ancestor(key),
357 }
358 }
359}
360
361#[derive(Debug)]
363pub enum TrieIter<'b, 'a, T> {
364 U8(TrieIterSized<'b, 'a, T, u8>),
366 U16(TrieIterSized<'b, 'a, T, u16>),
368 U32(TrieIterSized<'b, 'a, T, u32>),
370 U64(TrieIterSized<'b, 'a, T, u64>),
372 U128(TrieIterSized<'b, 'a, T, u128>),
374 U256(TrieIterSized<'b, 'a, T, U256>),
376}
377
378#[derive(Debug, Default)]
379struct TrieNodeIter {
380 node_index: usize,
381 stage: TrieNodeIterStage,
382}
383
384#[derive(Debug, Default)]
385enum TrieNodeIterStage {
386 #[default]
387 Inner,
388 Child(usize, usize),
389}
390
391#[derive(Debug)]
394pub struct TrieIterSized<'b, 'a, T, I> {
395 stack: Vec<TrieNodeIter>,
396 trie: &'b TrieHardSized<'a, T, I>,
397}
398
399impl<'b, 'a, T, I> TrieIterSized<'b, 'a, T, I> {
400 fn empty(trie: &'b TrieHardSized<'a, T, I>) -> Self {
401 Self {
402 stack: Default::default(),
403 trie,
404 }
405 }
406
407 fn new(trie: &'b TrieHardSized<'a, T, I>, node_index: usize) -> Self {
408 Self {
409 stack: vec![TrieNodeIter {
410 node_index,
411 stage: Default::default(),
412 }],
413 trie,
414 }
415 }
416}
417
418impl<'b, 'a, T> Iterator for TrieIter<'b, 'a, T>
419where
420 T: Copy,
421{
422 type Item = (&'a [u8], T);
423
424 fn next(&mut self) -> Option<Self::Item> {
425 match self {
426 TrieIter::U8(iter) => iter.next(),
427 TrieIter::U16(iter) => iter.next(),
428 TrieIter::U32(iter) => iter.next(),
429 TrieIter::U64(iter) => iter.next(),
430 TrieIter::U128(iter) => iter.next(),
431 TrieIter::U256(iter) => iter.next(),
432 }
433 }
434}
435
436impl<'a, T> FromIterator<&'a T> for TrieHard<'a, &'a T>
437where
438 T: 'a + AsRef<[u8]> + ?Sized,
439{
440 fn from_iter<I: IntoIterator<Item = &'a T>>(values: I) -> Self {
441 let values = values
442 .into_iter()
443 .map(|v| (v.as_ref(), v))
444 .collect::<Vec<_>>();
445
446 Self::new(values)
447 }
448}
449
450macro_rules! trie_impls {
451 ($($int_type:ty),+) => {
452 $(
453 trie_impls!(_impl $int_type);
454 )+
455 };
456
457 (_impl $int_type:ty) => {
458
459 impl SearchNode<$int_type> {
460 fn evaluate<T>(&self, c: u8, trie: &TrieHardSized<'_, T, $int_type>) -> Option<usize> {
461 let c_mask = trie.masks.0[c as usize];
462 let mask_res = self.mask & c_mask;
463 (mask_res > 0).then(|| {
464 let smaller_bits = mask_res - 1;
465 let smaller_bits_mask = smaller_bits & self.mask;
466 let index_offset = smaller_bits_mask.count_ones() as usize;
467 self.edge_start + index_offset
468 })
469 }
470 }
471
472 impl<'a, T> TrieHardSized<'a, T, $int_type>
473 where
474 T: Copy
475 {
476
477 pub fn get<K: AsRef<[u8]>>(&self, key: K) -> Option<T> {
495 self.get_from_bytes(key.as_ref())
496 }
497
498 pub fn get_from_bytes(&self, key: &[u8]) -> Option<T> {
514 let mut state = self.nodes.get(0)?;
515
516 for (i, c) in key.iter().enumerate() {
517
518 let next_state_opt = match state {
519 TrieState::Leaf(k, value) => {
520 return (
521 k.len() == key.len()
522 && k[i..] == key[i..]
523 ).then_some(*value)
524 }
525 TrieState::Search(search)
526 | TrieState::SearchOrLeaf(_, _, search) => {
527 search.evaluate(*c, self)
528 }
529 };
530
531 if let Some(next_state_index) = next_state_opt {
532 state = &self.nodes[next_state_index];
533 } else {
534 return None;
535 }
536 }
537
538 if let TrieState::Leaf(k, value)
539 | TrieState::SearchOrLeaf(k, value, _) = state
540 {
541 (k.len() == key.len()).then_some(*value)
542 } else {
543 None
544 }
545 }
546
547 pub fn iter(&self) -> TrieIterSized<'_, 'a, T, $int_type> {
566 TrieIterSized {
567 stack: vec![TrieNodeIter::default()],
568 trie: self
569 }
570 }
571
572
573 pub fn prefix_search<K: AsRef<[u8]>>(&self, prefix: K) -> TrieIterSized<'_, 'a, T, $int_type> {
592 let key = prefix.as_ref();
593 let mut node_index = 0;
594 let Some(mut state) = self.nodes.get(node_index) else {
595 return TrieIterSized::empty(self);
596 };
597
598 for (i, c) in key.iter().enumerate() {
599 let next_state_opt = match state {
600 TrieState::Leaf(k, _) => {
601 if k.len() == key.len() && k[i..] == key[i..] {
602 return TrieIterSized::new(self, node_index);
603 } else {
604 return TrieIterSized::empty(self);
605 }
606 }
607 TrieState::Search(search)
608 | TrieState::SearchOrLeaf(_, _, search) => {
609 search.evaluate(*c, self)
610 }
611 };
612
613 if let Some(next_state_index) = next_state_opt {
614 node_index = next_state_index;
615 state = &self.nodes[next_state_index];
616 } else {
617 return TrieIterSized::empty(self);
618 }
619 }
620
621 TrieIterSized::new(self, node_index)
622 }
623
624 pub fn ancestor<K: AsRef<[u8]>>(
647 &self,
648 key: K,
649 ) -> Option<(&[u8], T)> {
650 self.ancestor_recurse(0, key.as_ref(), self.nodes.get(0)?)
651 }
652
653 fn ancestor_recurse(
654 &self,
655 i: usize,
656 key: &[u8],
657 state: &TrieState<'a, T, $int_type>,
658 ) -> Option<(&[u8], T)> {
659 match state {
660 TrieState::Leaf(k, value) => {
661 (
662 k.len() <= key.len()
663 && k[i..] == key[i..k.len()]
664 ).then_some((k, *value))
665 }
666 TrieState::Search(search) => {
667 let c = key.get(i)?;
668 let next_state_index = search.evaluate(*c, self)?;
669 self.ancestor_recurse(i + 1, key, &self.nodes[next_state_index])
670 }
671 TrieState::SearchOrLeaf(k, value, search) => {
672 let search = || {
674 let c = key.get(i)?;
675 let next_state_index = search.evaluate(*c, self)?;
676 self.ancestor_recurse(i + 1, key, &self.nodes[next_state_index])
677 };
678
679 search().or_else(|| {
680 (
681 k.len() <= key.len()
682 && k[i..] == key[i..k.len()]
683 ).then_some((k, *value))
684 })
685 }
686 }
687 }
688 }
689
690 impl<'a, T> TrieHardSized<'a, T, $int_type> where T: 'a + Copy {
691 fn new(masks: MasksByByteSized<$int_type>, values: Vec<(&'a [u8], T)>) -> Self {
692 let values = values.into_iter().collect::<Vec<_>>();
693 let sorted = values
694 .iter()
695 .map(|(k, v)| (*k, *v))
696 .collect::<BTreeMap<_, _>>();
697
698 let mut nodes = Vec::new();
699 let mut next_index = 1;
700
701 let root_state_spec = StateSpec {
702 prefix: &[],
703 index: 0,
704 };
705
706 let mut spec_queue = VecDeque::new();
707 spec_queue.push_back(root_state_spec);
708
709 while let Some(spec) = spec_queue.pop_front() {
710 debug_assert_eq!(spec.index, nodes.len());
711 let (state, next_specs) = TrieState::<'_, _, $int_type>::new(
712 spec,
713 next_index,
714 &masks.0,
715 &sorted,
716 );
717
718 next_index += next_specs.len();
719 spec_queue.extend(next_specs);
720 nodes.push(state);
721 }
722
723 TrieHardSized {
724 nodes,
725 masks,
726 }
727 }
728 }
729
730
731 impl <'a, T> TrieState<'a, T, $int_type> where T: 'a + Copy {
732 fn new(
733 spec: StateSpec<'a>,
734 edge_start: usize,
735 byte_masks: &[$int_type; 256],
736 sorted: &BTreeMap<&'a [u8], T>,
737 ) -> (Self, Vec<StateSpec<'a>>) {
738 let StateSpec { prefix, .. } = spec;
739
740 let prefix_len = prefix.len();
741 let next_prefix_len = prefix_len + 1;
742
743 let mut prefix_match = None;
744 let mut children_seen = 0;
745 let mut last_seen = None;
746
747 let next_states_paired = sorted
748 .range(RangeFrom { start: prefix })
749 .take_while(|(key, _)| key.starts_with(prefix))
750 .filter_map(|(key, val)| {
751 children_seen += 1;
752 last_seen = Some((key, *val));
753
754 if *key == prefix {
755 prefix_match = Some((key, *val));
756 None
757 } else {
758 let next_c = key.get(prefix_len).unwrap();
761 let next_prefix = &key[..next_prefix_len];
762
763 Some((
764 *next_c,
765 StateSpec {
766 prefix: next_prefix,
767 index: 0,
768 },
769 ))
770 }
771 })
772 .collect::<BTreeMap<_, _>>()
773 .into_iter()
774 .collect::<Vec<_>>();
775
776 let (last_k, last_v) = last_seen.unwrap();
779
780 if children_seen == 1 {
781 return (TrieState::Leaf(last_k, last_v), vec![]);
782 }
783
784 if next_states_paired.is_empty() {
786 return (TrieState::Leaf(last_k, last_v), vec![], );
787 }
788
789 let mut mask = Default::default();
790
791 let next_state_specs = next_states_paired
793 .into_iter()
794 .enumerate()
795 .map(|(i, (c, mut next_state))| {
796 let next_node = edge_start + i;
797 next_state.index = next_node;
798 mask |= byte_masks[c as usize];
799 next_state
800 })
801 .collect();
802
803 let search_node = SearchNode { mask, edge_start };
804 let state = match prefix_match {
805 Some((key, value)) => {
806 TrieState::SearchOrLeaf(key, value, search_node)
807 }
808 _ => TrieState::Search(search_node),
809 };
810
811 (state, next_state_specs)
812 }
813 }
814
815 impl MasksByByteSized<$int_type> {
816 fn new(used_bytes: BTreeSet<u8>) -> Self {
817 let mut mask = Default::default();
818 mask += 1;
819
820 let mut byte_masks = [Default::default(); 256];
821
822 for c in used_bytes.into_iter() {
823 byte_masks[c as usize] = mask;
824 mask <<= 1;
825
826 }
827
828 Self(byte_masks)
829 }
830 }
831
832 impl <'b, 'a, T> Iterator for TrieIterSized<'b, 'a, T, $int_type>
833 where
834 T: Copy
835 {
836 type Item = (&'a [u8], T);
837
838 fn next(&mut self) -> Option<Self::Item> {
839
840 use TrieState as T;
841 use TrieNodeIterStage as S;
842
843 while let Some((node, node_index, stage)) = self.stack.pop()
844 .and_then(|TrieNodeIter { node_index, stage }| {
845 self.trie.nodes.get(node_index).map(|node| (node, node_index, stage))
846 })
847 {
848 match (node, stage) {
849 (T::Leaf(key, value), S::Inner) => return Some((*key, *value)),
850 (T::SearchOrLeaf(key, value, search), S::Inner) => {
851 self.stack.push(TrieNodeIter {
852 node_index,
853 stage: TrieNodeIterStage::Child(0, search.mask.count_ones() as usize)
854 });
855 self.stack.push(TrieNodeIter {
856 node_index: search.edge_start,
857 stage: Default::default()
858 });
859 return Some((*key, *value));
860 }
861 (T::Search(search), S::Inner) => {
862 self.stack.push(TrieNodeIter {
863 node_index,
864 stage: TrieNodeIterStage::Child(0, search.mask.count_ones() as usize)
865 });
866 self.stack.push(TrieNodeIter {
867 node_index: search.edge_start,
868 stage: Default::default()
869 });
870 }
871 (
872 T::SearchOrLeaf(_, _, search) | T::Search(search),
873 S::Child(mut child, child_count)
874 ) => {
875 child += 1;
876 if child < child_count {
877 self.stack.push(TrieNodeIter {
878 node_index,
879 stage: TrieNodeIterStage::Child(child, child_count)
880 });
881 self.stack.push(TrieNodeIter {
882 node_index: search.edge_start + child,
883 stage: Default::default()
884 });
885 }
886 }
887 _ => unreachable!()
888 }
889 }
890
891 None
892 }
893 }
894 }
895}
896
897trie_impls! {u8, u16, u32, u64, u128, U256}
898
899#[cfg(test)]
900mod tests {
901 use rstest::rstest;
902
903 use super::*;
904
905 #[test]
906 fn test_trivial() {
907 let empty: Vec<&str> = vec![];
908 let empty_trie = empty.iter().collect::<TrieHard<'_, _>>();
909
910 assert_eq!(None, empty_trie.get("anything"));
911 }
912
913 #[rstest]
914 #[case("", Some(""))]
915 #[case("a", Some("a"))]
916 #[case("ab", Some("ab"))]
917 #[case("abc", None)]
918 #[case("aac", Some("aac"))]
919 #[case("aa", None)]
920 #[case("aab", None)]
921 #[case("adddd", Some("adddd"))]
922 fn test_small_get(#[case] key: &str, #[case] expected: Option<&str>) {
923 let trie = ["", "a", "ab", "aac", "adddd", "addde"]
924 .into_iter()
925 .collect::<TrieHard<'_, _>>();
926 assert_eq!(expected, trie.get(key));
927 }
928
929 #[test]
930 fn test_skip_to_leaf() {
931 let trie = ["a", "aa", "aaa"].into_iter().collect::<TrieHard<'_, _>>();
932
933 assert_eq!(trie.get("aa"), Some("aa"))
934 }
935
936 #[rstest]
937 #[case(8)]
938 #[case(16)]
939 #[case(32)]
940 #[case(64)]
941 #[case(128)]
942 #[case(256)]
943 fn test_sizes(#[case] bits: usize) {
944 let range = 0..bits;
945 let bytes = range.map(|b| [b as u8]).collect::<Vec<_>>();
946 let trie = bytes.iter().collect::<TrieHard<'_, _>>();
947
948 use TrieHard as T;
949
950 match (bits, trie) {
951 (8, T::U8(_)) => (),
952 (16, T::U16(_)) => (),
953 (32, T::U32(_)) => (),
954 (64, T::U64(_)) => (),
955 (128, T::U128(_)) => (),
956 (256, T::U256(_)) => (),
957 _ => panic!("Mismatched trie sizes"),
958 }
959 }
960
961 #[rstest]
962 #[case(include_str!("../data/1984.txt"))]
963 #[case(include_str!("../data/sun-rising.txt"))]
964 fn test_full_text(#[case] text: &str) {
965 let words: Vec<&str> =
966 text.split(|c: char| c.is_whitespace()).collect();
967 let trie: TrieHard<'_, _> = words.iter().copied().collect();
968
969 let unique_words = words
970 .into_iter()
971 .collect::<BTreeSet<_>>()
972 .into_iter()
973 .collect::<Vec<_>>();
974
975 for word in &unique_words {
976 assert!(trie.get(word).is_some())
977 }
978
979 assert_eq!(
980 unique_words,
981 trie.iter().map(|(_, v)| v).collect::<Vec<_>>()
982 );
983 }
984
985 #[test]
986 fn test_unicode() {
987 let trie: TrieHard<'_, _> = ["bär", "bären"].into_iter().collect();
988
989 assert_eq!(trie.get("bär"), Some("bär"));
990 assert_eq!(trie.get("bä"), None);
991 assert_eq!(trie.get("bären"), Some("bären"));
992 assert_eq!(trie.get("bärën"), None);
993 }
994
995 #[rstest]
996 #[case(&[], &[])]
997 #[case(&[""], &[""])]
998 #[case(&["aaa", "a", ""], &["", "a", "aaa"])]
999 #[case(&["aaa", "a", ""], &["", "a", "aaa"])]
1000 #[case(&["", "a", "ab", "aac", "adddd", "addde"], &["", "a", "aac", "ab", "adddd", "addde"])]
1001 fn test_iter(#[case] input: &[&str], #[case] output: &[&str]) {
1002 let trie = input.iter().copied().collect::<TrieHard<'_, _>>();
1003 let emitted = trie.iter().map(|(_, v)| v).collect::<Vec<_>>();
1004 assert_eq!(emitted, output);
1005 }
1006
1007 #[rstest]
1008 #[case(&[], "", &[])]
1009 #[case(&[""], "", &[""])]
1010 #[case(&["aaa", "a", ""], "", &["", "a", "aaa"])]
1011 #[case(&["aaa", "a", ""], "a", &["a", "aaa"])]
1012 #[case(&["aaa", "a", ""], "aa", &["aaa"])]
1013 #[case(&["aaa", "a", ""], "aab", &[])]
1014 #[case(&["aaa", "a", ""], "aaa", &["aaa"])]
1015 #[case(&["aaa", "a", ""], "b", &[])]
1016 #[case(&["dad", "ant", "and", "dot", "do"], "d", &["dad", "do", "dot"])]
1017 fn test_prefix_search(
1018 #[case] input: &[&str],
1019 #[case] prefix: &str,
1020 #[case] output: &[&str],
1021 ) {
1022 let trie = input.iter().copied().collect::<TrieHard<'_, _>>();
1023 let emitted = trie
1024 .prefix_search(prefix)
1025 .map(|(_, v)| v)
1026 .collect::<Vec<_>>();
1027 assert_eq!(emitted, output);
1028 }
1029
1030 #[rstest]
1031 #[case(&[], "", None)]
1032 #[case(&[""], "", Some(""))]
1033 #[case(&["aaa", "a", ""], "", Some(""))]
1034 #[case(&["aaa", "a", ""], "a", Some("a"))]
1035 #[case(&["aaa", "a", ""], "aa", Some("a"))]
1036 #[case(&["aaa", "a", ""], "aab", Some("a"))]
1037 #[case(&["aaa", "a", ""], "aaa", Some("aaa"))]
1038 #[case(&["aaa", "a", ""], "b", Some(""))]
1039 #[case(&["dad", "ant", "and", "dot", "do"], "d", None)]
1040 #[case(&["dad", "ant", "and", "dot", "do"], "dad", Some("dad"))]
1041 #[case(&["dad", "ant", "and", "dot", "do"], "dada", Some("dad"))]
1042 #[case(&["dad", "ant", "and", "dot", "do"], "do", Some("do"))]
1043 #[case(&["dad", "ant", "and", "dot", "do"], "dot", Some("dot"))]
1044 #[case(&["dad", "ant", "and", "dot", "do"], "dob", Some("do"))]
1045 #[case(&["dad", "ant", "and", "dot", "do"], "doto", Some("dot"))]
1046 fn test_ancestor(
1047 #[case] input: &[&str],
1048 #[case] key: &str,
1049 #[case] output: Option<&str>,
1050 ) {
1051 let trie = input.iter().copied().collect::<TrieHard<'_, _>>();
1052 let emitted = trie.ancestor(key).map(|(_, v)| v);
1053 assert_eq!(emitted, output);
1054 }
1055}