1use std::ops::Range;
4use std::sync::Arc;
5
6use arrow_array::{Array, ArrayRef, ListArray};
7use arrow_buffer::OffsetBuffer;
8use arrow_schema::{DataType, Field};
9use sample_std::{Random, Sample, Shrunk};
10
11use crate::{generate_validity, ArrowSampler, Bitmap, SampleLen, SetLen};
12
13pub struct ListSampler<V> {
14 pub data_type: DataType,
15 pub len: Range<usize>,
16 pub null: Option<V>,
17 pub inner: ArrowSampler,
18}
19
20impl<V> Sample for ListSampler<V>
21where
22 V: Sample<Output = bool> + Send + Sync + 'static,
23{
24 type Output = ArrayRef;
25
26 fn generate(&mut self, g: &mut Random) -> Self::Output {
27 let values = self.inner.generate(g);
28 let len = g.gen_range(self.len.clone());
29 let mut ix = 0;
30 let mut offsets = vec![0i32];
31
32 for outer_ix in 0..len {
33 if outer_ix + 1 != len {
34 let remaining = values.len() - ix;
35 let fair = std::cmp::max(2, remaining / (len - outer_ix));
36 let upper = std::cmp::min(values.len() - ix, fair);
37 let count = g.gen_range(0..=upper);
38 ix += count;
39 offsets.push(ix as i32);
40 } else {
41 offsets.push(values.len() as i32);
42 }
43 }
44
45 let validity = generate_validity(&mut self.null, g, len);
46 let null_buffer = validity.map(|bitmap| bitmap);
48
49 let scalar_buffer = arrow_buffer::ScalarBuffer::from(offsets);
51 let offsets_buffer = OffsetBuffer::new(scalar_buffer);
52
53 let field = if let DataType::List(field) = &self.data_type {
54 field.clone()
55 } else {
56 panic!("Expected List data type")
57 };
58
59 Arc::new(ListArray::new(field, offsets_buffer, values, null_buffer))
61 }
62
63 fn shrink(&self, _v: Self::Output) -> Shrunk<Self::Output> {
64 Box::new(std::iter::empty())
65 }
66}
67
68pub struct ListWithLen<V, C, A, N> {
69 pub len: usize,
70 pub validity: V,
71 pub count: C,
72 pub inner: A,
73 pub inner_name: N,
74}
75
76impl<V: SetLen, C, A, N> SetLen for ListWithLen<V, C, A, N> {
77 fn set_len(&mut self, len: usize) {
78 self.len = len;
79 self.validity.set_len(len);
80 }
81}
82
83impl<V, C, A, N> Sample for ListWithLen<V, C, A, N>
84where
85 V: Sample<Output = Option<Bitmap>> + SetLen,
86 C: Sample<Output = i32>,
87 A: Sample<Output = ArrayRef> + SetLen,
88 N: Sample<Output = String>,
89{
90 type Output = ArrayRef;
91
92 fn generate(&mut self, g: &mut Random) -> Self::Output {
93 let mut offsets = vec![0];
94 let mut inner_len: i32 = 0;
95 for _ in 0..self.len {
96 let count = self.count.generate(g);
97 assert!(count >= 0);
98 inner_len += count;
99 offsets.push(inner_len);
100 }
101
102 self.inner.set_len(inner_len as usize);
103 let values = self.inner.generate(g);
104 let is_nullable = values.nulls().is_some();
105 let inner_name = self.inner_name.generate(g);
106 let field = Arc::new(Field::new(
107 inner_name,
108 values.data_type().clone(),
109 is_nullable,
110 ));
111
112 let scalar_buffer = arrow_buffer::ScalarBuffer::from(offsets);
114 let offsets_buffer = OffsetBuffer::new(scalar_buffer);
115
116 let null_buffer = self.validity.generate(g).map(|bitmap| bitmap);
118
119 Arc::new(ListArray::new(field, offsets_buffer, values, null_buffer))
121 }
122
123 fn shrink(&self, _: Self::Output) -> Shrunk<Self::Output> {
124 Box::new(std::iter::empty())
125 }
126}
127
128impl<V, C, A, N> SampleLen for ListWithLen<V, C, A, N>
129where
130 V: Sample<Output = Option<Bitmap>> + SetLen,
131 C: Sample<Output = i32>,
132 A: Sample<Output = ArrayRef> + SetLen,
133 N: Sample<Output = String>,
134{
135}