Skip to main content

trie_hard/
lib.rs

1// Copyright 2024 Cloudflare, Inc.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7// http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15#![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/// Inner representation of a trie-hard trie that is generic to a specific size
90/// of integer.
91#[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/// Enumeration of all the possible sizes of trie-hard tries. An instance of
129/// this enum can be created from any set of arbitrary string or byte slices.
130/// The variant returned will depend on the number of distinct bytes contained
131/// in the set.
132///
133/// ```
134/// # use trie_hard::TrieHard;
135/// let trie = ["and", "ant", "dad", "do", "dot"]
136///     .into_iter()
137///     .collect::<TrieHard<'_, _>>();
138///
139/// assert!(trie.get("dad").is_some());
140/// assert!(trie.get("do").is_some());
141/// assert!(trie.get("don't").is_none());
142/// ```
143///
144/// _Note_: This enum has a very large variant which dominates the size for
145/// the enum. That means that a small trie using `u8`s for storage will take up
146/// way (32x) more storage than it needs to. If you are concerned about extra
147/// space (and you know ahead of time the trie size needed), you should extract
148/// the inner, `[TrieHardSized]` which will use only the size required.
149#[allow(clippy::large_enum_variant)]
150#[derive(Debug, Clone)]
151pub enum TrieHard<'a, T> {
152    /// Trie-hard using u8s for storage. For sets with 1..=8 unique bytes
153    U8(TrieHardSized<'a, T, u8>),
154    /// Trie-hard using u16s for storage. For sets with 9..=16 unique bytes
155    U16(TrieHardSized<'a, T, u16>),
156    /// Trie-hard using u32s for storage. For sets with 17..=32 unique bytes
157    U32(TrieHardSized<'a, T, u32>),
158    /// Trie-hard using u64s for storage. For sets with 33..=64 unique bytes
159    U64(TrieHardSized<'a, T, u64>),
160    /// Trie-hard using u128s for storage. For sets with 65..=126 unique bytes
161    U128(TrieHardSized<'a, T, u128>),
162    /// Trie-hard using U256s for storage. For sets with 129.. unique bytes
163    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    /// Create an instance of a trie-hard trie with the given keys and values.
177    /// The variant returned will be determined based on the number of unique
178    /// bytes in the keys.
179    ///
180    /// ```
181    /// # use trie_hard::TrieHard;
182    /// let trie = TrieHard::new(vec![
183    ///     (b"and", 0),
184    ///     (b"ant", 1),
185    ///     (b"dad", 2),
186    ///     (b"do", 3),
187    ///     (b"dot", 4)
188    /// ]);
189    ///
190    /// // Only 5 unique characters produces a u8 TrieHard
191    /// assert!(matches!(trie, TrieHard::U8(_)));
192    ///
193    /// assert_eq!(trie.get("dad"), Some(2));
194    /// assert_eq!(trie.get("do"), Some(3));
195    /// assert!(trie.get("don't").is_none());
196    /// ```
197    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    /// Get the value stored for the given key. Any key type can be used here as
233    /// long as the type implements `AsRef<[u8]>`. The byte slice referenced
234    /// will serve as the actual key.
235    /// ```
236    /// # use trie_hard::TrieHard;
237    /// let trie = ["and", "ant", "dad", "do", "dot"]
238    ///     .into_iter()
239    ///     .collect::<TrieHard<'_, _>>();
240    ///
241    /// assert!(trie.get("dad".to_owned()).is_some());
242    /// assert!(trie.get(b"do").is_some());
243    /// assert!(trie.get(b"don't".to_vec()).is_none());
244    /// ```
245    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    /// Get the value stored for the given byte-slice key
257    /// ```
258    /// # use trie_hard::TrieHard;
259    /// let trie = ["and", "ant", "dad", "do", "dot"]
260    ///     .into_iter()
261    ///     .collect::<TrieHard<'_, _>>();
262    ///
263    /// assert!(trie.get_from_bytes(b"dad").is_some());
264    /// assert!(trie.get_from_bytes(b"do").is_some());
265    /// assert!(trie.get_from_bytes(b"don't").is_none());
266    /// ```
267    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    /// Create an iterator over the entire trie. Emitted items will be ordered
279    /// by their keys
280    ///
281    /// ```
282    /// # use trie_hard::TrieHard;
283    /// let trie = ["dad", "ant", "and", "dot", "do"]
284    ///     .into_iter()
285    ///     .collect::<TrieHard<'_, _>>();
286    ///
287    /// assert_eq!(
288    ///     trie.iter().map(|(_, v)| v).collect::<Vec<_>>(),
289    ///     ["and", "ant", "dad", "do", "dot"]
290    /// );
291    /// ```
292    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    /// Create an iterator over the portion of the trie starting with the given
304    /// prefix
305    ///
306    /// ```
307    /// # use trie_hard::TrieHard;
308    /// let trie = ["dad", "ant", "and", "dot", "do"]
309    ///     .into_iter()
310    ///     .collect::<TrieHard<'_, _>>();
311    ///
312    /// assert_eq!(
313    ///     trie.prefix_search("d").map(|(_, v)| v).collect::<Vec<_>>(),
314    ///     ["dad", "do", "dot"]
315    /// );
316    /// ```
317    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    /// Find the closest ancestor to the given key, where an ancestor is defined as the longest
332    /// string present in the trie that appears as a prefix of the given key.
333    ///
334    /// ```
335    /// # use trie_hard::TrieHard;
336    /// let trie = ["dad", "ant", "and", "dot", "do"]
337    ///     .into_iter()
338    ///     .collect::<TrieHard<'_, _>>();
339    ///
340    /// assert_eq!(
341    ///     trie.ancestor("dada").map(|(_, v)| v),
342    ///     Some("dad")
343    /// );
344    /// assert_eq!(
345    ///     trie.ancestor("an").map(|(_, v)| v),
346    ///     None
347    /// );
348    /// ```
349    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/// Structure used for iterative over the contents of trie
362#[derive(Debug)]
363pub enum TrieIter<'b, 'a, T> {
364    /// Variant for iterating over trie-hard tries built on u8
365    U8(TrieIterSized<'b, 'a, T, u8>),
366    /// Variant for iterating over trie-hard tries built on u16
367    U16(TrieIterSized<'b, 'a, T, u16>),
368    /// Variant for iterating over trie-hard tries built on u32
369    U32(TrieIterSized<'b, 'a, T, u32>),
370    /// Variant for iterating over trie-hard tries built on u64
371    U64(TrieIterSized<'b, 'a, T, u64>),
372    /// Variant for iterating over trie-hard tries built on u128
373    U128(TrieIterSized<'b, 'a, T, u128>),
374    /// Variant for iterating over trie-hard tries built on u256
375    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/// Structure for iterating of a trie-hard trie built on specific a specific
392/// integer size
393#[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            /// Get the value stored for the given key. Any key type can be used
478            /// here as long as the type implements `AsRef<[u8]>`. The byte slice
479            /// referenced will serve as the actual key.
480            /// ```
481            /// # use trie_hard::TrieHard;
482            /// let trie = ["and", "ant", "dad", "do", "dot"]
483            ///     .into_iter()
484            ///     .collect::<TrieHard<'_, _>>();
485            ///
486            /// let TrieHard::U8(sized_trie) = trie else {
487            ///     unreachable!()
488            /// };
489            ///
490            /// assert!(sized_trie.get("dad".to_owned()).is_some());
491            /// assert!(sized_trie.get(b"do").is_some());
492            /// assert!(sized_trie.get(b"don't".to_vec()).is_none());
493            /// ```
494            pub fn get<K: AsRef<[u8]>>(&self, key: K) -> Option<T> {
495                self.get_from_bytes(key.as_ref())
496            }
497
498            /// Get the value stored for the given byte-slice key.
499            /// ```
500            /// # use trie_hard::TrieHard;
501            /// let trie = ["and", "ant", "dad", "do", "dot"]
502            ///     .into_iter()
503            ///     .collect::<TrieHard<'_, _>>();
504            ///
505            /// let TrieHard::U8(sized_trie) = trie else {
506            ///     unreachable!()
507            /// };
508            ///
509            /// assert!(sized_trie.get_from_bytes(b"dad").is_some());
510            /// assert!(sized_trie.get_from_bytes(b"do").is_some());
511            /// assert!(sized_trie.get_from_bytes(b"don't").is_none());
512            /// ```
513            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            /// Create an iterator over the entire trie. Emitted items will be
548            /// ordered by their keys
549            ///
550            /// ```
551            /// # use trie_hard::TrieHard;
552            /// let trie = ["dad", "ant", "and", "dot", "do"]
553            ///     .into_iter()
554            ///     .collect::<TrieHard<'_, _>>();
555            ///
556            /// let TrieHard::U8(sized_trie) = trie else {
557            ///     unreachable!()
558            /// };
559            ///
560            /// assert_eq!(
561            ///     sized_trie.iter().map(|(_, v)| v).collect::<Vec<_>>(),
562            ///     ["and", "ant", "dad", "do", "dot"]
563            /// );
564            /// ```
565            pub fn iter(&self) -> TrieIterSized<'_, 'a, T, $int_type> {
566                TrieIterSized {
567                    stack: vec![TrieNodeIter::default()],
568                    trie: self
569                }
570            }
571
572
573            /// Create an iterator over the portion of the trie starting with the given
574            /// prefix
575            ///
576            /// ```
577            /// # use trie_hard::TrieHard;
578            /// let trie = ["dad", "ant", "and", "dot", "do"]
579            ///     .into_iter()
580            ///     .collect::<TrieHard<'_, _>>();
581            ///
582            /// let TrieHard::U8(sized_trie) = trie else {
583            ///     unreachable!()
584            /// };
585            ///
586            /// assert_eq!(
587            ///     sized_trie.prefix_search("d").map(|(_, v)| v).collect::<Vec<_>>(),
588            ///     ["dad", "do", "dot"]
589            /// );
590            /// ```
591            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            /// Find the closest ancestor to the given key, where an ancestor is defined as the
625            /// longest string present in the trie that appears as a prefix of the given key.
626            ///
627            /// ```
628            /// # use trie_hard::TrieHard;
629            /// let trie = ["dad", "ant", "and", "dot", "do"]
630            ///     .into_iter()
631            ///     .collect::<TrieHard<'_, _>>();
632            ///
633            /// let TrieHard::U8(sized_trie) = trie else {
634            ///     unreachable!()
635            /// };
636            ///
637            /// assert_eq!(
638            ///     sized_trie.ancestor("dada").map(|(_, v)| v),
639            ///     Some("dad")
640            /// );
641            /// assert_eq!(
642            ///     sized_trie.ancestor("an").map(|(_, v)| v),
643            ///     None
644            /// );
645            /// ```
646            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                        // lambda to enable using `?` operator
673                        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                            // Safety: The byte at prefix_len must exist otherwise we
759                            // would have ended up in the other branch of this statement
760                            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                // Safety: last_seen will be present because we saw at least one
777                //         entry must be present for this function to be called
778                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                // No next_states means we hit a leaf node
785                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                // Update the index for the next state now that we have ordered by
792                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}