zerotrie/byte_phf/
builder.rs1use super::*;
6use crate::error::ZeroTrieBuildError;
7use alloc::vec;
8use alloc::vec::Vec;
9
10const MAX_L2_SEARCH_MISSES: usize = 24;
15
16#[allow(unused_labels)] #[allow(clippy::indexing_slicing)] pub fn find(bytes: &[u8]) -> Result<(u8, Vec<u8>), ZeroTrieBuildError> {
22 let n_usize = bytes.len();
23
24 let mut p = 0u8;
25 let mut qq = vec![0u8; n_usize];
26
27 let mut bqs = vec![0u8; n_usize];
28 let mut seen = vec![false; n_usize];
29 let max_allowable_p = P_FAST_MAX;
30 let mut max_allowable_q = Q_FAST_MAX;
31
32 #[allow(non_snake_case)]
33 let N = if n_usize > 0 && n_usize < 256 {
34 n_usize as u8
35 } else {
36 debug_assert!(n_usize == 0 || n_usize == 256);
37 return Ok((p, qq));
38 };
39
40 'p_loop: loop {
41 let mut buckets: Vec<(usize, Vec<u8>)> = (0..n_usize).map(|i| (i, vec![])).collect();
43 for byte in bytes {
44 let l1 = f1(*byte, p, N) as usize;
45 buckets[l1].1.push(*byte);
46 }
47 buckets.sort_by_key(|(_, v)| -(v.len() as isize));
48 let mut i = 0;
50 let mut num_max_q = 0;
51 bqs.fill(0);
52 seen.fill(false);
53 'q_loop: loop {
54 if i == buckets.len() {
56 for (local_j, real_j) in buckets.iter().map(|(j, _)| *j).enumerate() {
57 debug_assert!(local_j < n_usize); debug_assert!(real_j < n_usize); qq[real_j] = bqs[local_j];
60 }
61 return Ok((p, qq));
66 }
67 let mut bucket = buckets[i].1.as_slice();
68 'byte_loop: for (j, byte) in bucket.iter().enumerate() {
69 let l2 = f2(*byte, bqs[i], N) as usize;
70 if seen[l2] {
71 for k_byte in &bucket[0..j] {
73 let l2 = f2(*k_byte, bqs[i], N) as usize;
74 assert!(seen[l2]);
75 seen[l2] = false;
76 }
77 'reset_loop: loop {
78 if bqs[i] < max_allowable_q {
79 bqs[i] += 1;
80 continue 'q_loop;
81 }
82 num_max_q += 1;
83 bqs[i] = 0;
84 if i == 0 || num_max_q > MAX_L2_SEARCH_MISSES {
85 if p == max_allowable_p && max_allowable_q != Q_REAL_MAX {
86 max_allowable_q = Q_REAL_MAX;
88 p = 0;
89 continue 'p_loop;
90 } else if p == max_allowable_p {
91 debug_assert_eq!(max_allowable_p, P_REAL_MAX);
94 return Err(ZeroTrieBuildError::CouldNotSolvePerfectHash);
96 } else {
97 p += 1;
98 continue 'p_loop;
99 }
100 }
101 i -= 1;
102 bucket = buckets[i].1.as_slice();
103 for byte in bucket {
104 let l2 = f2(*byte, bqs[i], N) as usize;
105 assert!(seen[l2]);
106 seen[l2] = false;
107 }
108 }
109 } else {
110 let l2 = f2(*byte, bqs[i], N) as usize;
112 seen[l2] = true;
113 }
114 }
115 i += 1;
117 }
118 }
119}
120
121impl PerfectByteHashMap<Vec<u8>> {
122 #[allow(clippy::indexing_slicing)] pub fn try_new(keys: &[u8]) -> Result<Self, ZeroTrieBuildError> {
127 let n_usize = keys.len();
128 let n = n_usize as u8;
129 let (p, mut qq) = find(keys)?;
130 let mut keys_permuted = vec![0; n_usize];
131 for key in keys {
132 let l1 = f1(*key, p, n) as usize;
133 let q = qq[l1];
134 let l2 = f2(*key, q, n) as usize;
135 keys_permuted[l2] = *key;
136 }
137 let mut result = Vec::with_capacity(n_usize * 2 + 1);
138 result.push(p);
139 result.append(&mut qq);
140 result.append(&mut keys_permuted);
141 Ok(Self(result))
142 }
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148
149 extern crate std;
150 use std::print;
151 use std::println;
152
153 fn print_byte_to_stdout(byte: u8) {
154 let c = char::from(byte);
155 if c.is_ascii_alphanumeric() {
156 print!("'{c}'");
157 } else {
158 print!("0x{byte:X}");
159 }
160 }
161
162 fn random_alphanums(seed: u64, len: usize) -> Vec<u8> {
163 use rand::seq::SliceRandom;
164 use rand::SeedableRng;
165 let mut bytes: Vec<u8> =
166 b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789".into();
167 let mut rng = rand_pcg::Lcg64Xsh32::seed_from_u64(seed);
168 bytes.partial_shuffle(&mut rng, len).0.into()
169 }
170
171 #[test]
172 fn test_random_distributions() {
173 let mut p_distr = vec![0; 256];
174 let mut q_distr = vec![0; 256];
175 for len in 0..50 {
176 for seed in 0..50 {
177 let bytes = random_alphanums(seed, len);
178 let (p, qq) = find(bytes.as_slice()).unwrap();
179 p_distr[p as usize] += 1;
180 for q in qq {
181 q_distr[q as usize] += 1;
182 }
183 }
184 }
185 println!("p_distr: {p_distr:?}");
186 println!("q_distr: {q_distr:?}");
187
188 let fast_p = p_distr[0..=P_FAST_MAX as usize].iter().sum::<usize>();
189 let slow_p = p_distr[(P_FAST_MAX + 1) as usize..].iter().sum::<usize>();
190 let fast_q = q_distr[0..=Q_FAST_MAX as usize].iter().sum::<usize>();
191 let slow_q = q_distr[(Q_FAST_MAX + 1) as usize..].iter().sum::<usize>();
192
193 assert_eq!(2500, fast_p);
194 assert_eq!(0, slow_p);
195 assert_eq!(61243, fast_q);
196 assert_eq!(7, slow_q);
197
198 let bytes = random_alphanums(0, 16);
199
200 #[allow(non_snake_case)]
201 let N = u8::try_from(bytes.len()).unwrap();
202
203 let (p, qq) = find(bytes.as_slice()).unwrap();
204
205 println!("Results:");
206 for byte in bytes.iter() {
207 print_byte_to_stdout(*byte);
208 let l1 = f1(*byte, p, N) as usize;
209 let q = qq[l1];
210 let l2 = f2(*byte, q, N) as usize;
211 println!(" => l1 {l1} => q {q} => l2 {l2}");
212 }
213 }
214}