Skip to main content

sample_arrow_rs/
list.rs

1//! Samplers for generating an arrow [`ListArray`].
2
3use 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        // Convert our custom Bitmap type to arrow's NullBuffer
47        let null_buffer = validity.map(|bitmap| bitmap);
48
49        // Convert offsets to a ScalarBuffer first, then to an OffsetBuffer
50        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        // Create the ListArray and return as Arc<dyn Array>
60        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        // Convert offsets to a ScalarBuffer first, then to an OffsetBuffer
113        let scalar_buffer = arrow_buffer::ScalarBuffer::from(offsets);
114        let offsets_buffer = OffsetBuffer::new(scalar_buffer);
115
116        // Convert the validity bitmap to a NullBuffer if present
117        let null_buffer = self.validity.generate(g).map(|bitmap| bitmap);
118
119        // Create the ListArray
120        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}