Skip to main content

sample_arrow_rs/
datatypes.rs

1//! Samplers for generating an arrow [`DataType`].
2
3use arrow_schema::{DataType, Field, Fields};
4use sample_std::{sampler_choice, Always, Random, Sample, Shrunk};
5use std::sync::Arc;
6
7pub type DataTypeSampler = Box<dyn Sample<Output = DataType> + Send + Sync>;
8
9struct FieldSampler<N, V> {
10    names: N,
11    nullable: V,
12    inner: DataTypeSampler,
13}
14
15impl<N, V> Sample for FieldSampler<N, V>
16where
17    N: Sample<Output = String>,
18    V: Sample<Output = bool>,
19{
20    type Output = Field;
21
22    fn generate(&mut self, g: &mut Random) -> Self::Output {
23        // In arrow-rs, Field::new takes name, data_type, is_nullable
24        Field::new(
25            self.names.generate(g),
26            self.inner.generate(g),
27            self.nullable.generate(g),
28        )
29    }
30}
31
32struct StructDataTypeSampler<S, F> {
33    size: S,
34    field: F,
35}
36
37impl<S, F> Sample for StructDataTypeSampler<S, F>
38where
39    S: Sample<Output = usize>,
40    F: Sample<Output = Field>,
41{
42    type Output = DataType;
43
44    fn generate(&mut self, g: &mut Random) -> Self::Output {
45        let size = self.size.generate(g);
46        let fields = (0..size)
47            .map(|_| self.field.generate(g))
48            .collect::<Fields>();
49        // In arrow-rs, we use the struct constructor directly
50        DataType::Struct(fields)
51    }
52}
53
54pub fn sample_flat() -> DataTypeSampler {
55    Box::new(sampler_choice([
56        Always(DataType::Float32),
57        Always(DataType::Float64),
58        Always(DataType::Int8),
59        Always(DataType::Int16),
60        Always(DataType::Int32),
61        Always(DataType::Int64),
62        Always(DataType::UInt8),
63        Always(DataType::UInt16),
64        Always(DataType::UInt32),
65        Always(DataType::UInt64),
66    ]))
67}
68
69pub struct ArbitraryDataType<N, V, B, F> {
70    pub names: N,
71    pub nullable: V,
72    pub struct_branch: B,
73    pub flat: F,
74}
75
76impl<N, V, B, F> ArbitraryDataType<N, V, B, F>
77where
78    N: Sample<Output = String> + Clone + Send + Sync + 'static,
79    V: Sample<Output = bool> + Clone + Send + Sync + 'static,
80    B: Sample<Output = usize> + Clone + Send + Sync + 'static,
81    F: Fn() -> DataTypeSampler,
82{
83    pub fn sample_nested<IF>(&self, inner: IF) -> DataTypeSampler
84    where
85        IF: Fn() -> DataTypeSampler + Clone,
86    {
87        let names_clone = self.names.clone();
88        let nullable_clone = self.nullable.clone();
89        let inner_clone = inner.clone();
90
91        // Create a field constructor that can be used multiple times
92        let field_constructor = move || FieldSampler {
93            names: names_clone.clone(),
94            nullable: nullable_clone.clone(),
95            inner: inner_clone(),
96        };
97
98        // Create a list type sampler
99        let list_field_sampler = {
100            let names_clone = self.names.clone();
101            let nullable_clone = self.nullable.clone();
102            let inner_clone = inner.clone();
103
104            move || {
105                // Create a field sampler for the list element
106                let field_sampler = FieldSampler {
107                    names: names_clone.clone(),
108                    nullable: nullable_clone.clone(),
109                    inner: inner_clone(),
110                };
111
112                // Create a sampler that produces DataType::List with the generated field
113                Box::new(ListFieldSampler { field_sampler }) as DataTypeSampler
114            }
115        };
116
117        Box::new(sampler_choice([
118            Box::new((self.flat)()) as DataTypeSampler,
119            Box::new(StructDataTypeSampler {
120                size: self.struct_branch.clone(),
121                field: field_constructor(),
122            }),
123            list_field_sampler(),
124        ]))
125    }
126
127    pub fn sample_depth(&self, depth: usize) -> DataTypeSampler {
128        let flats = (self.flat)();
129        if depth == 0 {
130            flats
131        } else {
132            let inner = || self.sample_depth(depth - 1);
133            Box::new(sampler_choice([self.sample_nested(inner), flats]))
134        }
135    }
136}
137
138struct ListFieldSampler<F> {
139    field_sampler: F,
140}
141
142impl<F> Sample for ListFieldSampler<F>
143where
144    F: Sample<Output = Field>,
145{
146    type Output = DataType;
147
148    fn generate(&mut self, g: &mut Random) -> Self::Output {
149        let field = self.field_sampler.generate(g);
150        DataType::List(Arc::new(field))
151    }
152
153    fn shrink(&self, _: Self::Output) -> Shrunk<Self::Output> {
154        Box::new(std::iter::empty())
155    }
156}