1use std::time::{Duration, Instant};
24
25use arrow::array::{Array, AsArray};
26use arrow::datatypes::{
27 DataType, Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, UInt16Type, UInt32Type,
28 UInt64Type, UInt8Type,
29};
30use arrow::record_batch::RecordBatch;
31
32const FIRST: usize = 512;
33const GROWTH: usize = 4;
34const SAMPLE_MAX: usize = 65_536;
35const EXACT: usize = 4_096;
36const SEED: u64 = 0xA6_5EED;
37const SETTLED_F2: u64 = 8;
38const NULL_KEY: u64 = 0x6E75_6C6C_6B65_7931;
40
41#[derive(Debug, Clone, Copy, PartialEq)]
43#[non_exhaustive]
44pub struct GroupEstimate {
45 pub estimate: u64,
47 pub low: u64,
49 pub high: u64,
51 pub sample_rows: usize,
53 pub rows: usize,
55 pub exact: bool,
57 pub time: Duration,
59}
60
61fn splitmix(mut z: u64) -> u64 {
62 z = z.wrapping_add(0x9E37_79B9_7F4A_7C15);
63 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
64 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
65 z ^ (z >> 31)
66}
67
68fn fnv1a(b: &[u8]) -> u64 {
69 let mut h: u64 = 0xCBF2_9CE4_8422_2325;
70 for &x in b {
71 h ^= x as u64;
72 h = h.wrapping_mul(0x0100_0000_01B3);
73 }
74 h
75}
76
77fn positions(rows: usize, n: usize) -> Vec<usize> {
79 let stride = (rows / n).max(1);
80 (0..n)
81 .map(|i| i * stride + (splitmix(i as u64 + SEED) % stride as u64) as usize)
82 .filter(|&p| p < rows)
83 .collect()
84}
85
86fn fold_column(a: &dyn Array, idx: &[usize], out: &mut [u64], first: bool) {
91 fn put(out: &mut [u64], j: usize, v: u64, first: bool) {
92 let v = splitmix(v);
93 out[j] = if first { v } else { out[j].rotate_left(23).wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ v };
94 }
95 let nulls = a.nulls();
96 let valid = |i: usize| nulls.is_none_or(|n| n.is_valid(i));
97 macro_rules! prim {
98 ($t:ty, $f:expr) => {{
99 let v = a.as_primitive::<$t>().values();
100 for (j, &i) in idx.iter().enumerate() {
101 put(out, j, if valid(i) { $f(v[i]) } else { NULL_KEY }, first);
102 }
103 }};
104 }
105 macro_rules! strs {
106 ($arr:expr) => {{
107 let s = $arr;
108 for (j, &i) in idx.iter().enumerate() {
109 put(out, j, if valid(i) { fnv1a(s.value(i).as_bytes()) } else { NULL_KEY }, first);
110 }
111 }};
112 }
113 match a.data_type() {
114 DataType::Int8 => prim!(Int8Type, |x: i8| x as i64 as u64),
115 DataType::Int16 => prim!(Int16Type, |x: i16| x as i64 as u64),
116 DataType::Int32 => prim!(Int32Type, |x: i32| x as i64 as u64),
117 DataType::Int64 => prim!(Int64Type, |x: i64| x as u64),
118 DataType::UInt8 => prim!(UInt8Type, |x: u8| x as u64),
119 DataType::UInt16 => prim!(UInt16Type, |x: u16| x as u64),
120 DataType::UInt32 => prim!(UInt32Type, |x: u32| x as u64),
121 DataType::UInt64 => prim!(UInt64Type, |x: u64| x),
122 DataType::Float64 => prim!(Float64Type, |x: f64| (if x == 0.0 { 0.0 } else { x }).to_bits()),
123 DataType::Float32 => prim!(Float32Type, |x: f32| (if x == 0.0 { 0.0f32 } else { x }).to_bits() as u64),
124 DataType::Utf8 => strs!(a.as_string::<i32>()),
125 DataType::LargeUtf8 => strs!(a.as_string::<i64>()),
126 DataType::Utf8View => strs!(a.as_string_view()),
127 _ => {
130 for (j, &i) in idx.iter().enumerate() {
131 put(out, j, if valid(i) { 0 } else { NULL_KEY }, first);
132 }
133 }
134 }
135}
136
137fn sample(batches: &[&RecordBatch], keys: &[usize], pos: &[usize]) -> Vec<u64> {
140 let mut out = vec![0u64; pos.len()];
141 let mut start = 0usize;
142 let mut at = 0usize;
143 let mut idx = Vec::new();
144 for b in batches {
145 let end = start + b.num_rows();
146 idx.clear();
147 let from = at;
148 while at < pos.len() && pos[at] < end {
149 idx.push(pos[at] - start);
150 at += 1;
151 }
152 if !idx.is_empty() {
153 for (j, &k) in keys.iter().enumerate() {
154 fold_column(b.column(k).as_ref(), &idx, &mut out[from..at], j == 0);
155 }
156 }
157 start = end;
158 if at == pos.len() {
159 break;
160 }
161 }
162 out
163}
164
165fn stats(v: Vec<u64>) -> (u64, u64, u64) {
168 let size = (v.len() * 2).next_power_of_two().max(16);
169 let mask = size - 1;
170 let mut keys = vec![0u64; size];
171 let mut counts = vec![0u32; size];
172 let mut d = 0u64;
173 for x in v {
174 let mut i = (x as usize) & mask;
175 loop {
176 if counts[i] == 0 {
177 keys[i] = x;
178 counts[i] = 1;
179 d += 1;
180 break;
181 }
182 if keys[i] == x {
183 counts[i] += 1;
184 break;
185 }
186 i = (i + 1) & mask;
187 }
188 }
189 let f1 = counts.iter().filter(|&&c| c == 1).count() as u64;
190 let f2 = counts.iter().filter(|&&c| c == 2).count() as u64;
191 (d, f1, f2)
192}
193
194fn chao1(d: u64, f1: u64, f2: u64, rows: usize) -> (f64, f64, f64) {
198 let (d, f1, f2, rows) = (d as f64, f1 as f64, f2 as f64, rows as f64);
199 let s = f2.sqrt();
200 let at = |f: f64| rows.min(d.max(d + f1 * (f1 - 1.0) / (2.0 * (f + 1.0))));
201 let low_f2 = f2 - 2.0 * s - 1.0;
202 let hi = if low_f2 > 0.0 {
203 at(low_f2)
204 } else if f1 == 0.0 {
205 d
206 } else {
207 rows
208 };
209 (at(f2), at(f2 + 2.0 * s + 2.0), hi)
210}
211
212#[cfg(test)]
219pub(crate) fn estimate(
220 batches: &[&RecordBatch],
221 keys: &[usize],
222 settled: Option<&dyn Fn(u64, u64) -> bool>,
223 total: Option<usize>,
224) -> GroupEstimate {
225 estimate_up_to(batches, keys, settled, total, SAMPLE_MAX)
226}
227
228pub(crate) fn estimate_up_to(
230 batches: &[&RecordBatch],
231 keys: &[usize],
232 settled: Option<&dyn Fn(u64, u64) -> bool>,
233 total: Option<usize>,
234 max_sample: usize,
235) -> GroupEstimate {
236 let t = Instant::now();
237 let batches: Vec<&RecordBatch> = batches.iter().copied().filter(|b| b.num_rows() > 0).collect();
238 let avail: usize = batches.iter().map(|b| b.num_rows()).sum();
239 let rows = total.unwrap_or(avail).max(avail);
240 if avail == rows && rows <= EXACT {
241 let all: Vec<usize> = (0..rows).collect();
242 let (d, _, _) = stats(sample(&batches, keys, &all));
243 return GroupEstimate {
244 estimate: d,
245 low: d,
246 high: d,
247 sample_rows: rows,
248 rows,
249 exact: true,
250 time: t.elapsed(),
251 };
252 }
253 let cap = max_sample.min(SAMPLE_MAX).min(avail / 4).max(FIRST.min(avail));
254 let mut n = FIRST.min(avail);
255 loop {
256 let (d, f1, f2) = stats(sample(&batches, keys, &positions(avail, n)));
257 let (est, lo, hi) = chao1(d, f1, f2, rows);
258 let done = n * GROWTH > cap
259 || match settled {
260 Some(s) => s(lo.round() as u64, hi.round() as u64),
261 None => est <= 2.0 * d as f64 || f2 >= SETTLED_F2 || lo * 4.0 >= rows as f64,
262 };
263 if done {
264 return GroupEstimate {
265 estimate: est.round() as u64,
266 low: lo.round() as u64,
267 high: hi.round() as u64,
268 sample_rows: n,
269 rows,
270 exact: false,
271 time: t.elapsed(),
272 };
273 }
274 n *= GROWTH;
275 }
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281 use arrow::array::{Int32Array, StringArray};
282 use arrow::datatypes::{Field, Schema};
283 use std::sync::Arc;
284
285 fn batches(keys: Vec<i32>, chunk: usize) -> Vec<RecordBatch> {
286 let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int32, true)]));
287 keys.chunks(chunk)
288 .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(Int32Array::from(c.to_vec()))]).unwrap())
289 .collect()
290 }
291
292 fn uniform(rows: usize, groups: u64) -> Vec<i32> {
293 (0..rows as u64).map(|i| (splitmix(i ^ 0x55) % groups) as i32).collect()
294 }
295
296 #[test]
297 fn small_inputs_are_counted_exactly() {
298 let b = batches((0..4000).map(|i| i % 37).collect(), 1000);
299 let r: Vec<&RecordBatch> = b.iter().collect();
300 let e = estimate(&r, &[0], None, None);
301 assert!(e.exact);
302 assert_eq!((e.estimate, e.low, e.high), (37, 37, 37));
303 }
304
305 #[test]
306 fn estimates_land_near_the_truth() {
307 for (groups, rows) in [(200u64, 1_000_000usize), (10_000, 1_000_000), (100_000, 2_000_000), (1_000_000, 4_000_000)] {
309 let keys = uniform(rows, groups);
310 let mut seen = keys.clone();
311 seen.sort_unstable();
312 seen.dedup();
313 let truth = seen.len() as f64;
314 let b = batches(keys, 8192);
315 let r: Vec<&RecordBatch> = b.iter().collect();
316 let e = estimate(&r, &[0], None, None);
317 let ratio = e.estimate as f64 / truth;
318 assert!((0.5..2.0).contains(&ratio), "groups {groups}, rows {rows}: estimate {e:?}, truth {truth}");
319 assert!(e.low as f64 <= truth * 1.2 && e.high as f64 >= truth * 0.8, "{e:?} vs {truth}");
320 }
321 }
322
323 #[test]
324 fn sorted_keys_are_not_underestimated() {
325 let rows = 1_000_000;
327 let keys: Vec<i32> = (0..rows as i32).map(|i| i / 2).collect();
328 let b = batches(keys, 8192);
329 let r: Vec<&RecordBatch> = b.iter().collect();
330 let e = estimate(&r, &[0], None, None);
331 assert!(e.high as usize >= rows / 4, "{e:?}");
332 }
333
334 #[test]
335 fn two_keys_and_strings_hash_as_tuples() {
336 let n = 20_000;
337 let schema = Arc::new(Schema::new(vec![
338 Field::new("a", DataType::Int32, false),
339 Field::new("s", DataType::Utf8, true),
340 ]));
341 let a = Int32Array::from((0..n).map(|i| i % 10).collect::<Vec<i32>>());
342 let s = StringArray::from((0..n).map(|i| if i % 7 == 0 { None } else { Some(format!("s{}", i % 3)) }).collect::<Vec<_>>());
343 let b = RecordBatch::try_new(schema, vec![Arc::new(a), Arc::new(s)]).unwrap();
344 let e = estimate(&[&b], &[0, 1], None, None);
345 assert!((30..=48).contains(&e.estimate), "{e:?}");
347 }
348
349 #[test]
350 fn a_prefix_estimates_the_whole_input() {
351 let keys = uniform(2_000_000, 100_000);
353 let b = batches(keys[..262_144].to_vec(), 8192);
354 let r: Vec<&RecordBatch> = b.iter().collect();
355 let e = estimate(&r, &[0], None, Some(2_000_000));
356 assert_eq!(e.rows, 2_000_000);
357 assert!((50_000..200_000).contains(&e.estimate), "{e:?}");
358 }
359
360 #[test]
362 #[ignore]
363 fn probe_timing() {
364 let keys: Vec<i64> = (0..262_144u64).map(|i| (splitmix(i) % 1_000_000) as i64).collect();
365 let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int64, false)]));
366 let b: Vec<RecordBatch> = keys
367 .chunks(8192)
368 .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(arrow::array::Int64Array::from(c.to_vec()))]).unwrap())
369 .collect();
370 let r: Vec<&RecordBatch> = b.iter().collect();
371 for n in [512usize, 2048, 8192, 32768] {
372 let pos = positions(262_144, n);
373 let t = Instant::now();
374 let mut v = Vec::new();
375 for _ in 0..100 {
376 v = sample(&r, &[0], &pos);
377 }
378 let ts = t.elapsed() / 100;
379 let t = Instant::now();
380 for _ in 0..100 {
381 std::hint::black_box(stats(v.clone()));
382 }
383 let tt = t.elapsed() / 100;
384 println!("n {n}: sample {ts:?}, stats {tt:?}");
385 }
386 }
387
388 #[test]
389 fn same_input_same_estimate() {
390 let b = batches(uniform(500_000, 50_000), 8192);
391 let r: Vec<&RecordBatch> = b.iter().collect();
392 let (x, y) = (estimate(&r, &[0], None, None), estimate(&r, &[0], None, None));
393 assert_eq!((x.estimate, x.low, x.high, x.sample_rows), (y.estimate, y.low, y.high, y.sample_rows));
394 }
395}