Skip to main content

lance_datagen/
generator.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use std::{collections::HashMap, iter, marker::PhantomData, sync::Arc, sync::LazyLock};
5
6use arrow::{
7    array::{ArrayData, AsArray, Float32Builder, GenericBinaryBuilder, GenericStringBuilder},
8    buffer::{BooleanBuffer, Buffer, OffsetBuffer, ScalarBuffer},
9    datatypes::{
10        ArrowPrimitiveType, Float32Type, Int32Type, Int64Type, IntervalDayTime,
11        IntervalMonthDayNano, UInt32Type,
12    },
13};
14use arrow_array::{
15    Array, BinaryArray, FixedSizeBinaryArray, FixedSizeListArray, Float32Array, LargeListArray,
16    LargeStringArray, ListArray, MapArray, NullArray, OffsetSizeTrait, PrimitiveArray, RecordBatch,
17    RecordBatchOptions, RecordBatchReader, StringArray, StructArray, make_array,
18    types::{ArrowDictionaryKeyType, BinaryType, ByteArrayType, Utf8Type},
19};
20use arrow_schema::{ArrowError, DataType, Field, Fields, IntervalUnit, Schema, SchemaRef};
21use futures::{StreamExt, stream::BoxStream};
22use rand::{Rng, RngCore, SeedableRng, distr::Uniform};
23use rand_distr::Zipf;
24
25use self::array::rand_with_distribution;
26
27#[derive(Copy, Clone, Debug, Default)]
28pub struct RowCount(u64);
29#[derive(Copy, Clone, Debug, Default)]
30pub struct BatchCount(u32);
31#[derive(Copy, Clone, Debug, Default)]
32pub struct ByteCount(u64);
33#[derive(Copy, Clone, Debug, Default)]
34pub struct Dimension(u32);
35
36impl From<u32> for BatchCount {
37    fn from(n: u32) -> Self {
38        Self(n)
39    }
40}
41
42impl From<u64> for RowCount {
43    fn from(n: u64) -> Self {
44        Self(n)
45    }
46}
47
48impl From<u64> for ByteCount {
49    fn from(n: u64) -> Self {
50        Self(n)
51    }
52}
53
54impl From<u32> for Dimension {
55    fn from(n: u32) -> Self {
56        Self(n)
57    }
58}
59
60/// A trait for anything that can generate arrays of data
61pub trait ArrayGenerator: Send + Sync + std::fmt::Debug {
62    /// Generate an array of the given length
63    ///
64    /// # Arguments
65    ///
66    /// * `length` - The number of elements to generate
67    /// * `rng` - The random number generator to use
68    ///
69    /// # Returns
70    ///
71    /// An array of the given length
72    ///
73    /// Note: Not every generator needs an rng.  However, it is passed here because many do and this
74    /// lets us manage RNGs at the batch level instead of the array level.
75    fn generate(
76        &mut self,
77        length: RowCount,
78        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
79    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError>;
80
81    /// Generate an array of the given length using a new RNG with the default seed
82    ///
83    /// # Arguments
84    ///
85    /// * `length` - The number of elements to generate
86    ///
87    /// # Returns
88    ///
89    /// An array of the given length
90    fn generate_default(
91        &mut self,
92        length: RowCount,
93    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
94        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
95        Self::generate(self, length, &mut rng)
96    }
97    /// Get the data type of the array that this generator produces
98    ///
99    /// # Returns
100    ///
101    /// The data type of the array that this generator produces
102    fn data_type(&self) -> &DataType;
103    /// Gets metadata that should be associated with the field generated by this generator
104    fn metadata(&self) -> Option<HashMap<String, String>> {
105        None
106    }
107    /// Get the size of each element in bytes
108    ///
109    /// # Returns
110    ///
111    /// The size of each element in bytes.  Will be None if the size varies by element.
112    fn element_size_bytes(&self) -> Option<ByteCount>;
113}
114
115#[derive(Debug)]
116pub struct CycleNullGenerator {
117    generator: Box<dyn ArrayGenerator>,
118    validity: Vec<bool>,
119    idx: usize,
120}
121#[derive(Debug)]
122pub struct CycleNanGenerator {
123    generator: Box<dyn ArrayGenerator>,
124    nan_pattern: Vec<bool>,
125    idx: usize,
126}
127
128impl ArrayGenerator for CycleNanGenerator {
129    fn generate(
130        &mut self,
131        length: RowCount,
132        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
133    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
134        let array = self.generator.generate(length, rng)?;
135
136        // Only apply NaN pattern to float types
137        match array.data_type() {
138            DataType::Float16 => {
139                let float_array = array
140                    .as_any()
141                    .downcast_ref::<arrow_array::Float16Array>()
142                    .unwrap();
143                let mut values: Vec<half::f16> = float_array.values().to_vec();
144
145                for (i, &should_be_nan) in self
146                    .nan_pattern
147                    .iter()
148                    .cycle()
149                    .skip(self.idx)
150                    .take(length.0 as usize)
151                    .enumerate()
152                {
153                    if should_be_nan {
154                        values[i] = half::f16::NAN;
155                    }
156                }
157
158                self.idx = (self.idx + (length.0 as usize)) % self.nan_pattern.len();
159                Ok(Arc::new(arrow_array::Float16Array::from(values)))
160            }
161            DataType::Float32 => {
162                let float_array = array
163                    .as_any()
164                    .downcast_ref::<arrow_array::Float32Array>()
165                    .unwrap();
166                let mut values: Vec<f32> = float_array.values().to_vec();
167
168                for (i, &should_be_nan) in self
169                    .nan_pattern
170                    .iter()
171                    .cycle()
172                    .skip(self.idx)
173                    .take(length.0 as usize)
174                    .enumerate()
175                {
176                    if should_be_nan {
177                        values[i] = f32::NAN;
178                    }
179                }
180
181                self.idx = (self.idx + (length.0 as usize)) % self.nan_pattern.len();
182                Ok(Arc::new(arrow_array::Float32Array::from(values)))
183            }
184            DataType::Float64 => {
185                let float_array = array
186                    .as_any()
187                    .downcast_ref::<arrow_array::Float64Array>()
188                    .unwrap();
189                let mut values: Vec<f64> = float_array.values().to_vec();
190
191                for (i, &should_be_nan) in self
192                    .nan_pattern
193                    .iter()
194                    .cycle()
195                    .skip(self.idx)
196                    .take(length.0 as usize)
197                    .enumerate()
198                {
199                    if should_be_nan {
200                        values[i] = f64::NAN;
201                    }
202                }
203
204                self.idx = (self.idx + (length.0 as usize)) % self.nan_pattern.len();
205                Ok(Arc::new(arrow_array::Float64Array::from(values)))
206            }
207            _ => {
208                // For non-float types, just return the original array unchanged
209                Ok(array)
210            }
211        }
212    }
213
214    fn data_type(&self) -> &DataType {
215        self.generator.data_type()
216    }
217
218    fn element_size_bytes(&self) -> Option<ByteCount> {
219        self.generator.element_size_bytes()
220    }
221}
222
223impl ArrayGenerator for CycleNullGenerator {
224    fn generate(
225        &mut self,
226        length: RowCount,
227        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
228    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
229        let array = self.generator.generate(length, rng)?;
230        let data = array.to_data();
231        let validity_itr = self
232            .validity
233            .iter()
234            .cycle()
235            .skip(self.idx)
236            .take(length.0 as usize)
237            .copied();
238        let validity_bitmap = BooleanBuffer::from_iter(validity_itr);
239
240        self.idx = (self.idx + (length.0 as usize)) % self.validity.len();
241        unsafe {
242            let new_data = ArrayData::new_unchecked(
243                data.data_type().clone(),
244                data.len(),
245                None,
246                Some(validity_bitmap.into_inner()),
247                data.offset(),
248                data.buffers().to_vec(),
249                data.child_data().into(),
250            );
251            Ok(make_array(new_data))
252        }
253    }
254
255    fn data_type(&self) -> &DataType {
256        self.generator.data_type()
257    }
258
259    fn element_size_bytes(&self) -> Option<ByteCount> {
260        self.generator.element_size_bytes()
261    }
262}
263
264#[derive(Debug)]
265pub struct MetadataGenerator {
266    generator: Box<dyn ArrayGenerator>,
267    metadata: HashMap<String, String>,
268}
269
270impl ArrayGenerator for MetadataGenerator {
271    fn generate(
272        &mut self,
273        length: RowCount,
274        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
275    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
276        self.generator.generate(length, rng)
277    }
278
279    fn metadata(&self) -> Option<HashMap<String, String>> {
280        Some(self.metadata.clone())
281    }
282
283    fn data_type(&self) -> &DataType {
284        self.generator.data_type()
285    }
286
287    fn element_size_bytes(&self) -> Option<ByteCount> {
288        self.generator.element_size_bytes()
289    }
290}
291
292#[derive(Debug)]
293pub struct NullGenerator {
294    generator: Box<dyn ArrayGenerator>,
295    null_probability: f64,
296}
297
298impl ArrayGenerator for NullGenerator {
299    fn generate(
300        &mut self,
301        length: RowCount,
302        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
303    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
304        let array = self.generator.generate(length, rng)?;
305        let data = array.to_data();
306
307        if self.null_probability < 0.0 || self.null_probability > 1.0 {
308            return Err(ArrowError::InvalidArgumentError(format!(
309                "null_probability must be between 0 and 1, got {}",
310                self.null_probability
311            )));
312        }
313
314        let (null_count, new_validity) = if self.null_probability == 0.0 {
315            if data.null_count() == 0 {
316                return Ok(array);
317            } else {
318                (0_usize, None)
319            }
320        } else if self.null_probability == 1.0 {
321            if data.null_count() == data.len() {
322                return Ok(array);
323            } else {
324                let all_nulls = BooleanBuffer::new_unset(array.len());
325                (array.len(), Some(all_nulls.into_inner()))
326            }
327        } else {
328            let array_len = array.len();
329            let num_validity_bytes = array_len.div_ceil(8);
330            let mut null_count = 0;
331            // Sampling the RNG once per bit is kind of slow so we do this to sample once
332            // per byte.  We only get 8 bits of RNG resolution but that should be good enough.
333            let threshold = (self.null_probability * u8::MAX as f64) as u8;
334            let bytes = (0..num_validity_bytes)
335                .map(|byte_idx| {
336                    let mut sample = rng.random::<u64>();
337                    let mut byte: u8 = 0;
338                    for bit_idx in 0..8 {
339                        // We could probably overshoot and fill in extra bits with random data but
340                        // this is cleaner and that would mess up the null count
341                        byte <<= 1;
342                        let pos = byte_idx * 8 + (7 - bit_idx);
343                        if pos < array_len {
344                            let sample_piece = sample & 0xFF;
345                            let is_null = (sample_piece as u8) < threshold;
346                            byte |= (!is_null) as u8;
347                            null_count += is_null as usize;
348                        }
349                        sample >>= 8;
350                    }
351                    byte
352                })
353                .collect::<Vec<_>>();
354            let new_validity = Buffer::from_iter(bytes);
355            (null_count, Some(new_validity))
356        };
357
358        unsafe {
359            let new_data = ArrayData::new_unchecked(
360                data.data_type().clone(),
361                data.len(),
362                Some(null_count),
363                new_validity,
364                data.offset(),
365                data.buffers().to_vec(),
366                data.child_data().into(),
367            );
368            Ok(make_array(new_data))
369        }
370    }
371
372    fn metadata(&self) -> Option<HashMap<String, String>> {
373        self.generator.metadata()
374    }
375
376    fn data_type(&self) -> &DataType {
377        self.generator.data_type()
378    }
379
380    fn element_size_bytes(&self) -> Option<ByteCount> {
381        self.generator.element_size_bytes()
382    }
383}
384
385pub trait ArrayGeneratorExt {
386    /// Replaces the validity bitmap of generated arrays, inserting nulls with a given probability
387    fn with_random_nulls(self, null_probability: f64) -> Box<dyn ArrayGenerator>;
388    /// Replaces the validity bitmap of generated arrays with the inverse of `nulls`, cycling if needed
389    fn with_nulls(self, nulls: &[bool]) -> Box<dyn ArrayGenerator>;
390    /// Replaces the values of generated arrays with NaN values, cycling if needed
391    ///
392    /// Will have no effect if the data type is not a floating point data type
393    fn with_nans(self, nans: &[bool]) -> Box<dyn ArrayGenerator>;
394    /// Replaces the validity bitmap of generated arrays with `validity`, cycling if needed
395    fn with_validity(self, nulls: &[bool]) -> Box<dyn ArrayGenerator>;
396    fn with_metadata(self, metadata: HashMap<String, String>) -> Box<dyn ArrayGenerator>;
397}
398
399impl ArrayGeneratorExt for Box<dyn ArrayGenerator> {
400    fn with_random_nulls(self, null_probability: f64) -> Box<dyn ArrayGenerator> {
401        Box::new(NullGenerator {
402            generator: self,
403            null_probability,
404        })
405    }
406
407    fn with_nulls(self, nulls: &[bool]) -> Box<dyn ArrayGenerator> {
408        Box::new(CycleNullGenerator {
409            generator: self,
410            validity: nulls.iter().map(|v| !*v).collect(),
411            idx: 0,
412        })
413    }
414
415    fn with_nans(self, nans: &[bool]) -> Box<dyn ArrayGenerator> {
416        Box::new(CycleNanGenerator {
417            generator: self,
418            nan_pattern: nans.to_vec(),
419            idx: 0,
420        })
421    }
422
423    fn with_validity(self, validity: &[bool]) -> Box<dyn ArrayGenerator> {
424        Box::new(CycleNullGenerator {
425            generator: self,
426            validity: validity.to_vec(),
427            idx: 0,
428        })
429    }
430
431    fn with_metadata(self, metadata: HashMap<String, String>) -> Box<dyn ArrayGenerator> {
432        Box::new(MetadataGenerator {
433            generator: self,
434            metadata,
435        })
436    }
437}
438
439pub struct NTimesIter<I: Iterator>
440where
441    I::Item: Copy,
442{
443    iter: I,
444    n: u32,
445    cur: I::Item,
446    count: u32,
447}
448
449// Note: if this is used then there is a performance hit as the
450// inner loop cannot experience vectorization
451//
452// TODO: maybe faster to build the vec and then repeat it into
453// the destination array?
454impl<I: Iterator> Iterator for NTimesIter<I>
455where
456    I::Item: Copy,
457{
458    type Item = I::Item;
459
460    fn next(&mut self) -> Option<Self::Item> {
461        if self.count == 0 {
462            self.count = self.n - 1;
463            self.cur = self.iter.next()?;
464        } else {
465            self.count -= 1;
466        }
467        Some(self.cur)
468    }
469
470    fn size_hint(&self) -> (usize, Option<usize>) {
471        let (lower, upper) = self.iter.size_hint();
472        let lower = lower * self.n as usize;
473        let upper = upper.map(|u| u * self.n as usize);
474        (lower, upper)
475    }
476}
477
478pub struct FnGen<T, ArrayType, F: FnMut(&mut rand_xoshiro::Xoshiro256PlusPlus) -> T>
479where
480    T: Copy + Default,
481    ArrayType: arrow_array::Array + From<Vec<T>>,
482{
483    data_type: DataType,
484    generator: F,
485    array_type: PhantomData<ArrayType>,
486    repeat: u32,
487    leftover: T,
488    leftover_count: u32,
489    element_size_bytes: Option<ByteCount>,
490}
491
492impl<T, ArrayType, F: FnMut(&mut rand_xoshiro::Xoshiro256PlusPlus) -> T> std::fmt::Debug
493    for FnGen<T, ArrayType, F>
494where
495    T: Copy + Default,
496    ArrayType: arrow_array::Array + From<Vec<T>>,
497{
498    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
499        f.debug_struct("FnGen")
500            .field("data_type", &self.data_type)
501            .field("array_type", &self.array_type)
502            .field("repeat", &self.repeat)
503            .field("leftover_count", &self.leftover_count)
504            .field("element_size_bytes", &self.element_size_bytes)
505            .finish()
506    }
507}
508
509impl<T, ArrayType, F: FnMut(&mut rand_xoshiro::Xoshiro256PlusPlus) -> T> FnGen<T, ArrayType, F>
510where
511    T: Copy + Default,
512    ArrayType: arrow_array::Array + From<Vec<T>>,
513{
514    fn new_known_size(
515        data_type: DataType,
516        generator: F,
517        repeat: u32,
518        element_size_bytes: ByteCount,
519    ) -> Self {
520        Self {
521            data_type,
522            generator,
523            array_type: PhantomData,
524            repeat,
525            leftover: T::default(),
526            leftover_count: 0,
527            element_size_bytes: Some(element_size_bytes),
528        }
529    }
530
531    fn new_unknown_size(data_type: DataType, generator: F, repeat: u32) -> Self {
532        Self {
533            data_type,
534            generator,
535            array_type: PhantomData,
536            repeat,
537            leftover: T::default(),
538            leftover_count: 0,
539            element_size_bytes: None,
540        }
541    }
542}
543
544impl<T, ArrayType, F: FnMut(&mut rand_xoshiro::Xoshiro256PlusPlus) -> T> ArrayGenerator
545    for FnGen<T, ArrayType, F>
546where
547    T: Copy + Default + Send + Sync,
548    ArrayType: arrow_array::Array + From<Vec<T>> + 'static,
549    F: Send + Sync,
550{
551    fn generate(
552        &mut self,
553        length: RowCount,
554        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
555    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
556        let iter = (0..length.0).map(|_| (self.generator)(rng));
557        let values = if self.repeat > 1 {
558            Vec::from_iter(
559                NTimesIter {
560                    iter,
561                    n: self.repeat,
562                    cur: self.leftover,
563                    count: self.leftover_count,
564                }
565                .take(length.0 as usize),
566            )
567        } else {
568            Vec::from_iter(iter)
569        };
570        self.leftover_count = ((self.leftover_count as u64 + length.0) % self.repeat as u64) as u32;
571        self.leftover = values.last().copied().unwrap_or(T::default());
572        let array = ArrayType::from(values);
573        // `ArrayType::from` uses the primitive type's default metadata. For
574        // timezone-aware timestamps this drops the timezone, so restore the
575        // generator's declared type when it differs.
576        if array.data_type() == &self.data_type {
577            return Ok(Arc::new(array));
578        }
579        let data = array
580            .into_data()
581            .into_builder()
582            .data_type(self.data_type.clone())
583            .build()?;
584        Ok(make_array(data))
585    }
586
587    fn data_type(&self) -> &DataType {
588        &self.data_type
589    }
590
591    fn element_size_bytes(&self) -> Option<ByteCount> {
592        self.element_size_bytes
593    }
594}
595
596#[derive(Copy, Clone, Debug)]
597pub struct Seed(pub u64);
598pub const DEFAULT_SEED: Seed = Seed(42);
599
600impl From<u64> for Seed {
601    fn from(n: u64) -> Self {
602        Self(n)
603    }
604}
605
606#[derive(Debug)]
607pub struct CycleVectorGenerator {
608    underlying_gen: Box<dyn ArrayGenerator>,
609    dimension: Dimension,
610    data_type: DataType,
611}
612
613impl CycleVectorGenerator {
614    pub fn new(underlying_gen: Box<dyn ArrayGenerator>, dimension: Dimension) -> Self {
615        let data_type = DataType::FixedSizeList(
616            Arc::new(Field::new("item", underlying_gen.data_type().clone(), true)),
617            dimension.0 as i32,
618        );
619        Self {
620            underlying_gen,
621            dimension,
622            data_type,
623        }
624    }
625}
626
627impl ArrayGenerator for CycleVectorGenerator {
628    fn generate(
629        &mut self,
630        length: RowCount,
631        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
632    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
633        let values = self
634            .underlying_gen
635            .generate(RowCount::from(length.0 * self.dimension.0 as u64), rng)?;
636        let field = Arc::new(Field::new("item", values.data_type().clone(), true));
637        let values = Arc::new(values);
638
639        let array = FixedSizeListArray::try_new(field, self.dimension.0 as i32, values, None)?;
640
641        Ok(Arc::new(array))
642    }
643
644    fn data_type(&self) -> &DataType {
645        &self.data_type
646    }
647
648    fn element_size_bytes(&self) -> Option<ByteCount> {
649        self.underlying_gen
650            .element_size_bytes()
651            .map(|byte_count| ByteCount::from(byte_count.0 * self.dimension.0 as u64))
652    }
653}
654
655#[derive(Debug)]
656pub struct CycleListGenerator {
657    underlying_gen: Box<dyn ArrayGenerator>,
658    lengths_gen: Box<dyn ArrayGenerator>,
659    data_type: DataType,
660}
661
662impl CycleListGenerator {
663    pub fn new(
664        underlying_gen: Box<dyn ArrayGenerator>,
665        min_list_size: Dimension,
666        max_list_size: Dimension,
667    ) -> Self {
668        let data_type = DataType::List(Arc::new(Field::new(
669            "item",
670            underlying_gen.data_type().clone(),
671            true,
672        )));
673        let lengths_dist = Uniform::new(min_list_size.0, max_list_size.0).unwrap();
674        let lengths_gen = rand_with_distribution::<UInt32Type, Uniform<u32>>(lengths_dist);
675        Self {
676            underlying_gen,
677            lengths_gen,
678            data_type,
679        }
680    }
681}
682
683impl ArrayGenerator for CycleListGenerator {
684    fn generate(
685        &mut self,
686        length: RowCount,
687        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
688    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
689        let lengths = self.lengths_gen.generate(length, rng)?;
690        let lengths = lengths.as_primitive::<UInt32Type>();
691        let total_length = lengths.values().iter().map(|i| *i as u64).sum::<u64>();
692        let offsets = OffsetBuffer::from_lengths(lengths.values().iter().map(|v| *v as usize));
693        let values = self
694            .underlying_gen
695            .generate(RowCount::from(total_length), rng)?;
696        let field = Arc::new(Field::new("item", values.data_type().clone(), true));
697        let values = Arc::new(values);
698
699        let array = ListArray::try_new(field, offsets, values, None)?;
700
701        Ok(Arc::new(array))
702    }
703
704    fn data_type(&self) -> &DataType {
705        &self.data_type
706    }
707
708    fn element_size_bytes(&self) -> Option<ByteCount> {
709        None
710    }
711}
712
713#[derive(Debug, Default)]
714pub struct PseudoUuidGenerator {}
715
716impl ArrayGenerator for PseudoUuidGenerator {
717    fn generate(
718        &mut self,
719        length: RowCount,
720        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
721    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
722        Ok(Arc::new(FixedSizeBinaryArray::try_from_iter(
723            (0..length.0).map(|_| {
724                let mut data = vec![0; 16];
725                rng.fill_bytes(&mut data);
726                data
727            }),
728        )?))
729    }
730
731    fn data_type(&self) -> &DataType {
732        &DataType::FixedSizeBinary(16)
733    }
734
735    fn element_size_bytes(&self) -> Option<ByteCount> {
736        Some(ByteCount::from(16))
737    }
738}
739
740#[derive(Debug, Default)]
741pub struct PseudoUuidHexGenerator {}
742
743impl ArrayGenerator for PseudoUuidHexGenerator {
744    fn generate(
745        &mut self,
746        length: RowCount,
747        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
748    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
749        let mut data = vec![0; 16 * length.0 as usize];
750        rng.fill_bytes(&mut data);
751        let data_hex = hex::encode(data);
752
753        Ok(Arc::new(StringArray::from_iter_values(
754            (0..length.0 as usize).map(|i| data_hex.get(i * 32..(i + 1) * 32).unwrap()),
755        )))
756    }
757
758    fn data_type(&self) -> &DataType {
759        &DataType::Utf8
760    }
761
762    fn element_size_bytes(&self) -> Option<ByteCount> {
763        Some(ByteCount::from(16))
764    }
765}
766
767#[derive(Debug, Default)]
768pub struct RandomBooleanGenerator {}
769
770impl ArrayGenerator for RandomBooleanGenerator {
771    fn generate(
772        &mut self,
773        length: RowCount,
774        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
775    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
776        let num_bytes = length.0.div_ceil(8);
777        let mut bytes = vec![0; num_bytes as usize];
778        rng.fill_bytes(&mut bytes);
779        let bytes = BooleanBuffer::new(Buffer::from(bytes), 0, length.0 as usize);
780        Ok(Arc::new(arrow_array::BooleanArray::new(bytes, None)))
781    }
782
783    fn data_type(&self) -> &DataType {
784        &DataType::Boolean
785    }
786
787    fn element_size_bytes(&self) -> Option<ByteCount> {
788        // We can't say 1/8th of a byte and 1 byte would be a pretty extreme over-count so let's leave
789        // it at None until someone needs this.  Then we can probably special case this (e.g. make a ByteCount::ONE_BIT)
790        None
791    }
792}
793
794// Instead of using the "standard distribution" and generating values there are some cases (e.g. f16 / decimal)
795// where we just generate random bytes because there is no rand support
796pub struct RandomBytesGenerator<T: ArrowPrimitiveType + Send + Sync> {
797    phantom: PhantomData<T>,
798    data_type: DataType,
799}
800
801impl<T: ArrowPrimitiveType + Send + Sync> std::fmt::Debug for RandomBytesGenerator<T> {
802    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
803        f.debug_struct("RandomBytesGenerator")
804            .field("data_type", &self.data_type)
805            .finish()
806    }
807}
808
809impl<T: ArrowPrimitiveType + Send + Sync> RandomBytesGenerator<T> {
810    fn new(data_type: DataType) -> Self {
811        Self {
812            phantom: Default::default(),
813            data_type,
814        }
815    }
816
817    fn byte_width() -> Result<u64, ArrowError> {
818        T::DATA_TYPE.primitive_width().ok_or_else(|| ArrowError::InvalidArgumentError(format!("Cannot generate the data type {} with the RandomBytesGenerator because it is not a fixed-width bytes type", T::DATA_TYPE))).map(|val| val as u64)
819    }
820}
821
822impl<T: ArrowPrimitiveType + Send + Sync> ArrayGenerator for RandomBytesGenerator<T> {
823    fn generate(
824        &mut self,
825        length: RowCount,
826        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
827    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
828        let num_bytes = length.0 * Self::byte_width()?;
829        let mut bytes = vec![0; num_bytes as usize];
830        rng.fill_bytes(&mut bytes);
831        let bytes = ScalarBuffer::new(Buffer::from(bytes), 0, length.0 as usize);
832        Ok(Arc::new(
833            PrimitiveArray::<T>::new(bytes, None).with_data_type(self.data_type.clone()),
834        ))
835    }
836
837    fn data_type(&self) -> &DataType {
838        &self.data_type
839    }
840
841    fn element_size_bytes(&self) -> Option<ByteCount> {
842        Self::byte_width().map(ByteCount::from).ok()
843    }
844}
845
846// This is pretty much the same thing as RandomBinaryGenerator but we can't use that
847// because there is no ArrowPrimitiveType for FixedSizeBinary
848#[derive(Debug)]
849pub struct RandomFixedSizeBinaryGenerator {
850    data_type: DataType,
851    size: i32,
852}
853
854impl RandomFixedSizeBinaryGenerator {
855    fn new(size: i32) -> Self {
856        Self {
857            size,
858            data_type: DataType::FixedSizeBinary(size),
859        }
860    }
861}
862
863impl ArrayGenerator for RandomFixedSizeBinaryGenerator {
864    fn generate(
865        &mut self,
866        length: RowCount,
867        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
868    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
869        let num_bytes = length.0 * self.size as u64;
870        let mut bytes = vec![0; num_bytes as usize];
871        rng.fill_bytes(&mut bytes);
872        Ok(Arc::new(FixedSizeBinaryArray::new(
873            self.size,
874            Buffer::from(bytes),
875            None,
876        )))
877    }
878
879    fn data_type(&self) -> &DataType {
880        &self.data_type
881    }
882
883    fn element_size_bytes(&self) -> Option<ByteCount> {
884        Some(ByteCount::from(self.size as u64))
885    }
886}
887
888#[derive(Debug)]
889pub struct RandomIntervalGenerator {
890    unit: IntervalUnit,
891    data_type: DataType,
892}
893
894impl RandomIntervalGenerator {
895    pub fn new(unit: IntervalUnit) -> Self {
896        Self {
897            unit,
898            data_type: DataType::Interval(unit),
899        }
900    }
901}
902
903impl ArrayGenerator for RandomIntervalGenerator {
904    fn generate(
905        &mut self,
906        length: RowCount,
907        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
908    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
909        match self.unit {
910            IntervalUnit::YearMonth => {
911                let months = (0..length.0)
912                    .map(|_| rng.random::<i32>())
913                    .collect::<Vec<_>>();
914                Ok(Arc::new(arrow_array::IntervalYearMonthArray::from(months)))
915            }
916            IntervalUnit::MonthDayNano => {
917                let day_time_array = (0..length.0)
918                    .map(|_| IntervalMonthDayNano::new(rng.random(), rng.random(), rng.random()))
919                    .collect::<Vec<_>>();
920                Ok(Arc::new(arrow_array::IntervalMonthDayNanoArray::from(
921                    day_time_array,
922                )))
923            }
924            IntervalUnit::DayTime => {
925                let day_time_array = (0..length.0)
926                    .map(|_| IntervalDayTime::new(rng.random(), rng.random()))
927                    .collect::<Vec<_>>();
928                Ok(Arc::new(arrow_array::IntervalDayTimeArray::from(
929                    day_time_array,
930                )))
931            }
932        }
933    }
934
935    fn data_type(&self) -> &DataType {
936        &self.data_type
937    }
938
939    fn element_size_bytes(&self) -> Option<ByteCount> {
940        Some(ByteCount::from(12))
941    }
942}
943#[derive(Debug)]
944pub struct RandomBinaryGenerator {
945    bytes_per_element: ByteCount,
946    scale_to_utf8: bool,
947    is_large: bool,
948    data_type: DataType,
949}
950
951impl RandomBinaryGenerator {
952    pub fn new(bytes_per_element: ByteCount, scale_to_utf8: bool, is_large: bool) -> Self {
953        Self {
954            bytes_per_element,
955            scale_to_utf8,
956            is_large,
957            data_type: match (scale_to_utf8, is_large) {
958                (false, false) => DataType::Binary,
959                (false, true) => DataType::LargeBinary,
960                (true, false) => DataType::Utf8,
961                (true, true) => DataType::LargeUtf8,
962            },
963        }
964    }
965}
966
967impl ArrayGenerator for RandomBinaryGenerator {
968    fn generate(
969        &mut self,
970        length: RowCount,
971        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
972    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
973        let mut bytes = vec![0; (self.bytes_per_element.0 * length.0) as usize];
974        rng.fill_bytes(&mut bytes);
975        if self.scale_to_utf8 {
976            // This doesn't give us the full UTF-8 range and it isn't statistically correct but
977            // it's fast and probably good enough for most cases
978            bytes = bytes.into_iter().map(|val| (val % 95) + 32).collect();
979        }
980        let bytes = Buffer::from(bytes);
981        if self.is_large {
982            let offsets = OffsetBuffer::from_lengths(iter::repeat_n(
983                self.bytes_per_element.0 as usize,
984                length.0 as usize,
985            ));
986            if self.scale_to_utf8 {
987                // This is safe because we are only using printable characters
988                unsafe {
989                    Ok(Arc::new(arrow_array::LargeStringArray::new_unchecked(
990                        offsets, bytes, None,
991                    )))
992                }
993            } else {
994                unsafe {
995                    Ok(Arc::new(arrow_array::LargeBinaryArray::new_unchecked(
996                        offsets, bytes, None,
997                    )))
998                }
999            }
1000        } else {
1001            let offsets = OffsetBuffer::from_lengths(iter::repeat_n(
1002                self.bytes_per_element.0 as usize,
1003                length.0 as usize,
1004            ));
1005            if self.scale_to_utf8 {
1006                // This is safe because we are only using printable characters
1007                unsafe {
1008                    Ok(Arc::new(arrow_array::StringArray::new_unchecked(
1009                        offsets, bytes, None,
1010                    )))
1011                }
1012            } else {
1013                unsafe {
1014                    Ok(Arc::new(arrow_array::BinaryArray::new_unchecked(
1015                        offsets, bytes, None,
1016                    )))
1017                }
1018            }
1019        }
1020    }
1021
1022    fn data_type(&self) -> &DataType {
1023        &self.data_type
1024    }
1025
1026    fn element_size_bytes(&self) -> Option<ByteCount> {
1027        // Not exactly correct since there are N + 1 4-byte offsets and this only counts N
1028        Some(ByteCount::from(
1029            self.bytes_per_element.0 + std::mem::size_of::<i32>() as u64,
1030        ))
1031    }
1032}
1033
1034/// Generate a sequence of strings with a prefix and a counter
1035///
1036/// For example, if the prefix is "user_" the strings will be "user_0", "user_1", ...
1037#[derive(Debug)]
1038pub struct PrefixPlusCounterGenerator {
1039    prefix: String,
1040    is_large: bool,
1041    data_type: DataType,
1042    current_counter: u64,
1043}
1044
1045impl PrefixPlusCounterGenerator {
1046    pub fn new(prefix: String, is_large: bool) -> Self {
1047        Self {
1048            prefix,
1049            is_large,
1050            data_type: if is_large {
1051                DataType::LargeUtf8
1052            } else {
1053                DataType::Utf8
1054            },
1055            current_counter: 0,
1056        }
1057    }
1058
1059    fn generate_values<T: OffsetSizeTrait>(
1060        &self,
1061        start: u64,
1062        num_values: u64,
1063    ) -> Result<Arc<dyn Array>, ArrowError> {
1064        let max_counter = start + num_values;
1065        let max_digits_per_counter = (max_counter as f64).log10().ceil() as u64;
1066        let max_bytes_per_str = max_digits_per_counter + self.prefix.len() as u64;
1067        let max_bytes = max_bytes_per_str * num_values;
1068        let mut builder =
1069            GenericStringBuilder::<T>::with_capacity(num_values as usize, max_bytes as usize);
1070        let mut word = String::with_capacity(max_bytes_per_str as usize);
1071        word.push_str(&self.prefix);
1072        for i in 0..num_values {
1073            let counter = start + i;
1074            word.truncate(self.prefix.len());
1075            word.push_str(&counter.to_string());
1076            builder.append_value(&word);
1077        }
1078        Ok(Arc::new(builder.finish()))
1079    }
1080}
1081
1082impl ArrayGenerator for PrefixPlusCounterGenerator {
1083    fn generate(
1084        &mut self,
1085        length: RowCount,
1086        _rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1087    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1088        let start = self.current_counter;
1089        self.current_counter += length.0;
1090        if self.is_large {
1091            self.generate_values::<i64>(start, length.0)
1092        } else {
1093            self.generate_values::<i32>(start, length.0)
1094        }
1095    }
1096
1097    fn data_type(&self) -> &DataType {
1098        &self.data_type
1099    }
1100
1101    fn element_size_bytes(&self) -> Option<ByteCount> {
1102        // It's not consistent
1103        None
1104    }
1105}
1106
1107/// Generate a sequence of binary strings with a prefix and a counter
1108///
1109/// The counter will be encoded (little-endian) as a u8, u16, u32, or u64 and added to the prefix
1110/// As long as more than 256 values are generated then the resulting array will have
1111/// variable width
1112#[derive(Debug)]
1113pub struct BinaryPrefixPlusCounterGenerator {
1114    prefix: Arc<[u8]>,
1115    is_large: bool,
1116    data_type: DataType,
1117    current_counter: u64,
1118}
1119
1120impl BinaryPrefixPlusCounterGenerator {
1121    pub fn new(prefix: Arc<[u8]>, is_large: bool) -> Self {
1122        Self {
1123            prefix,
1124            is_large,
1125            data_type: if is_large {
1126                DataType::LargeBinary
1127            } else {
1128                DataType::Binary
1129            },
1130            current_counter: 0,
1131        }
1132    }
1133
1134    fn generate_values<T: OffsetSizeTrait>(
1135        &self,
1136        start: u64,
1137        num_values: u64,
1138    ) -> Result<Arc<dyn Array>, ArrowError> {
1139        let max_bytes = (self.prefix.len() + std::mem::size_of::<u64>()) * num_values as usize;
1140        let mut builder = GenericBinaryBuilder::<T>::with_capacity(num_values as usize, max_bytes);
1141        let mut word = Vec::with_capacity(self.prefix.len() + std::mem::size_of::<u64>());
1142        word.extend_from_slice(&self.prefix);
1143        for i in 0..num_values {
1144            let counter = start + i;
1145            word.truncate(self.prefix.len());
1146            if counter < u8::MAX as u64 {
1147                word.push(counter as u8);
1148            } else if counter < u16::MAX as u64 {
1149                word.extend_from_slice(&(counter as u16).to_le_bytes());
1150            } else if counter < u32::MAX as u64 {
1151                word.extend_from_slice(&(counter as u32).to_le_bytes());
1152            } else {
1153                word.extend_from_slice(&counter.to_le_bytes());
1154            }
1155            builder.append_value(&word);
1156        }
1157        Ok(Arc::new(builder.finish()))
1158    }
1159}
1160
1161impl ArrayGenerator for BinaryPrefixPlusCounterGenerator {
1162    fn generate(
1163        &mut self,
1164        length: RowCount,
1165        _rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1166    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1167        let start = self.current_counter;
1168        self.current_counter += length.0;
1169        if self.is_large {
1170            self.generate_values::<i64>(start, length.0)
1171        } else {
1172            self.generate_values::<i32>(start, length.0)
1173        }
1174    }
1175
1176    fn data_type(&self) -> &DataType {
1177        &self.data_type
1178    }
1179
1180    fn element_size_bytes(&self) -> Option<ByteCount> {
1181        // It's not consistent
1182        None
1183    }
1184}
1185
1186// Common English stop words placed at the front to be sampled more frequently.
1187const STOP_WORDS: &[&str] = &[
1188    "a", "an", "and", "are", "as", "at", "be", "but", "by", "for", "if", "in", "into", "is", "it",
1189    "no", "not", "of", "on", "or", "such", "that", "the", "their", "then", "there", "these",
1190    "they", "this", "to", "was", "will", "with",
1191];
1192
1193const ENGLISH_WORDS: &[&str] = &[
1194    "ability",
1195    "able",
1196    "about",
1197    "above",
1198    "accept",
1199    "access",
1200    "account",
1201    "across",
1202    "action",
1203    "active",
1204    "activity",
1205    "actual",
1206    "address",
1207    "adjust",
1208    "admin",
1209    "advance",
1210    "agent",
1211    "align",
1212    "allow",
1213    "amount",
1214    "analysis",
1215    "answer",
1216    "application",
1217    "archive",
1218    "array",
1219    "asset",
1220    "async",
1221    "attribute",
1222    "available",
1223    "balance",
1224    "batch",
1225    "binary",
1226    "bitmap",
1227    "block",
1228    "branch",
1229    "buffer",
1230    "build",
1231    "cache",
1232    "capacity",
1233    "catalog",
1234    "change",
1235    "chunk",
1236    "client",
1237    "cluster",
1238    "column",
1239    "commit",
1240    "common",
1241    "compare",
1242    "compile",
1243    "compute",
1244    "condition",
1245    "config",
1246    "connect",
1247    "content",
1248    "context",
1249    "control",
1250    "convert",
1251    "copy",
1252    "core",
1253    "count",
1254    "create",
1255    "current",
1256    "cursor",
1257    "data",
1258    "dataset",
1259    "decode",
1260    "default",
1261    "delete",
1262    "delta",
1263    "depend",
1264    "derive",
1265    "design",
1266    "detail",
1267    "detect",
1268    "device",
1269    "direct",
1270    "display",
1271    "document",
1272    "domain",
1273    "drive",
1274    "dynamic",
1275    "encode",
1276    "engine",
1277    "error",
1278    "event",
1279    "example",
1280    "execute",
1281    "expand",
1282    "expect",
1283    "export",
1284    "extend",
1285    "feature",
1286    "field",
1287    "filter",
1288    "final",
1289    "finish",
1290    "format",
1291    "fragment",
1292    "future",
1293    "generate",
1294    "global",
1295    "group",
1296    "handle",
1297    "header",
1298    "index",
1299    "input",
1300    "insert",
1301    "inspect",
1302    "instance",
1303    "integer",
1304    "internal",
1305    "item",
1306    "join",
1307    "kernel",
1308    "large",
1309    "layer",
1310    "layout",
1311    "length",
1312    "level",
1313    "limit",
1314    "linear",
1315    "local",
1316    "logical",
1317    "lookup",
1318    "manage",
1319    "manifest",
1320    "memory",
1321    "merge",
1322    "metric",
1323    "model",
1324    "module",
1325    "namespace",
1326    "native",
1327    "node",
1328    "normal",
1329    "number",
1330    "object",
1331    "offset",
1332    "option",
1333    "output",
1334    "package",
1335    "page",
1336    "parallel",
1337    "parse",
1338    "partition",
1339    "pattern",
1340    "physical",
1341    "plan",
1342    "policy",
1343    "prefix",
1344    "prepare",
1345    "primary",
1346    "process",
1347    "profile",
1348    "project",
1349    "property",
1350    "query",
1351    "range",
1352    "reader",
1353    "record",
1354    "region",
1355    "registry",
1356    "request",
1357    "resolve",
1358    "resource",
1359    "result",
1360    "return",
1361    "row",
1362    "runtime",
1363    "scalar",
1364    "scan",
1365    "schema",
1366    "search",
1367    "segment",
1368    "select",
1369    "session",
1370    "setting",
1371    "source",
1372    "stable",
1373    "stage",
1374    "state",
1375    "static",
1376    "storage",
1377    "stream",
1378    "string",
1379    "struct",
1380    "table",
1381    "target",
1382    "task",
1383    "thread",
1384    "token",
1385    "trace",
1386    "transform",
1387    "type",
1388    "update",
1389    "upload",
1390    "value",
1391    "vector",
1392    "version",
1393    "view",
1394    "write",
1395    "writer",
1396];
1397
1398/// Word list with stop words at the front for Zipf sampling, computed once.
1399static SENTENCE_WORDS: LazyLock<Vec<&'static str>> = LazyLock::new(|| {
1400    let mut words = Vec::with_capacity(STOP_WORDS.len() + ENGLISH_WORDS.len());
1401    words.extend(STOP_WORDS.iter().copied());
1402    words.extend(ENGLISH_WORDS.iter().copied());
1403    words
1404});
1405
1406struct RandomSentenceGenerator {
1407    min_words: usize,
1408    max_words: usize,
1409    /// Zipf distribution for word selection (favors lower indices)
1410    zipf: Zipf<f64>,
1411    is_large: bool,
1412}
1413
1414impl std::fmt::Debug for RandomSentenceGenerator {
1415    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1416        f.debug_struct("RandomSentenceGenerator")
1417            .field("min_words", &self.min_words)
1418            .field("max_words", &self.max_words)
1419            .field("num_words", &SENTENCE_WORDS.len())
1420            .field("is_large", &self.is_large)
1421            .finish()
1422    }
1423}
1424
1425impl RandomSentenceGenerator {
1426    pub fn new(min_words: usize, max_words: usize, is_large: bool) -> Self {
1427        // Zipf distribution with exponent ~1.0 approximates natural language
1428        let zipf = Zipf::new(SENTENCE_WORDS.len() as f64, 1.0).unwrap();
1429
1430        Self {
1431            min_words,
1432            max_words,
1433            zipf,
1434            is_large,
1435        }
1436    }
1437}
1438
1439impl ArrayGenerator for RandomSentenceGenerator {
1440    fn generate(
1441        &mut self,
1442        length: RowCount,
1443        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1444    ) -> Result<Arc<dyn Array>, ArrowError> {
1445        let mut values = Vec::with_capacity(length.0 as usize);
1446
1447        for _ in 0..length.0 {
1448            let num_words = rng.random_range(self.min_words..=self.max_words);
1449            let sentence: String = (0..num_words)
1450                .map(|_| {
1451                    // Zipf returns 1-indexed values, subtract 1 for 0-indexed array
1452                    let idx = rng.sample(self.zipf) as usize - 1;
1453                    SENTENCE_WORDS[idx]
1454                })
1455                .collect::<Vec<_>>()
1456                .join(" ");
1457            values.push(sentence);
1458        }
1459
1460        if self.is_large {
1461            Ok(Arc::new(LargeStringArray::from(values)))
1462        } else {
1463            Ok(Arc::new(StringArray::from(values)))
1464        }
1465    }
1466
1467    fn data_type(&self) -> &DataType {
1468        if self.is_large {
1469            &DataType::LargeUtf8
1470        } else {
1471            &DataType::Utf8
1472        }
1473    }
1474
1475    fn element_size_bytes(&self) -> Option<ByteCount> {
1476        // Estimate average word length as 5, plus space
1477        // See https://arxiv.org/pdf/1208.6109
1478        let avg_word_length = 6;
1479        let avg_words = (self.min_words + self.max_words) / 2;
1480        Some(ByteCount::from((avg_word_length * avg_words) as u64))
1481    }
1482}
1483
1484#[derive(Debug)]
1485struct RandomWordGenerator {
1486    words: &'static [&'static str],
1487    is_large: bool,
1488}
1489
1490impl RandomWordGenerator {
1491    pub fn new(is_large: bool) -> Self {
1492        let words = ENGLISH_WORDS;
1493        Self { words, is_large }
1494    }
1495}
1496
1497impl ArrayGenerator for RandomWordGenerator {
1498    fn generate(
1499        &mut self,
1500        length: RowCount,
1501        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1502    ) -> Result<Arc<dyn Array>, ArrowError> {
1503        let mut values = Vec::with_capacity(length.0 as usize);
1504
1505        for _ in 0..length.0 {
1506            let word = self.words[rng.random_range(0..self.words.len())];
1507            values.push(word.to_string());
1508        }
1509
1510        if self.is_large {
1511            Ok(Arc::new(LargeStringArray::from(values)))
1512        } else {
1513            Ok(Arc::new(StringArray::from(values)))
1514        }
1515    }
1516
1517    fn data_type(&self) -> &DataType {
1518        if self.is_large {
1519            &DataType::LargeUtf8
1520        } else {
1521            &DataType::Utf8
1522        }
1523    }
1524
1525    fn element_size_bytes(&self) -> Option<ByteCount> {
1526        // Average English word length is ~5 characters
1527        Some(ByteCount::from(5))
1528    }
1529}
1530
1531#[derive(Debug)]
1532pub struct VariableRandomBinaryGenerator {
1533    lengths_gen: Box<dyn ArrayGenerator>,
1534    data_type: DataType,
1535}
1536
1537impl VariableRandomBinaryGenerator {
1538    pub fn new(min_bytes_per_element: ByteCount, max_bytes_per_element: ByteCount) -> Self {
1539        let lengths_dist = Uniform::new_inclusive(
1540            min_bytes_per_element.0 as i32,
1541            max_bytes_per_element.0 as i32,
1542        )
1543        .unwrap();
1544        let lengths_gen = rand_with_distribution::<Int32Type, Uniform<i32>>(lengths_dist);
1545
1546        Self {
1547            lengths_gen,
1548            data_type: DataType::Binary,
1549        }
1550    }
1551}
1552
1553impl ArrayGenerator for VariableRandomBinaryGenerator {
1554    fn generate(
1555        &mut self,
1556        length: RowCount,
1557        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1558    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1559        let lengths = self.lengths_gen.generate(length, rng)?;
1560        let lengths = lengths.as_primitive::<Int32Type>();
1561        let total_length = lengths.values().iter().map(|i| *i as usize).sum::<usize>();
1562        let offsets = OffsetBuffer::from_lengths(lengths.values().iter().map(|v| *v as usize));
1563        let mut bytes = vec![0; total_length];
1564        rng.fill_bytes(&mut bytes);
1565        let bytes = Buffer::from(bytes);
1566        Ok(Arc::new(BinaryArray::try_new(offsets, bytes, None)?))
1567    }
1568
1569    fn data_type(&self) -> &DataType {
1570        &self.data_type
1571    }
1572
1573    fn element_size_bytes(&self) -> Option<ByteCount> {
1574        None
1575    }
1576}
1577
1578pub struct CycleBinaryGenerator<T: ByteArrayType> {
1579    values: Vec<u8>,
1580    lengths: Vec<usize>,
1581    data_type: DataType,
1582    array_type: PhantomData<T>,
1583    width: Option<ByteCount>,
1584    idx: usize,
1585}
1586
1587impl<T: ByteArrayType> std::fmt::Debug for CycleBinaryGenerator<T> {
1588    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1589        f.debug_struct("CycleBinaryGenerator")
1590            .field("values", &self.values)
1591            .field("lengths", &self.lengths)
1592            .field("data_type", &self.data_type)
1593            .field("width", &self.width)
1594            .field("idx", &self.idx)
1595            .finish()
1596    }
1597}
1598
1599impl<T: ByteArrayType> CycleBinaryGenerator<T> {
1600    pub fn from_strings(values: &[&str]) -> Self {
1601        if values.is_empty() {
1602            panic!("Attempt to create a cycle generator with no values");
1603        }
1604        let lengths = values.iter().map(|s| s.len()).collect::<Vec<_>>();
1605        let typical_length = lengths[0];
1606        let width = if lengths.iter().all(|item| *item == typical_length) {
1607            Some(ByteCount::from(
1608                typical_length as u64 + std::mem::size_of::<i32>() as u64,
1609            ))
1610        } else {
1611            None
1612        };
1613        let values = values
1614            .iter()
1615            .flat_map(|s| s.as_bytes().iter().copied())
1616            .collect::<Vec<_>>();
1617        Self {
1618            values,
1619            lengths,
1620            data_type: T::DATA_TYPE,
1621            array_type: PhantomData,
1622            width,
1623            idx: 0,
1624        }
1625    }
1626}
1627
1628impl<T: ByteArrayType> ArrayGenerator for CycleBinaryGenerator<T> {
1629    fn generate(
1630        &mut self,
1631        length: RowCount,
1632        _: &mut rand_xoshiro::Xoshiro256PlusPlus,
1633    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1634        let lengths = self
1635            .lengths
1636            .iter()
1637            .copied()
1638            .cycle()
1639            .skip(self.idx)
1640            .take(length.0 as usize);
1641        let num_bytes = lengths.clone().sum();
1642        let byte_offset = self.lengths[0..self.idx].iter().sum();
1643        let bytes = self
1644            .values
1645            .iter()
1646            .cycle()
1647            .skip(byte_offset)
1648            .copied()
1649            .take(num_bytes)
1650            .collect::<Vec<_>>();
1651        let bytes = Buffer::from(bytes);
1652        let offsets = OffsetBuffer::from_lengths(lengths);
1653        self.idx = (self.idx + length.0 as usize) % self.lengths.len();
1654        Ok(Arc::new(arrow_array::GenericByteArray::<T>::new(
1655            offsets, bytes, None,
1656        )))
1657    }
1658
1659    fn data_type(&self) -> &DataType {
1660        &self.data_type
1661    }
1662
1663    fn element_size_bytes(&self) -> Option<ByteCount> {
1664        self.width
1665    }
1666}
1667
1668pub struct FixedBinaryGenerator<T: ByteArrayType> {
1669    value: Vec<u8>,
1670    data_type: DataType,
1671    array_type: PhantomData<T>,
1672}
1673
1674impl<T: ByteArrayType> std::fmt::Debug for FixedBinaryGenerator<T> {
1675    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1676        f.debug_struct("FixedBinaryGenerator")
1677            .field("value", &self.value)
1678            .field("data_type", &self.data_type)
1679            .finish()
1680    }
1681}
1682
1683impl<T: ByteArrayType> FixedBinaryGenerator<T> {
1684    pub fn new(value: Vec<u8>) -> Self {
1685        Self {
1686            value,
1687            data_type: T::DATA_TYPE,
1688            array_type: PhantomData,
1689        }
1690    }
1691}
1692
1693impl<T: ByteArrayType> ArrayGenerator for FixedBinaryGenerator<T> {
1694    fn generate(
1695        &mut self,
1696        length: RowCount,
1697        _: &mut rand_xoshiro::Xoshiro256PlusPlus,
1698    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1699        let bytes = Buffer::from(Vec::from_iter(
1700            self.value
1701                .iter()
1702                .cycle()
1703                .take((length.0 * self.value.len() as u64) as usize)
1704                .copied(),
1705        ));
1706        let offsets =
1707            OffsetBuffer::from_lengths(iter::repeat_n(self.value.len(), length.0 as usize));
1708        Ok(Arc::new(arrow_array::GenericByteArray::<T>::new(
1709            offsets, bytes, None,
1710        )))
1711    }
1712
1713    fn data_type(&self) -> &DataType {
1714        &self.data_type
1715    }
1716
1717    fn element_size_bytes(&self) -> Option<ByteCount> {
1718        // Not exactly correct since there are N + 1 4-byte offsets and this only counts N
1719        Some(ByteCount::from(
1720            self.value.len() as u64 + std::mem::size_of::<i32>() as u64,
1721        ))
1722    }
1723}
1724
1725pub struct DictionaryGenerator<K: ArrowDictionaryKeyType> {
1726    generator: Box<dyn ArrayGenerator>,
1727    data_type: DataType,
1728    key_type: PhantomData<K>,
1729    key_width: u64,
1730}
1731
1732impl<K: ArrowDictionaryKeyType> std::fmt::Debug for DictionaryGenerator<K> {
1733    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1734        f.debug_struct("DictionaryGenerator")
1735            .field("generator", &self.generator)
1736            .field("data_type", &self.data_type)
1737            .field("key_width", &self.key_width)
1738            .finish()
1739    }
1740}
1741
1742impl<K: ArrowDictionaryKeyType> DictionaryGenerator<K> {
1743    fn new(generator: Box<dyn ArrayGenerator>) -> Self {
1744        let key_type = Box::new(K::DATA_TYPE);
1745        let key_width = key_type
1746            .primitive_width()
1747            .expect("dictionary key types should have a known width")
1748            as u64;
1749        let val_type = Box::new(generator.data_type().clone());
1750        let dict_type = DataType::Dictionary(key_type, val_type);
1751        Self {
1752            generator,
1753            data_type: dict_type,
1754            key_type: PhantomData,
1755            key_width,
1756        }
1757    }
1758}
1759
1760impl<K: ArrowDictionaryKeyType + Send + Sync> ArrayGenerator for DictionaryGenerator<K> {
1761    fn generate(
1762        &mut self,
1763        length: RowCount,
1764        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1765    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
1766        let underlying = self.generator.generate(length, rng)?;
1767        arrow_cast::cast::cast(&underlying, &self.data_type)
1768    }
1769
1770    fn data_type(&self) -> &DataType {
1771        &self.data_type
1772    }
1773
1774    fn element_size_bytes(&self) -> Option<ByteCount> {
1775        self.generator
1776            .element_size_bytes()
1777            .map(|size_bytes| ByteCount::from(size_bytes.0 + self.key_width))
1778    }
1779}
1780
1781/// Generator that produces low-cardinality data by generating a fixed set of
1782/// unique values and then randomly selecting from them.
1783struct LowCardinalityGenerator {
1784    inner: Box<dyn ArrayGenerator>,
1785    cardinality: usize,
1786    /// Cached unique values, generated on first call
1787    unique_values: Option<Arc<dyn Array>>,
1788}
1789
1790impl std::fmt::Debug for LowCardinalityGenerator {
1791    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1792        f.debug_struct("LowCardinalityGenerator")
1793            .field("inner", &self.inner)
1794            .field("cardinality", &self.cardinality)
1795            .field("initialized", &self.unique_values.is_some())
1796            .finish()
1797    }
1798}
1799
1800impl LowCardinalityGenerator {
1801    fn new(inner: Box<dyn ArrayGenerator>, cardinality: usize) -> Self {
1802        Self {
1803            inner,
1804            cardinality,
1805            unique_values: None,
1806        }
1807    }
1808}
1809
1810impl ArrayGenerator for LowCardinalityGenerator {
1811    fn generate(
1812        &mut self,
1813        length: RowCount,
1814        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1815    ) -> Result<Arc<dyn Array>, ArrowError> {
1816        // Generate unique values on first call
1817        if self.unique_values.is_none() {
1818            self.unique_values = Some(
1819                self.inner
1820                    .generate(RowCount::from(self.cardinality as u64), rng)?,
1821            );
1822        }
1823
1824        let unique_values = self.unique_values.as_ref().unwrap();
1825
1826        // Generate random indices into the unique values
1827        let indices: Vec<usize> = (0..length.0)
1828            .map(|_| rng.random_range(0..self.cardinality))
1829            .collect();
1830
1831        // Use arrow's take to select values
1832        let indices_array =
1833            arrow_array::UInt32Array::from(indices.iter().map(|&i| i as u32).collect::<Vec<_>>());
1834        arrow::compute::take(unique_values.as_ref(), &indices_array, None)
1835            .map(|arr| arr as Arc<dyn Array>)
1836    }
1837
1838    fn data_type(&self) -> &DataType {
1839        self.inner.data_type()
1840    }
1841
1842    fn element_size_bytes(&self) -> Option<ByteCount> {
1843        self.inner.element_size_bytes()
1844    }
1845}
1846
1847#[derive(Debug)]
1848struct RandomListGenerator {
1849    field: Arc<Field>,
1850    child_field: Arc<Field>,
1851    items_gen: Box<dyn ArrayGenerator>,
1852    lengths_gen: Box<dyn ArrayGenerator>,
1853    is_large: bool,
1854}
1855
1856impl RandomListGenerator {
1857    // Creates a list generator that generates random lists with lengths between 0 and 10 (inclusive)
1858    fn new(items_gen: Box<dyn ArrayGenerator>, is_large: bool) -> Self {
1859        let child_field = Arc::new(Field::new("item", items_gen.data_type().clone(), true));
1860        let list_type = if is_large {
1861            DataType::LargeList(child_field.clone())
1862        } else {
1863            DataType::List(child_field.clone())
1864        };
1865        let field = Field::new("", list_type, true);
1866        let lengths_gen = if is_large {
1867            let lengths_dist = Uniform::new_inclusive(0, 10).unwrap();
1868            rand_with_distribution::<Int64Type, Uniform<i64>>(lengths_dist)
1869        } else {
1870            let lengths_dist = Uniform::new_inclusive(0, 10).unwrap();
1871            rand_with_distribution::<Int32Type, Uniform<i32>>(lengths_dist)
1872        };
1873        Self {
1874            field: Arc::new(field),
1875            child_field,
1876            items_gen,
1877            lengths_gen,
1878            is_large,
1879        }
1880    }
1881}
1882
1883impl ArrayGenerator for RandomListGenerator {
1884    fn generate(
1885        &mut self,
1886        length: RowCount,
1887        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1888    ) -> Result<Arc<dyn Array>, ArrowError> {
1889        let lengths = self.lengths_gen.generate(length, rng)?;
1890        if self.is_large {
1891            let lengths = lengths.as_primitive::<Int64Type>();
1892            let total_length = lengths.values().iter().sum::<i64>() as u64;
1893            let offsets = OffsetBuffer::from_lengths(lengths.values().iter().map(|v| *v as usize));
1894            let items = self.items_gen.generate(RowCount::from(total_length), rng)?;
1895            Ok(Arc::new(LargeListArray::try_new(
1896                self.child_field.clone(),
1897                offsets,
1898                items,
1899                None,
1900            )?))
1901        } else {
1902            let lengths = lengths.as_primitive::<Int32Type>();
1903            let total_length = lengths.values().iter().sum::<i32>() as u64;
1904            let offsets = OffsetBuffer::from_lengths(lengths.values().iter().map(|v| *v as usize));
1905            let items = self.items_gen.generate(RowCount::from(total_length), rng)?;
1906            Ok(Arc::new(ListArray::try_new(
1907                self.child_field.clone(),
1908                offsets,
1909                items,
1910                None,
1911            )?))
1912        }
1913    }
1914
1915    fn data_type(&self) -> &DataType {
1916        self.field.data_type()
1917    }
1918
1919    fn element_size_bytes(&self) -> Option<ByteCount> {
1920        None
1921    }
1922}
1923
1924/// Generates random map arrays where each map has 0-4 entries.
1925#[derive(Debug)]
1926struct RandomMapGenerator {
1927    field: Arc<Field>,
1928    entries_field: Arc<Field>,
1929    keys_gen: Box<dyn ArrayGenerator>,
1930    values_gen: Box<dyn ArrayGenerator>,
1931    lengths_gen: Box<dyn ArrayGenerator>,
1932}
1933
1934impl RandomMapGenerator {
1935    fn new(keys_gen: Box<dyn ArrayGenerator>, values_gen: Box<dyn ArrayGenerator>) -> Self {
1936        let entries_fields = Fields::from(vec![
1937            Field::new("keys", keys_gen.data_type().clone(), false),
1938            Field::new("values", values_gen.data_type().clone(), true),
1939        ]);
1940        let entries_field = Arc::new(Field::new(
1941            "entries",
1942            DataType::Struct(entries_fields),
1943            false,
1944        ));
1945        let map_type = DataType::Map(entries_field.clone(), false);
1946        let field = Arc::new(Field::new("", map_type, true));
1947        let lengths_dist = Uniform::new_inclusive(0_i32, 4).unwrap();
1948        let lengths_gen = rand_with_distribution::<Int32Type, Uniform<i32>>(lengths_dist);
1949
1950        Self {
1951            field,
1952            entries_field,
1953            keys_gen,
1954            values_gen,
1955            lengths_gen,
1956        }
1957    }
1958}
1959
1960impl ArrayGenerator for RandomMapGenerator {
1961    fn generate(
1962        &mut self,
1963        length: RowCount,
1964        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
1965    ) -> Result<Arc<dyn Array>, ArrowError> {
1966        let lengths = self.lengths_gen.generate(length, rng)?;
1967        let lengths = lengths.as_primitive::<Int32Type>();
1968        let total_entries = lengths.values().iter().sum::<i32>() as u64;
1969        let offsets = OffsetBuffer::from_lengths(lengths.values().iter().map(|v| *v as usize));
1970
1971        let keys = self.keys_gen.generate(RowCount::from(total_entries), rng)?;
1972        let values = self
1973            .values_gen
1974            .generate(RowCount::from(total_entries), rng)?;
1975
1976        let entries = StructArray::new(
1977            Fields::from(vec![
1978                Field::new("keys", keys.data_type().clone(), false),
1979                Field::new("values", values.data_type().clone(), true),
1980            ]),
1981            vec![keys, values],
1982            None,
1983        );
1984
1985        Ok(Arc::new(MapArray::try_new(
1986            self.entries_field.clone(),
1987            offsets,
1988            entries,
1989            None,
1990            false,
1991        )?))
1992    }
1993
1994    fn data_type(&self) -> &DataType {
1995        self.field.data_type()
1996    }
1997
1998    fn element_size_bytes(&self) -> Option<ByteCount> {
1999        None
2000    }
2001}
2002
2003#[derive(Debug)]
2004struct NullArrayGenerator {}
2005
2006impl ArrayGenerator for NullArrayGenerator {
2007    fn generate(
2008        &mut self,
2009        length: RowCount,
2010        _: &mut rand_xoshiro::Xoshiro256PlusPlus,
2011    ) -> Result<Arc<dyn Array>, ArrowError> {
2012        Ok(Arc::new(NullArray::new(length.0 as usize)))
2013    }
2014
2015    fn data_type(&self) -> &DataType {
2016        &DataType::Null
2017    }
2018
2019    fn element_size_bytes(&self) -> Option<ByteCount> {
2020        None
2021    }
2022}
2023
2024/// Generates 2 dimensional vectors along the unit circle, with a configurable number of steps per circle.
2025#[derive(Debug)]
2026struct RadialStepGenerator {
2027    num_steps_per_circle: u32,
2028    data_field: Arc<Field>,
2029    data_type: DataType,
2030    current_step: u32,
2031}
2032
2033impl RadialStepGenerator {
2034    fn new(num_steps_per_circle: u32) -> Self {
2035        let data_field = Arc::new(Field::new("item", DataType::Float32, false));
2036        let data_type = DataType::FixedSizeList(data_field.clone(), 2);
2037        Self {
2038            num_steps_per_circle,
2039            data_field,
2040            data_type,
2041            current_step: 0,
2042        }
2043    }
2044}
2045
2046impl ArrayGenerator for RadialStepGenerator {
2047    fn generate(
2048        &mut self,
2049        length: RowCount,
2050        _rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
2051    ) -> Result<Arc<dyn Array>, ArrowError> {
2052        let mut values_builder = Float32Builder::with_capacity(length.0 as usize * 2);
2053        for _ in 0..length.0 {
2054            let angle = (self.current_step as f32) / (self.num_steps_per_circle as f32)
2055                * 2.0
2056                * std::f32::consts::PI;
2057            values_builder.append_value(angle.cos());
2058            values_builder.append_value(angle.sin());
2059            self.current_step = (self.current_step + 1) % self.num_steps_per_circle;
2060        }
2061        let values = values_builder.finish();
2062        let vectors =
2063            FixedSizeListArray::try_new(self.data_field.clone(), 2, Arc::new(values), None)?;
2064        Ok(Arc::new(vectors))
2065    }
2066
2067    fn data_type(&self) -> &DataType {
2068        &self.data_type
2069    }
2070
2071    fn element_size_bytes(&self) -> Option<ByteCount> {
2072        Some(ByteCount::from(8))
2073    }
2074}
2075
2076/// Cycles through a set of centroids, adding noise to each point
2077#[derive(Debug)]
2078struct JitterCentroidsGenerator {
2079    centroids: Float32Array,
2080    dimension: u32,
2081    noise_level: f32,
2082    data_type: DataType,
2083    data_field: Arc<Field>,
2084
2085    offset: usize,
2086}
2087
2088impl JitterCentroidsGenerator {
2089    fn try_new(centroids: Arc<dyn Array>, noise_level: f32) -> Result<Self, ArrowError> {
2090        let DataType::FixedSizeList(values_field, dimension) = centroids.data_type() else {
2091            return Err(ArrowError::InvalidArgumentError(
2092                "Centroids must be a FixedSizeList".to_string(),
2093            ));
2094        };
2095        if values_field.data_type() != &DataType::Float32 {
2096            return Err(ArrowError::InvalidArgumentError(
2097                "Centroids values must be a Float32".to_string(),
2098            ));
2099        }
2100        let data_type = DataType::FixedSizeList(values_field.clone(), *dimension);
2101        Ok(Self {
2102            centroids: centroids
2103                .as_fixed_size_list()
2104                .values()
2105                .as_primitive::<Float32Type>()
2106                .clone(),
2107            dimension: *dimension as u32,
2108            noise_level,
2109            data_type,
2110            data_field: values_field.clone(),
2111            offset: 0,
2112        })
2113    }
2114}
2115
2116impl ArrayGenerator for JitterCentroidsGenerator {
2117    fn generate(
2118        &mut self,
2119        length: RowCount,
2120        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
2121    ) -> Result<Arc<dyn Array>, ArrowError> {
2122        let mut values_builder =
2123            Float32Builder::with_capacity(length.0 as usize * self.dimension as usize);
2124        for _ in 0..length.0 {
2125            // Generate random N dimensional point
2126            let mut noise = (0..self.dimension as usize)
2127                .map(|_| rng.random::<f32>())
2128                .collect::<Vec<_>>();
2129            // Scale point to noise_level length
2130            let scale = self.noise_level / noise.iter().map(|v| v * v).sum::<f32>().sqrt();
2131            noise.iter_mut().for_each(|v| *v *= scale);
2132
2133            // Add noise to centroid and store in values
2134            for (i, noise) in noise.into_iter().enumerate() {
2135                let centroid_val = self.centroids.value(self.offset + i);
2136                let jittered_val = centroid_val + noise;
2137                values_builder.append_value(jittered_val);
2138            }
2139            // Advance to next centroid
2140            self.offset = (self.offset + self.dimension as usize) % self.centroids.len();
2141        }
2142        let values = values_builder.finish();
2143        let vectors = FixedSizeListArray::try_new(
2144            self.data_field.clone(),
2145            self.dimension as i32,
2146            Arc::new(values),
2147            None,
2148        )?;
2149        Ok(Arc::new(vectors))
2150    }
2151
2152    fn data_type(&self) -> &DataType {
2153        &self.data_type
2154    }
2155
2156    fn element_size_bytes(&self) -> Option<ByteCount> {
2157        Some(ByteCount::from(self.dimension as u64 * 4))
2158    }
2159}
2160#[derive(Debug)]
2161struct RandomStructGenerator {
2162    fields: Fields,
2163    data_type: DataType,
2164    child_gens: Vec<Box<dyn ArrayGenerator>>,
2165}
2166
2167impl RandomStructGenerator {
2168    fn new(fields: Fields, child_gens: Vec<Box<dyn ArrayGenerator>>) -> Self {
2169        let data_type = DataType::Struct(fields.clone());
2170        Self {
2171            fields,
2172            data_type,
2173            child_gens,
2174        }
2175    }
2176}
2177
2178impl ArrayGenerator for RandomStructGenerator {
2179    fn generate(
2180        &mut self,
2181        length: RowCount,
2182        rng: &mut rand_xoshiro::Xoshiro256PlusPlus,
2183    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
2184        if self.child_gens.is_empty() {
2185            // Have to create empty struct arrays specially to ensure they have the correct
2186            // row count
2187            let struct_arr = StructArray::new_empty_fields(length.0 as usize, None);
2188            return Ok(Arc::new(struct_arr));
2189        }
2190        let child_arrays = self
2191            .child_gens
2192            .iter_mut()
2193            .map(|genn| genn.generate(length, rng))
2194            .collect::<Result<Vec<_>, ArrowError>>()?;
2195        let struct_arr = StructArray::new(self.fields.clone(), child_arrays, None);
2196        Ok(Arc::new(struct_arr))
2197    }
2198
2199    fn data_type(&self) -> &DataType {
2200        &self.data_type
2201    }
2202
2203    fn element_size_bytes(&self) -> Option<ByteCount> {
2204        let mut sum = 0;
2205        for child_gen in &self.child_gens {
2206            sum += child_gen.element_size_bytes()?.0;
2207        }
2208        Some(ByteCount::from(sum))
2209    }
2210}
2211
2212/// A RecordBatchReader that generates batches of the given size from the given array generators
2213pub struct FixedSizeBatchGenerator {
2214    rng: rand_xoshiro::Xoshiro256PlusPlus,
2215    generators: Vec<Box<dyn ArrayGenerator>>,
2216    batch_size: RowCount,
2217    num_batches: BatchCount,
2218    schema: SchemaRef,
2219}
2220
2221impl FixedSizeBatchGenerator {
2222    fn new(
2223        generators: Vec<(Option<String>, Box<dyn ArrayGenerator>)>,
2224        batch_size: RowCount,
2225        num_batches: BatchCount,
2226        seed: Option<Seed>,
2227        default_null_probability: Option<f64>,
2228    ) -> Self {
2229        let mut fields = Vec::with_capacity(generators.len());
2230        for (field_index, field_gen) in generators.iter().enumerate() {
2231            let (name, genn) = field_gen;
2232            let default_name = format!("field_{}", field_index);
2233            let name = name.clone().unwrap_or(default_name);
2234            let mut field = Field::new(name, genn.data_type().clone(), true);
2235            if let Some(metadata) = genn.metadata() {
2236                field = field.with_metadata(metadata);
2237            }
2238            fields.push(field);
2239        }
2240        let mut generators = generators
2241            .into_iter()
2242            .map(|(_, genn)| genn)
2243            .collect::<Vec<_>>();
2244        if let Some(null_probability) = default_null_probability {
2245            generators = generators
2246                .into_iter()
2247                .map(|genn| genn.with_random_nulls(null_probability))
2248                .collect();
2249        }
2250        let schema = Arc::new(Schema::new(fields));
2251        Self {
2252            rng: rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(
2253                seed.map(|s| s.0).unwrap_or(DEFAULT_SEED.0),
2254            ),
2255            generators,
2256            batch_size,
2257            num_batches,
2258            schema,
2259        }
2260    }
2261
2262    fn gen_next(&mut self) -> Result<RecordBatch, ArrowError> {
2263        let mut arrays = Vec::with_capacity(self.generators.len());
2264        for genn in self.generators.iter_mut() {
2265            let arr = genn.generate(self.batch_size, &mut self.rng)?;
2266            arrays.push(arr);
2267        }
2268        self.num_batches.0 -= 1;
2269        Ok(RecordBatch::try_new_with_options(
2270            self.schema.clone(),
2271            arrays,
2272            &RecordBatchOptions::new().with_row_count(Some(self.batch_size.0 as usize)),
2273        )
2274        .unwrap())
2275    }
2276}
2277
2278impl Iterator for FixedSizeBatchGenerator {
2279    type Item = Result<RecordBatch, ArrowError>;
2280
2281    fn next(&mut self) -> Option<Self::Item> {
2282        if self.num_batches.0 == 0 {
2283            return None;
2284        }
2285        Some(self.gen_next())
2286    }
2287}
2288
2289impl RecordBatchReader for FixedSizeBatchGenerator {
2290    fn schema(&self) -> SchemaRef {
2291        self.schema.clone()
2292    }
2293}
2294
2295/// A builder to create a record batch reader with generated data
2296///
2297/// This type is meant to be used in a fluent builder style to define the schema and generators
2298/// for a record batch reader.
2299#[derive(Default)]
2300pub struct BatchGeneratorBuilder {
2301    generators: Vec<(Option<String>, Box<dyn ArrayGenerator>)>,
2302    default_null_probability: Option<f64>,
2303    seed: Option<Seed>,
2304}
2305
2306pub enum RoundingBehavior {
2307    ExactOrErr,
2308    RoundUp,
2309    RoundDown,
2310}
2311
2312impl BatchGeneratorBuilder {
2313    /// Create a new BatchGeneratorBuilder with a default random seed
2314    pub fn new() -> Self {
2315        Default::default()
2316    }
2317
2318    /// Create a new BatchGeneratorBuilder with the given seed
2319    pub fn new_with_seed(seed: Seed) -> Self {
2320        Self {
2321            seed: Some(seed),
2322            ..Default::default()
2323        }
2324    }
2325
2326    /// Adds a new column to the generator
2327    ///
2328    /// See [`crate::generator::array`] for methods to create generators
2329    pub fn col(mut self, name: impl Into<String>, genn: Box<dyn ArrayGenerator>) -> Self {
2330        self.generators.push((Some(name.into()), genn));
2331        self
2332    }
2333
2334    /// Adds a new column to the generator with a generated unique name
2335    ///
2336    /// See [`crate::generator::array`] for methods to create generators
2337    pub fn anon_col(mut self, genn: Box<dyn ArrayGenerator>) -> Self {
2338        self.generators.push((None, genn));
2339        self
2340    }
2341
2342    pub fn into_batch_rows(self, batch_size: RowCount) -> Result<RecordBatch, ArrowError> {
2343        let mut reader = self.into_reader_rows(batch_size, BatchCount::from(1));
2344        reader
2345            .next()
2346            .expect("Asked for 1 batch but reader was empty")
2347    }
2348
2349    pub fn into_batch_bytes(
2350        self,
2351        batch_size: ByteCount,
2352        rounding: RoundingBehavior,
2353    ) -> Result<RecordBatch, ArrowError> {
2354        let mut reader = self.into_reader_bytes(batch_size, BatchCount::from(1), rounding)?;
2355        reader
2356            .next()
2357            .expect("Asked for 1 batch but reader was empty")
2358    }
2359
2360    /// Create a RecordBatchReader that generates batches of the given size (in rows)
2361    pub fn into_reader_rows(
2362        self,
2363        batch_size: RowCount,
2364        num_batches: BatchCount,
2365    ) -> impl RecordBatchReader {
2366        FixedSizeBatchGenerator::new(
2367            self.generators,
2368            batch_size,
2369            num_batches,
2370            self.seed,
2371            self.default_null_probability,
2372        )
2373    }
2374
2375    pub fn into_reader_stream(
2376        self,
2377        batch_size: RowCount,
2378        num_batches: BatchCount,
2379    ) -> (
2380        BoxStream<'static, Result<RecordBatch, ArrowError>>,
2381        Arc<Schema>,
2382    ) {
2383        // TODO: this is pretty lazy and could be optimized
2384        let reader = self.into_reader_rows(batch_size, num_batches);
2385        let schema = reader.schema();
2386        let batches = reader.collect::<Vec<_>>();
2387        (futures::stream::iter(batches).boxed(), schema)
2388    }
2389
2390    /// Create a RecordBatchReader that generates batches of the given size (in bytes)
2391    pub fn into_reader_bytes(
2392        self,
2393        batch_size_bytes: ByteCount,
2394        num_batches: BatchCount,
2395        rounding: RoundingBehavior,
2396    ) -> Result<impl RecordBatchReader, ArrowError> {
2397        let bytes_per_row = self
2398            .generators
2399            .iter()
2400            .map(|genn| genn.1.element_size_bytes().map(|byte_count| byte_count.0).ok_or(
2401                        ArrowError::NotYetImplemented("The function into_reader_bytes currently requires each array generator to have a fixed element size".to_string())
2402                )
2403            )
2404            .sum::<Result<u64, ArrowError>>()?;
2405        let mut num_rows = RowCount::from(batch_size_bytes.0 / bytes_per_row);
2406        if !batch_size_bytes.0.is_multiple_of(bytes_per_row) {
2407            match rounding {
2408                RoundingBehavior::ExactOrErr => {
2409                    return Err(ArrowError::NotYetImplemented(format!(
2410                        "Exact rounding requested but not possible.  Batch size requested {}, row size: {}",
2411                        batch_size_bytes.0, bytes_per_row
2412                    )));
2413                }
2414                RoundingBehavior::RoundUp => {
2415                    num_rows = RowCount::from(num_rows.0 + 1);
2416                }
2417                RoundingBehavior::RoundDown => (),
2418            }
2419        }
2420        Ok(self.into_reader_rows(num_rows, num_batches))
2421    }
2422
2423    /// Set the seed for the generator
2424    pub fn with_seed(mut self, seed: Seed) -> Self {
2425        self.seed = Some(seed);
2426        self
2427    }
2428
2429    /// Adds nulls (with the given probability) to all columns
2430    pub fn with_random_nulls(&mut self, default_null_probability: f64) {
2431        self.default_null_probability = Some(default_null_probability);
2432    }
2433}
2434
2435/// Factory for creating a single random array
2436pub struct ArrayGeneratorBuilder {
2437    generator: Box<dyn ArrayGenerator>,
2438    seed: Option<Seed>,
2439}
2440
2441impl ArrayGeneratorBuilder {
2442    fn new(generator: Box<dyn ArrayGenerator>) -> Self {
2443        Self {
2444            generator,
2445            seed: None,
2446        }
2447    }
2448
2449    /// Use the given seed for the generator
2450    pub fn with_seed(mut self, seed: Seed) -> Self {
2451        self.seed = Some(seed);
2452        self
2453    }
2454
2455    /// Generate a single array with the given length
2456    pub fn into_array_rows(
2457        mut self,
2458        length: RowCount,
2459    ) -> Result<Arc<dyn arrow_array::Array>, ArrowError> {
2460        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(
2461            self.seed.map(|s| s.0).unwrap_or(DEFAULT_SEED.0),
2462        );
2463        self.generator.generate(length, &mut rng)
2464    }
2465}
2466
2467const MS_PER_DAY: i64 = 86400000;
2468
2469pub mod array {
2470
2471    use arrow::datatypes::{Int8Type, Int16Type, Int64Type};
2472    use arrow_array::types::{
2473        Decimal128Type, Decimal256Type, DurationMicrosecondType, DurationMillisecondType,
2474        DurationNanosecondType, DurationSecondType, Float16Type, Float32Type, Float64Type,
2475        UInt8Type, UInt16Type, UInt32Type, UInt64Type,
2476    };
2477    use arrow_array::{
2478        ArrowNativeTypeOp, BooleanArray, Date32Array, Date64Array, Time32MillisecondArray,
2479        Time32SecondArray, Time64MicrosecondArray, Time64NanosecondArray,
2480        TimestampMicrosecondArray, TimestampMillisecondArray, TimestampNanosecondArray,
2481        TimestampSecondArray,
2482    };
2483    use arrow_schema::{IntervalUnit, TimeUnit};
2484    use chrono::Utc;
2485    use rand::prelude::Distribution;
2486
2487    use super::*;
2488
2489    /// Create a generator of vectors by continuously calling the given generator
2490    ///
2491    /// For example, given a step generator and a dimension of 3 this will generate vectors like
2492    /// [0, 1, 2], [3, 4, 5], [6, 7, 8], ...
2493    pub fn cycle_vec(
2494        generator: Box<dyn ArrayGenerator>,
2495        dimension: Dimension,
2496    ) -> Box<dyn ArrayGenerator> {
2497        Box::new(CycleVectorGenerator::new(generator, dimension))
2498    }
2499
2500    /// Create a generator of list vectors by continuously calling the given generator
2501    ///
2502    /// The lists will have lengths uniformly distributed between `min_list_size` (inclusive) and
2503    /// `max_list_size` (exclusive).
2504    pub fn cycle_vec_var(
2505        generator: Box<dyn ArrayGenerator>,
2506        min_list_size: Dimension,
2507        max_list_size: Dimension,
2508    ) -> Box<dyn ArrayGenerator> {
2509        Box::new(CycleListGenerator::new(
2510            generator,
2511            min_list_size,
2512            max_list_size,
2513        ))
2514    }
2515
2516    /// Create a generator of vectors around unit circle
2517    ///
2518    /// Vectors will be equally spaced around the unit circle so that there are num_steps
2519    /// vectors per circle.
2520    pub fn cycle_unit_circle(num_steps: u32) -> Box<dyn ArrayGenerator> {
2521        Box::new(RadialStepGenerator::new(num_steps))
2522    }
2523
2524    /// Create a generator of vectors by cycling through a given set of vectors
2525    ///
2526    /// Each value will be spaced in slightly away from the previous value on a ball of radius jitter
2527    pub fn jitter_centroids(centroids: Arc<dyn Array>, jitter: f32) -> Box<dyn ArrayGenerator> {
2528        Box::new(JitterCentroidsGenerator::try_new(centroids, jitter).unwrap())
2529    }
2530
2531    /// Create a generator from a vector of values
2532    ///
2533    /// If more rows are requested than the length of values then it will restart
2534    /// from the beginning of the vector.
2535    pub fn cycle<DataType>(values: Vec<DataType::Native>) -> Box<dyn ArrayGenerator>
2536    where
2537        DataType::Native: Copy + 'static,
2538        DataType: ArrowPrimitiveType,
2539        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2540    {
2541        let mut values_idx = 0;
2542        Box::new(
2543            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2544                DataType::DATA_TYPE,
2545                move |_| {
2546                    let y = values[values_idx];
2547                    values_idx = (values_idx + 1) % values.len();
2548                    y
2549                },
2550                1,
2551                DataType::DATA_TYPE
2552                    .primitive_width()
2553                    .map(|width| ByteCount::from(width as u64))
2554                    .expect("Primitive types should have a fixed width"),
2555            ),
2556        )
2557    }
2558
2559    /// Create a generator from a vector of booleans
2560    ///
2561    /// If more rows are requested than the length of values then it will restart from
2562    /// the beginning of the vector
2563    pub fn cycle_bool(values: Vec<bool>) -> Box<dyn ArrayGenerator> {
2564        let mut values_idx = 0;
2565        Box::new(FnGen::<bool, BooleanArray, _>::new_unknown_size(
2566            DataType::Boolean,
2567            move |_| {
2568                let val = values[values_idx];
2569                values_idx = (values_idx + 1) % values.len();
2570                val
2571            },
2572            1,
2573        ))
2574    }
2575
2576    /// Create a generator that starts at 0 and increments by 1 for each element
2577    pub fn step<DataType>() -> Box<dyn ArrayGenerator>
2578    where
2579        DataType::Native: Copy + Default + std::ops::AddAssign<DataType::Native> + 'static,
2580        DataType: ArrowPrimitiveType,
2581        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2582    {
2583        let mut x = DataType::Native::default();
2584        Box::new(
2585            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2586                DataType::DATA_TYPE,
2587                move |_| {
2588                    let y = x;
2589                    x += DataType::Native::ONE;
2590                    y
2591                },
2592                1,
2593                DataType::DATA_TYPE
2594                    .primitive_width()
2595                    .map(|width| ByteCount::from(width as u64))
2596                    .expect("Primitive types should have a fixed width"),
2597            ),
2598        )
2599    }
2600
2601    pub fn blob() -> Box<dyn ArrayGenerator> {
2602        let mut blob_meta = HashMap::new();
2603        blob_meta.insert("lance-encoding:blob".to_string(), "true".to_string());
2604        rand_fixedbin(ByteCount::from(4 * 1024 * 1024), true).with_metadata(blob_meta)
2605    }
2606
2607    /// Create a generator that starts at a given value and increments by a given step for each element
2608    pub fn step_custom<DataType>(
2609        start: DataType::Native,
2610        step: DataType::Native,
2611    ) -> Box<dyn ArrayGenerator>
2612    where
2613        DataType::Native: Copy + Default + std::ops::AddAssign<DataType::Native> + 'static,
2614        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2615        DataType: ArrowPrimitiveType,
2616    {
2617        let mut x = start;
2618        Box::new(
2619            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2620                DataType::DATA_TYPE,
2621                move |_| {
2622                    let y = x;
2623                    x += step;
2624                    y
2625                },
2626                1,
2627                DataType::DATA_TYPE
2628                    .primitive_width()
2629                    .map(|width| ByteCount::from(width as u64))
2630                    .expect("Primitive types should have a fixed width"),
2631            ),
2632        )
2633    }
2634
2635    /// Create a generator that fills each element with the given primitive value
2636    pub fn fill<DataType>(value: DataType::Native) -> Box<dyn ArrayGenerator>
2637    where
2638        DataType::Native: Copy + 'static,
2639        DataType: ArrowPrimitiveType,
2640        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2641    {
2642        Box::new(
2643            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2644                DataType::DATA_TYPE,
2645                move |_| value,
2646                1,
2647                DataType::DATA_TYPE
2648                    .primitive_width()
2649                    .map(|width| ByteCount::from(width as u64))
2650                    .expect("Primitive types should have a fixed width"),
2651            ),
2652        )
2653    }
2654
2655    /// Create a generator that fills each element with the given binary value
2656    pub fn fill_varbin(value: Vec<u8>) -> Box<dyn ArrayGenerator> {
2657        Box::new(FixedBinaryGenerator::<BinaryType>::new(value))
2658    }
2659
2660    /// Create a generator that fills each element with the given string value
2661    pub fn fill_utf8(value: String) -> Box<dyn ArrayGenerator> {
2662        Box::new(FixedBinaryGenerator::<Utf8Type>::new(value.into_bytes()))
2663    }
2664
2665    pub fn cycle_utf8_literals(values: &[&'static str]) -> Box<dyn ArrayGenerator> {
2666        Box::new(CycleBinaryGenerator::<Utf8Type>::from_strings(values))
2667    }
2668
2669    /// Create a generator of primitive values that are randomly sampled from the entire range available for the value
2670    pub fn rand<DataType>() -> Box<dyn ArrayGenerator>
2671    where
2672        DataType::Native: Copy + 'static,
2673        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2674        DataType: ArrowPrimitiveType,
2675        rand::distr::StandardUniform: rand::distr::Distribution<DataType::Native>,
2676    {
2677        Box::new(
2678            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2679                DataType::DATA_TYPE,
2680                move |rng| rng.random(),
2681                1,
2682                DataType::DATA_TYPE
2683                    .primitive_width()
2684                    .map(|width| ByteCount::from(width as u64))
2685                    .expect("Primitive types should have a fixed width"),
2686            ),
2687        )
2688    }
2689
2690    /// Create a generator of primitive values that are randomly sampled from the entire range available for the value
2691    pub fn rand_with_distribution<
2692        DataType,
2693        Dist: rand::distr::Distribution<DataType::Native> + Clone + Send + Sync + 'static,
2694    >(
2695        dist: Dist,
2696    ) -> Box<dyn ArrayGenerator>
2697    where
2698        DataType::Native: Copy + 'static,
2699        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2700        DataType: ArrowPrimitiveType,
2701    {
2702        Box::new(
2703            FnGen::<DataType::Native, PrimitiveArray<DataType>, _>::new_known_size(
2704                DataType::DATA_TYPE,
2705                move |rng| rng.sample(dist.clone()),
2706                1,
2707                DataType::DATA_TYPE
2708                    .primitive_width()
2709                    .map(|width| ByteCount::from(width as u64))
2710                    .expect("Primitive types should have a fixed width"),
2711            ),
2712        )
2713    }
2714
2715    /// Create a generator of 1d vectors (of a primitive type) consisting of randomly sampled primitive values
2716    pub fn rand_vec<DataType>(dimension: Dimension) -> Box<dyn ArrayGenerator>
2717    where
2718        DataType::Native: Copy + 'static,
2719        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2720        DataType: ArrowPrimitiveType,
2721        rand::distr::StandardUniform: rand::distr::Distribution<DataType::Native>,
2722    {
2723        let underlying = rand::<DataType>();
2724        cycle_vec(underlying, dimension)
2725    }
2726
2727    /// Create a generator of 1d vectors (of a primitive type) consisting of randomly sampled nullable values
2728    pub fn rand_vec_nullable<DataType>(
2729        dimension: Dimension,
2730        null_probability: f64,
2731    ) -> Box<dyn ArrayGenerator>
2732    where
2733        DataType::Native: Copy + 'static,
2734        PrimitiveArray<DataType>: From<Vec<DataType::Native>> + 'static,
2735        DataType: ArrowPrimitiveType,
2736        rand::distr::StandardUniform: rand::distr::Distribution<DataType::Native>,
2737    {
2738        let underlying = rand::<DataType>().with_random_nulls(null_probability);
2739        cycle_vec(underlying, dimension)
2740    }
2741
2742    /// Create a generator of randomly sampled time32 values covering the entire
2743    /// range of 1 day
2744    pub fn rand_time32(resolution: &TimeUnit) -> Box<dyn ArrayGenerator> {
2745        let start = 0;
2746        let end = match resolution {
2747            TimeUnit::Second => 86_400,
2748            TimeUnit::Millisecond => 86_400_000,
2749            _ => panic!(),
2750        };
2751
2752        let data_type = DataType::Time32(*resolution);
2753        let size = ByteCount::from(data_type.primitive_width().unwrap() as u64);
2754        let dist = Uniform::new(start, end).unwrap();
2755        let sample_fn = move |rng: &mut _| dist.sample(rng);
2756
2757        match resolution {
2758            TimeUnit::Second => Box::new(FnGen::<i32, Time32SecondArray, _>::new_known_size(
2759                data_type, sample_fn, 1, size,
2760            )),
2761            TimeUnit::Millisecond => {
2762                Box::new(FnGen::<i32, Time32MillisecondArray, _>::new_known_size(
2763                    data_type, sample_fn, 1, size,
2764                ))
2765            }
2766            _ => panic!(),
2767        }
2768    }
2769
2770    /// Create a generator of randomly sampled time64 values covering the entire
2771    /// range of 1 day
2772    pub fn rand_time64(resolution: &TimeUnit) -> Box<dyn ArrayGenerator> {
2773        let start = 0_i64;
2774        let end: i64 = match resolution {
2775            TimeUnit::Microsecond => 86_400_000,
2776            TimeUnit::Nanosecond => 86_400_000_000,
2777            _ => panic!(),
2778        };
2779
2780        let data_type = DataType::Time64(*resolution);
2781        let size = ByteCount::from(data_type.primitive_width().unwrap() as u64);
2782        let dist = Uniform::new(start, end).unwrap();
2783        let sample_fn = move |rng: &mut _| dist.sample(rng);
2784
2785        match resolution {
2786            TimeUnit::Microsecond => {
2787                Box::new(FnGen::<i64, Time64MicrosecondArray, _>::new_known_size(
2788                    data_type, sample_fn, 1, size,
2789                ))
2790            }
2791            TimeUnit::Nanosecond => {
2792                Box::new(FnGen::<i64, Time64NanosecondArray, _>::new_known_size(
2793                    data_type, sample_fn, 1, size,
2794                ))
2795            }
2796            _ => panic!(),
2797        }
2798    }
2799
2800    /// Create a generator of random UUIDs, stored as fixed size binary values
2801    ///
2802    /// Note, these are "pseudo UUIDs".  They are 16-byte randomish values but they
2803    /// are not guaranteed to be unique.  We use a simplistic RNG that trades uniqueness
2804    /// for speed.
2805    pub fn rand_pseudo_uuid() -> Box<dyn ArrayGenerator> {
2806        Box::<PseudoUuidGenerator>::default()
2807    }
2808
2809    /// Create a generator of random UUIDs, stored as 32-character strings (hex encoding
2810    /// of the 16-byte binary value)
2811    ///
2812    /// Note, these are "pseudo UUIDs".  They are 16-byte randomish values but they
2813    /// are not guaranteed to be unique.  We use a simplistic RNG that trades uniqueness
2814    /// for speed.
2815    pub fn rand_pseudo_uuid_hex() -> Box<dyn ArrayGenerator> {
2816        Box::<PseudoUuidHexGenerator>::default()
2817    }
2818
2819    pub fn rand_primitive<T: ArrowPrimitiveType + Send + Sync>(
2820        data_type: DataType,
2821    ) -> Box<dyn ArrayGenerator> {
2822        Box::new(RandomBytesGenerator::<T>::new(data_type))
2823    }
2824
2825    pub fn rand_fsb(size: i32) -> Box<dyn ArrayGenerator> {
2826        Box::new(RandomFixedSizeBinaryGenerator::new(size))
2827    }
2828
2829    pub fn rand_interval(unit: IntervalUnit) -> Box<dyn ArrayGenerator> {
2830        Box::new(RandomIntervalGenerator::new(unit))
2831    }
2832
2833    /// The default sampling range for temporal generators: the 365 days ending at
2834    /// 2024-01-01T00:00:00Z (exclusive)
2835    ///
2836    /// The range must be a fixed anchor and not derived from the wall clock
2837    /// (e.g. `Utc::now()`), otherwise the same RNG seed would generate different
2838    /// values depending on when the generator was created, breaking
2839    /// reproducibility (e.g. of saved fuzz inputs).  Callers that need a
2840    /// time-relative range can use the `*_in_range` variants.
2841    fn default_temporal_range() -> (chrono::DateTime<Utc>, chrono::DateTime<Utc>) {
2842        let end = chrono::DateTime::<Utc>::from_timestamp(1_704_067_200, 0)
2843            .expect("2024-01-01T00:00:00Z is a valid timestamp");
2844        let start = end - chrono::TimeDelta::try_days(365).expect("TimeDelta try_days");
2845        (start, end)
2846    }
2847
2848    /// Create a generator of randomly sampled date32 values
2849    ///
2850    /// Instead of sampling the entire range, all values will be drawn from a fixed
2851    /// one-year range (the 365 days ending at 2024-01-01 UTC) as this is a more
2852    /// common use pattern.  Use [`rand_date32_in_range`] to control the range.
2853    pub fn rand_date32() -> Box<dyn ArrayGenerator> {
2854        let (start, end) = default_temporal_range();
2855        rand_date32_in_range(start, end)
2856    }
2857
2858    /// Create a generator of randomly sampled date32 values in the given range
2859    pub fn rand_date32_in_range(
2860        start: chrono::DateTime<Utc>,
2861        end: chrono::DateTime<Utc>,
2862    ) -> Box<dyn ArrayGenerator> {
2863        let data_type = DataType::Date32;
2864        let end_ms = end.timestamp_millis();
2865        let end_days = (end_ms / MS_PER_DAY) as i32;
2866        let start_ms = start.timestamp_millis();
2867        let start_days = (start_ms / MS_PER_DAY) as i32;
2868        let dist = Uniform::new(start_days, end_days).unwrap();
2869
2870        Box::new(FnGen::<i32, Date32Array, _>::new_known_size(
2871            data_type,
2872            move |rng| dist.sample(rng),
2873            1,
2874            DataType::Date32
2875                .primitive_width()
2876                .map(|width| ByteCount::from(width as u64))
2877                .expect("Date32 should have a fixed width"),
2878        ))
2879    }
2880
2881    /// Create a generator of randomly sampled date64 values
2882    ///
2883    /// Instead of sampling the entire range, all values will be drawn from a fixed
2884    /// one-year range (the 365 days ending at 2024-01-01 UTC) as this is a more
2885    /// common use pattern.  Use [`rand_date64_in_range`] to control the range.
2886    pub fn rand_date64() -> Box<dyn ArrayGenerator> {
2887        let (start, end) = default_temporal_range();
2888        rand_date64_in_range(start, end)
2889    }
2890
2891    /// Create a generator of randomly sampled timestamp values in the given range
2892    ///
2893    /// Currently just samples the entire range of u64 values and casts to timestamp
2894    pub fn rand_timestamp_in_range(
2895        start: chrono::DateTime<Utc>,
2896        end: chrono::DateTime<Utc>,
2897        data_type: &DataType,
2898    ) -> Box<dyn ArrayGenerator> {
2899        let end_ms = end.timestamp_millis();
2900        let start_ms = start.timestamp_millis();
2901        let (start_ticks, end_ticks) = match data_type {
2902            DataType::Timestamp(TimeUnit::Nanosecond, _) => {
2903                (start_ms * 1000 * 1000, end_ms * 1000 * 1000)
2904            }
2905            DataType::Timestamp(TimeUnit::Microsecond, _) => (start_ms * 1000, end_ms * 1000),
2906            DataType::Timestamp(TimeUnit::Millisecond, _) => (start_ms, end_ms),
2907            DataType::Timestamp(TimeUnit::Second, _) => (start.timestamp(), end.timestamp()),
2908            _ => panic!(),
2909        };
2910        let dist = Uniform::new(start_ticks, end_ticks).unwrap();
2911
2912        let data_type = data_type.clone();
2913        let sample_fn = move |rng: &mut _| dist.sample(rng);
2914        let width = data_type
2915            .primitive_width()
2916            .map(|width| ByteCount::from(width as u64))
2917            .unwrap();
2918
2919        match data_type {
2920            DataType::Timestamp(TimeUnit::Nanosecond, _) => {
2921                Box::new(FnGen::<i64, TimestampNanosecondArray, _>::new_known_size(
2922                    data_type, sample_fn, 1, width,
2923                ))
2924            }
2925            DataType::Timestamp(TimeUnit::Microsecond, _) => {
2926                Box::new(FnGen::<i64, TimestampMicrosecondArray, _>::new_known_size(
2927                    data_type, sample_fn, 1, width,
2928                ))
2929            }
2930            DataType::Timestamp(TimeUnit::Millisecond, _) => {
2931                Box::new(FnGen::<i64, TimestampMillisecondArray, _>::new_known_size(
2932                    data_type, sample_fn, 1, width,
2933                ))
2934            }
2935            DataType::Timestamp(TimeUnit::Second, _) => {
2936                Box::new(FnGen::<i64, TimestampSecondArray, _>::new_known_size(
2937                    data_type, sample_fn, 1, width,
2938                ))
2939            }
2940            _ => panic!(),
2941        }
2942    }
2943
2944    /// Create a generator of randomly sampled timestamp values
2945    ///
2946    /// Instead of sampling the entire range, all values will be drawn from a fixed
2947    /// one-year range (the 365 days ending at 2024-01-01 UTC) as this is a more
2948    /// common use pattern.  Use [`rand_timestamp_in_range`] to control the range.
2949    pub fn rand_timestamp(data_type: &DataType) -> Box<dyn ArrayGenerator> {
2950        let (start, end) = default_temporal_range();
2951        rand_timestamp_in_range(start, end, data_type)
2952    }
2953
2954    /// Create a generator of randomly sampled date64 values
2955    ///
2956    /// Instead of sampling the entire range, all values will be drawn from the last year as this
2957    /// is a more common use pattern
2958    pub fn rand_date64_in_range(
2959        start: chrono::DateTime<Utc>,
2960        end: chrono::DateTime<Utc>,
2961    ) -> Box<dyn ArrayGenerator> {
2962        let data_type = DataType::Date64;
2963        let end_ms = end.timestamp_millis();
2964        let end_days = end_ms / MS_PER_DAY;
2965        let start_ms = start.timestamp_millis();
2966        let start_days = start_ms / MS_PER_DAY;
2967        let dist = Uniform::new(start_days, end_days).unwrap();
2968
2969        Box::new(FnGen::<i64, Date64Array, _>::new_known_size(
2970            data_type,
2971            move |rng| (dist.sample(rng)) * MS_PER_DAY,
2972            1,
2973            DataType::Date64
2974                .primitive_width()
2975                .map(|width| ByteCount::from(width as u64))
2976                .expect("Date64 should have a fixed width"),
2977        ))
2978    }
2979
2980    /// Create a generator of random binary values where each value has a fixed number of bytes
2981    pub fn rand_fixedbin(bytes_per_element: ByteCount, is_large: bool) -> Box<dyn ArrayGenerator> {
2982        Box::new(RandomBinaryGenerator::new(
2983            bytes_per_element,
2984            false,
2985            is_large,
2986        ))
2987    }
2988
2989    /// Create a generator of random binary values where each value has a variable number of bytes
2990    ///
2991    /// The number of bytes per element will be randomly sampled from the given (inclusive) range
2992    pub fn rand_varbin(
2993        min_bytes_per_element: ByteCount,
2994        max_bytes_per_element: ByteCount,
2995    ) -> Box<dyn ArrayGenerator> {
2996        Box::new(VariableRandomBinaryGenerator::new(
2997            min_bytes_per_element,
2998            max_bytes_per_element,
2999        ))
3000    }
3001
3002    /// Create a generator of random strings
3003    ///
3004    /// All strings will consist entirely of printable ASCII characters
3005    pub fn rand_utf8(bytes_per_element: ByteCount, is_large: bool) -> Box<dyn ArrayGenerator> {
3006        Box::new(RandomBinaryGenerator::new(
3007            bytes_per_element,
3008            true,
3009            is_large,
3010        ))
3011    }
3012
3013    /// Creates a generator of strings with a prefix and a counter
3014    ///
3015    /// For example, if the prefix is "user_" the strings will be "user_0", "user_1", ...
3016    pub fn utf8_prefix_plus_counter(
3017        prefix: impl Into<String>,
3018        is_large: bool,
3019    ) -> Box<dyn ArrayGenerator> {
3020        Box::new(PrefixPlusCounterGenerator::new(prefix.into(), is_large))
3021    }
3022
3023    pub fn binary_prefix_plus_counter(
3024        prefix: Arc<[u8]>,
3025        is_large: bool,
3026    ) -> Box<dyn ArrayGenerator> {
3027        Box::new(BinaryPrefixPlusCounterGenerator::new(prefix, is_large))
3028    }
3029
3030    /// Create a random generator of boolean values
3031    pub fn rand_boolean() -> Box<dyn ArrayGenerator> {
3032        Box::<RandomBooleanGenerator>::default()
3033    }
3034
3035    /// Create a generator of random sentences
3036    ///
3037    /// Generates strings containing between min_words and max_words random English words joined by spaces
3038    pub fn random_sentence(
3039        min_words: usize,
3040        max_words: usize,
3041        is_large: bool,
3042    ) -> Box<dyn ArrayGenerator> {
3043        Box::new(RandomSentenceGenerator::new(min_words, max_words, is_large))
3044    }
3045
3046    /// Create a generator of random words (one word per row)
3047    ///
3048    /// Generates strings containing a single random English word per row
3049    pub fn random_word(is_large: bool) -> Box<dyn ArrayGenerator> {
3050        Box::new(RandomWordGenerator::new(is_large))
3051    }
3052
3053    pub fn rand_list(item_type: &DataType, is_large: bool) -> Box<dyn ArrayGenerator> {
3054        let child_gen = rand_type(item_type);
3055        Box::new(RandomListGenerator::new(child_gen, is_large))
3056    }
3057
3058    pub fn rand_list_any(
3059        item_gen: Box<dyn ArrayGenerator>,
3060        is_large: bool,
3061    ) -> Box<dyn ArrayGenerator> {
3062        Box::new(RandomListGenerator::new(item_gen, is_large))
3063    }
3064
3065    /// Generates random map arrays where each map has 0-4 entries.
3066    pub fn rand_map(key_type: &DataType, value_type: &DataType) -> Box<dyn ArrayGenerator> {
3067        let keys_gen = rand_type(key_type);
3068        let values_gen = rand_type(value_type);
3069        Box::new(RandomMapGenerator::new(keys_gen, values_gen))
3070    }
3071
3072    pub fn rand_struct(fields: Fields) -> Box<dyn ArrayGenerator> {
3073        let child_gens = fields
3074            .iter()
3075            .map(|f| rand_type(f.data_type()))
3076            .collect::<Vec<_>>();
3077        Box::new(RandomStructGenerator::new(fields, child_gens))
3078    }
3079
3080    pub fn null_type() -> Box<dyn ArrayGenerator> {
3081        Box::new(NullArrayGenerator {})
3082    }
3083
3084    /// Create a generator of random values
3085    pub fn rand_type(data_type: &DataType) -> Box<dyn ArrayGenerator> {
3086        match data_type {
3087            DataType::Boolean => rand_boolean(),
3088            DataType::Int8 => rand::<Int8Type>(),
3089            DataType::Int16 => rand::<Int16Type>(),
3090            DataType::Int32 => rand::<Int32Type>(),
3091            DataType::Int64 => rand::<Int64Type>(),
3092            DataType::UInt8 => rand::<UInt8Type>(),
3093            DataType::UInt16 => rand::<UInt16Type>(),
3094            DataType::UInt32 => rand::<UInt32Type>(),
3095            DataType::UInt64 => rand::<UInt64Type>(),
3096            DataType::Float16 => rand_primitive::<Float16Type>(data_type.clone()),
3097            DataType::Float32 => rand::<Float32Type>(),
3098            DataType::Float64 => rand::<Float64Type>(),
3099            DataType::Decimal128(_, _) => rand_primitive::<Decimal128Type>(data_type.clone()),
3100            DataType::Decimal256(_, _) => rand_primitive::<Decimal256Type>(data_type.clone()),
3101            DataType::Utf8 => rand_utf8(ByteCount::from(12), false),
3102            DataType::LargeUtf8 => rand_utf8(ByteCount::from(12), true),
3103            DataType::Binary => rand_fixedbin(ByteCount::from(12), false),
3104            DataType::LargeBinary => rand_fixedbin(ByteCount::from(12), true),
3105            DataType::Dictionary(key_type, value_type) => {
3106                dict_type(rand_type(value_type), key_type)
3107            }
3108            DataType::FixedSizeList(child, dimension) => cycle_vec(
3109                rand_type(child.data_type()),
3110                Dimension::from(*dimension as u32),
3111            ),
3112            DataType::FixedSizeBinary(size) => rand_fsb(*size),
3113            DataType::List(child) => rand_list(child.data_type(), false),
3114            DataType::LargeList(child) => rand_list(child.data_type(), true),
3115            DataType::Map(entries_field, _) => {
3116                let DataType::Struct(fields) = entries_field.data_type() else {
3117                    panic!("Map entries field must be a struct");
3118                };
3119                let key_type = fields[0].data_type();
3120                let value_type = fields[1].data_type();
3121                rand_map(key_type, value_type)
3122            }
3123            DataType::Duration(unit) => match unit {
3124                TimeUnit::Second => rand::<DurationSecondType>(),
3125                TimeUnit::Millisecond => rand::<DurationMillisecondType>(),
3126                TimeUnit::Microsecond => rand::<DurationMicrosecondType>(),
3127                TimeUnit::Nanosecond => rand::<DurationNanosecondType>(),
3128            },
3129            DataType::Interval(unit) => rand_interval(*unit),
3130            DataType::Date32 => rand_date32(),
3131            DataType::Date64 => rand_date64(),
3132            DataType::Time32(resolution) => rand_time32(resolution),
3133            DataType::Time64(resolution) => rand_time64(resolution),
3134            DataType::Timestamp(_, _) => rand_timestamp(data_type),
3135            DataType::Struct(fields) => rand_struct(fields.clone()),
3136            DataType::Null => null_type(),
3137            _ => unimplemented!("random generation of {}", data_type),
3138        }
3139    }
3140
3141    /// Encodes arrays generated by the underlying generator as dictionaries with the given key type
3142    ///
3143    /// Note that this may not be very realistic if the underlying generator is something like a random
3144    /// generator since most of the underlying values will be unique and the common case for dictionary
3145    /// encoding is when there is a small set of possible values.
3146    pub fn dict<K: ArrowDictionaryKeyType + Send + Sync>(
3147        generator: Box<dyn ArrayGenerator>,
3148    ) -> Box<dyn ArrayGenerator> {
3149        Box::new(DictionaryGenerator::<K>::new(generator))
3150    }
3151
3152    /// Encodes arrays generated by the underlying generator as dictionaries with the given key type
3153    pub fn dict_type(
3154        generator: Box<dyn ArrayGenerator>,
3155        key_type: &DataType,
3156    ) -> Box<dyn ArrayGenerator> {
3157        match key_type {
3158            DataType::Int8 => dict::<Int8Type>(generator),
3159            DataType::Int16 => dict::<Int16Type>(generator),
3160            DataType::Int32 => dict::<Int32Type>(generator),
3161            DataType::Int64 => dict::<Int64Type>(generator),
3162            DataType::UInt8 => dict::<UInt8Type>(generator),
3163            DataType::UInt16 => dict::<UInt16Type>(generator),
3164            DataType::UInt32 => dict::<UInt32Type>(generator),
3165            DataType::UInt64 => dict::<UInt64Type>(generator),
3166            _ => unimplemented!(),
3167        }
3168    }
3169
3170    /// Wraps a generator to produce low-cardinality data.
3171    ///
3172    /// Generates `cardinality` unique values on first call, then randomly
3173    /// selects from them for all subsequent rows.
3174    pub fn low_cardinality(
3175        generator: Box<dyn ArrayGenerator>,
3176        cardinality: usize,
3177    ) -> Box<dyn ArrayGenerator> {
3178        Box::new(LowCardinalityGenerator::new(generator, cardinality))
3179    }
3180}
3181
3182/// Create a BatchGeneratorBuilder to start generating batch data
3183pub fn gen_batch() -> BatchGeneratorBuilder {
3184    BatchGeneratorBuilder::default()
3185}
3186
3187/// Create an ArrayGeneratorBuilder to start generating array data
3188pub fn gen_array(genn: Box<dyn ArrayGenerator>) -> ArrayGeneratorBuilder {
3189    ArrayGeneratorBuilder::new(genn)
3190}
3191
3192/// Metadata key to specify content type for string generation.
3193/// Set to "sentence" to use the sentence generator with Zipf distribution.
3194pub const CONTENT_TYPE_KEY: &str = "lance-datagen:content-type";
3195
3196/// Metadata key to specify cardinality for low-cardinality data generation.
3197/// Set to a numeric string (e.g., "100") to limit unique values.
3198pub const CARDINALITY_KEY: &str = "lance-datagen:cardinality";
3199
3200/// Create a generator for a field, checking metadata for content type hints.
3201///
3202/// Supported metadata keys:
3203/// - `lance-datagen:content-type`: Set to "sentence" for Utf8/LargeUtf8 fields
3204///   to use the sentence generator with Zipf distribution.
3205/// - `lance-datagen:cardinality`: Set to a number to limit unique values.
3206///   The generator will produce only that many unique values and randomly
3207///   select from them.
3208pub fn rand_field(field: &Field) -> Box<dyn ArrayGenerator> {
3209    let mut generator = if let Some(content_type) = field.metadata().get(CONTENT_TYPE_KEY) {
3210        match (content_type.as_str(), field.data_type()) {
3211            ("sentence", DataType::Utf8) => array::random_sentence(1, 10, false),
3212            ("sentence", DataType::LargeUtf8) => array::random_sentence(1, 10, true),
3213            _ => array::rand_type(field.data_type()),
3214        }
3215    } else {
3216        array::rand_type(field.data_type())
3217    };
3218
3219    if let Some(cardinality_str) = field.metadata().get(CARDINALITY_KEY)
3220        && let Ok(cardinality) = cardinality_str.parse::<usize>()
3221        && cardinality > 0
3222    {
3223        generator = array::low_cardinality(generator, cardinality);
3224    }
3225
3226    generator
3227}
3228
3229/// Create a BatchGeneratorBuilder with the given schema
3230///
3231/// You can add more columns or convert this into a reader immediately.
3232///
3233/// Supported field metadata:
3234/// - `lance-datagen:content-type` = `"sentence"`: Use sentence generator with
3235///   Zipf distribution for more realistic text (Utf8/LargeUtf8 only).
3236/// - `lance-datagen:cardinality` = `"<number>"`: Limit to N unique values.
3237pub fn rand(schema: &Schema) -> BatchGeneratorBuilder {
3238    let mut builder = BatchGeneratorBuilder::default();
3239    for field in schema.fields() {
3240        builder = builder.col(field.name(), rand_field(field));
3241    }
3242    builder
3243}
3244
3245#[cfg(test)]
3246mod tests {
3247
3248    use arrow::datatypes::{Float32Type, Int8Type, Int16Type, TimeUnit, UInt32Type};
3249    use arrow_array::{
3250        BooleanArray, Date32Array, Date64Array, Float32Array, Int8Array, Int16Array, Int32Array,
3251        TimestampMicrosecondArray, TimestampMillisecondArray, TimestampNanosecondArray,
3252        TimestampSecondArray, UInt32Array,
3253    };
3254
3255    use super::*;
3256
3257    #[test]
3258    fn test_timestamp_timezone_is_preserved() {
3259        let data_type = DataType::Timestamp(TimeUnit::Millisecond, Some("UTC".into()));
3260        let mut generator = array::rand_type(&data_type);
3261        let generated = generator.generate_default(RowCount::from(2)).unwrap();
3262        assert_eq!(generated.data_type(), &data_type);
3263
3264        let fields = Fields::from(vec![Field::new("timestamp", data_type, true)]);
3265        let mut generator = array::rand_struct(fields.clone());
3266        let generated = generator.generate_default(RowCount::from(2)).unwrap();
3267        assert_eq!(generated.data_type(), &DataType::Struct(fields));
3268    }
3269
3270    #[test]
3271    fn test_fn_gen_propagates_array_data_build_error() {
3272        // FnGen constructors are internal. Use an incompatible declared type to
3273        // verify that ArrayDataBuilder validation failures are propagated.
3274        let mut generator = FnGen::<i32, Int32Array, _>::new_unknown_size(DataType::Utf8, |_| 0, 1);
3275
3276        assert!(matches!(
3277            generator.generate_default(RowCount::from(1)),
3278            Err(ArrowError::InvalidArgumentError(_))
3279        ));
3280    }
3281
3282    #[test]
3283    fn test_step() {
3284        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3285        let mut genn = array::step::<Int32Type>();
3286        assert_eq!(
3287            *genn.generate(RowCount::from(5), &mut rng).unwrap(),
3288            Int32Array::from_iter([0, 1, 2, 3, 4])
3289        );
3290        assert_eq!(
3291            *genn.generate(RowCount::from(5), &mut rng).unwrap(),
3292            Int32Array::from_iter([5, 6, 7, 8, 9])
3293        );
3294
3295        let mut genn = array::step::<Int8Type>();
3296        assert_eq!(
3297            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3298            Int8Array::from_iter([0, 1, 2])
3299        );
3300
3301        let mut genn = array::step::<Float32Type>();
3302        assert_eq!(
3303            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3304            Float32Array::from_iter([0.0, 1.0, 2.0])
3305        );
3306
3307        let mut genn = array::step_custom::<Int16Type>(4, 8);
3308        assert_eq!(
3309            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3310            Int16Array::from_iter([4, 12, 20])
3311        );
3312        assert_eq!(
3313            *genn.generate(RowCount::from(2), &mut rng).unwrap(),
3314            Int16Array::from_iter([28, 36])
3315        );
3316    }
3317
3318    #[test]
3319    fn test_cycle() {
3320        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3321        let mut genn = array::cycle::<Int32Type>(vec![1, 2, 3]);
3322        assert_eq!(
3323            *genn.generate(RowCount::from(5), &mut rng).unwrap(),
3324            Int32Array::from_iter([1, 2, 3, 1, 2])
3325        );
3326
3327        let mut genn = array::cycle_utf8_literals(&["abc", "def", "xyz"]);
3328        assert_eq!(
3329            *genn.generate(RowCount::from(5), &mut rng).unwrap(),
3330            StringArray::from_iter_values(["abc", "def", "xyz", "abc", "def"])
3331        );
3332        assert_eq!(
3333            *genn.generate(RowCount::from(1), &mut rng).unwrap(),
3334            StringArray::from_iter_values(["xyz"])
3335        );
3336
3337        let mut genn = array::cycle_bool(vec![false, false, true]);
3338        assert_eq!(
3339            *genn.generate(RowCount::from(5), &mut rng).unwrap(),
3340            BooleanArray::from_iter(vec![false, false, true, false, false].into_iter().map(Some))
3341        );
3342        assert_eq!(
3343            *genn.generate(RowCount::from(1), &mut rng).unwrap(),
3344            BooleanArray::from_iter(vec![Some(true)])
3345        )
3346    }
3347
3348    #[test]
3349    fn test_fill() {
3350        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3351        let mut genn = array::fill::<Int32Type>(42);
3352        assert_eq!(
3353            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3354            Int32Array::from_iter([42, 42, 42])
3355        );
3356        assert_eq!(
3357            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3358            Int32Array::from_iter([42, 42, 42])
3359        );
3360
3361        let mut genn = array::fill_varbin(vec![0, 1, 2]);
3362        assert_eq!(
3363            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3364            arrow_array::BinaryArray::from_iter_values([
3365                "\x00\x01\x02",
3366                "\x00\x01\x02",
3367                "\x00\x01\x02"
3368            ])
3369        );
3370
3371        let mut genn = array::fill_utf8("xyz".to_string());
3372        assert_eq!(
3373            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3374            arrow_array::StringArray::from_iter_values(["xyz", "xyz", "xyz"])
3375        );
3376    }
3377
3378    #[test]
3379    fn test_utf8_prefix_plus_counter() {
3380        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3381        let mut genn = array::utf8_prefix_plus_counter("user_", false);
3382        assert_eq!(
3383            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3384            arrow_array::StringArray::from_iter_values(["user_0", "user_1", "user_2"])
3385        );
3386
3387        let mut genn = array::utf8_prefix_plus_counter("user_", true);
3388        assert_eq!(
3389            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3390            arrow_array::LargeStringArray::from_iter_values(["user_0", "user_1", "user_2"])
3391        );
3392    }
3393
3394    #[test]
3395    fn test_rng() {
3396        // Note: these tests are heavily dependent on the default seed.
3397        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3398        let mut genn = array::rand::<Int32Type>();
3399        assert_eq!(
3400            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3401            Int32Array::from_iter([-797553329, 1369325940, -69174021])
3402        );
3403
3404        let mut genn = array::rand_fixedbin(ByteCount::from(3), false);
3405        assert_eq!(
3406            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3407            arrow_array::BinaryArray::from_iter_values([
3408                [184, 53, 216],
3409                [12, 96, 159],
3410                [125, 179, 56]
3411            ])
3412        );
3413
3414        let mut genn = array::rand_utf8(ByteCount::from(3), false);
3415        assert_eq!(
3416            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3417            arrow_array::StringArray::from_iter_values([">@p", "n `", "NWa"])
3418        );
3419
3420        let mut genn = array::random_sentence(1, 5, false);
3421        let words = genn.generate(RowCount::from(10), &mut rng).unwrap();
3422        assert_eq!(words.data_type(), &DataType::Utf8);
3423        let words_array = words.as_any().downcast_ref::<StringArray>().unwrap();
3424        // Verify each string contains 1-5 words
3425        for i in 0..10 {
3426            let sentence = words_array.value(i);
3427            let word_count = sentence.split_whitespace().count();
3428            assert!((1..=5).contains(&word_count));
3429        }
3430
3431        let mut genn = array::rand_date32();
3432        let days_32 = genn.generate(RowCount::from(3), &mut rng).unwrap();
3433        assert_eq!(days_32.data_type(), &DataType::Date32);
3434
3435        let mut genn = array::rand_date64();
3436        let days_64 = genn.generate(RowCount::from(3), &mut rng).unwrap();
3437        assert_eq!(days_64.data_type(), &DataType::Date64);
3438
3439        let mut genn = array::rand_boolean();
3440        let bools = genn.generate(RowCount::from(1024), &mut rng).unwrap();
3441        assert_eq!(bools.data_type(), &DataType::Boolean);
3442        let bools = bools.as_any().downcast_ref::<BooleanArray>().unwrap();
3443        // Sanity check to ensure we're getting at least some rng
3444        assert!(bools.false_count() > 100);
3445        assert!(bools.true_count() > 100);
3446
3447        let mut genn = array::rand_varbin(ByteCount::from(2), ByteCount::from(4));
3448        assert_eq!(
3449            *genn.generate(RowCount::from(3), &mut rng).unwrap(),
3450            arrow_array::BinaryArray::from_iter_values([
3451                vec![111, 9, 80],
3452                vec![86, 118, 13, 209],
3453                vec![68, 33, 202]
3454            ])
3455        );
3456    }
3457
3458    #[test]
3459    fn test_rng_temporal_deterministic() {
3460        // The default temporal generators must not depend on the wall clock: the
3461        // same seed must produce the same values no matter when the generator is
3462        // created (https://github.com/lance-format/lance/issues/7913).  These
3463        // exact values pin both the RNG stream and the fixed default sampling
3464        // range (the 365 days ending at 2024-01-01 UTC).
3465        fn gen_values(mut genn: Box<dyn ArrayGenerator>) -> Arc<dyn Array> {
3466            let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3467            genn.generate(RowCount::from(3), &mut rng).unwrap()
3468        }
3469
3470        assert_eq!(
3471            *gen_values(array::rand_date32()),
3472            Date32Array::from(vec![19655, 19474, 19717])
3473        );
3474        assert_eq!(
3475            *gen_values(array::rand_date64()),
3476            Date64Array::from(vec![
3477                1_698_192_000_000,
3478                1_682_553_600_000,
3479                1_703_548_800_000
3480            ])
3481        );
3482        assert_eq!(
3483            *gen_values(array::rand_timestamp(&DataType::Timestamp(
3484                TimeUnit::Second,
3485                None
3486            ))),
3487            TimestampSecondArray::from(vec![1_698_211_127, 1_682_585_540, 1_703_559_286])
3488        );
3489        assert_eq!(
3490            *gen_values(array::rand_timestamp(&DataType::Timestamp(
3491                TimeUnit::Millisecond,
3492                None
3493            ))),
3494            TimestampMillisecondArray::from(vec![
3495                1_698_211_127_056,
3496                1_682_585_540_319,
3497                1_703_559_286_487
3498            ])
3499        );
3500        assert_eq!(
3501            *gen_values(array::rand_timestamp(&DataType::Timestamp(
3502                TimeUnit::Microsecond,
3503                None
3504            ))),
3505            TimestampMicrosecondArray::from(vec![
3506                1_698_211_127_056_596,
3507                1_682_585_540_319_384,
3508                1_703_559_286_487_645
3509            ])
3510        );
3511        assert_eq!(
3512            *gen_values(array::rand_timestamp(&DataType::Timestamp(
3513                TimeUnit::Nanosecond,
3514                None
3515            ))),
3516            TimestampNanosecondArray::from(vec![
3517                1_698_211_127_056_596_085,
3518                1_682_585_540_319_384_548,
3519                1_703_559_286_487_645_287
3520            ])
3521        );
3522    }
3523
3524    #[test]
3525    fn test_rng_list() {
3526        // Note: these tests are heavily dependent on the default seed.
3527        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3528        let mut genn = array::rand_list(&DataType::Int32, false);
3529        let arr = genn.generate(RowCount::from(100), &mut rng).unwrap();
3530        // Make sure we can generate empty lists (note, test is dependent on seed)
3531        let arr = arr.as_list::<i32>();
3532        assert!(arr.iter().any(|l| l.unwrap().is_empty()));
3533        // Shouldn't generate any giant lists (don't kill performance in normal datagen)
3534        assert!(arr.iter().any(|l| l.unwrap().len() < 11));
3535    }
3536
3537    #[test]
3538    fn test_rng_distribution() {
3539        // Sanity test to make sure we our RNG is giving us well distributed values
3540        // We generates some 4-byte integers, histogram them into 8 buckets, and make
3541        // sure each bucket has a good # of values
3542        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3543        let mut genn = array::rand::<UInt32Type>();
3544        for _ in 0..10 {
3545            let arr = genn.generate(RowCount::from(10000), &mut rng).unwrap();
3546            let int_arr = arr.as_any().downcast_ref::<UInt32Array>().unwrap();
3547            let mut buckets = vec![0_u32; 256];
3548            for val in int_arr.values() {
3549                buckets[(*val >> 24) as usize] += 1;
3550            }
3551            for bucket in buckets {
3552                // Perfectly even distribution would have 10000 / 256 values (~40) per bucket
3553                // We test for 15 which should be "good enough" and statistically unlikely to fail
3554                assert!(bucket > 15);
3555            }
3556        }
3557    }
3558
3559    #[test]
3560    fn test_nulls() {
3561        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3562        let mut genn = array::rand::<Int32Type>().with_random_nulls(0.3);
3563
3564        let arr = genn.generate(RowCount::from(1000), &mut rng).unwrap();
3565
3566        // This assert depends on the default seed
3567        assert_eq!(arr.null_count(), 297);
3568
3569        for len in 0..100 {
3570            let arr = genn.generate(RowCount::from(len), &mut rng).unwrap();
3571            // Make sure the null count we came up with matches the actual # of unset bits
3572            assert_eq!(
3573                arr.null_count(),
3574                arr.nulls()
3575                    .map(|nulls| (len as usize)
3576                        - nulls.buffer().count_set_bits_offset(0, len as usize))
3577                    .unwrap_or(0)
3578            );
3579        }
3580
3581        let mut genn = array::rand::<Int32Type>().with_random_nulls(0.0);
3582        let arr = genn.generate(RowCount::from(10), &mut rng).unwrap();
3583
3584        assert_eq!(arr.null_count(), 0);
3585
3586        let mut genn = array::rand::<Int32Type>().with_random_nulls(1.0);
3587        let arr = genn.generate(RowCount::from(10), &mut rng).unwrap();
3588
3589        assert_eq!(arr.null_count(), 10);
3590        assert!((0..10).all(|idx| arr.is_null(idx)));
3591
3592        let mut genn = array::rand::<Int32Type>().with_nulls(&[false, false, true]);
3593        let arr = genn.generate(RowCount::from(7), &mut rng).unwrap();
3594        assert!((0..2).all(|idx| arr.is_valid(idx)));
3595        assert!(arr.is_null(2));
3596        assert!((3..5).all(|idx| arr.is_valid(idx)));
3597        assert!(arr.is_null(5));
3598        assert!(arr.is_valid(6));
3599    }
3600
3601    #[test]
3602    fn test_unit_circle() {
3603        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3604        let mut genn = array::cycle_unit_circle(4);
3605        let arr = genn.generate(RowCount::from(6), &mut rng).unwrap();
3606
3607        let arr_values = arr
3608            .as_fixed_size_list()
3609            .values()
3610            .as_primitive::<Float32Type>()
3611            .values()
3612            .to_vec();
3613        assert_eq!(arr_values.len(), 12);
3614        let expected_values = [1.0, 0.0, 0.0, 1.0, -1.0, 0.0, 0.0, -1.0, 1.0, 0.0, 0.0, 1.0];
3615        for (actual, expected) in arr_values.iter().zip(expected_values.iter()) {
3616            assert!((actual - expected).abs() < 0.0001);
3617        }
3618    }
3619
3620    #[test]
3621    fn test_jitter_centroids() {
3622        let mut rng = rand_xoshiro::Xoshiro256PlusPlus::seed_from_u64(DEFAULT_SEED.0);
3623        let mut centroids_gen = array::cycle_unit_circle(4);
3624        let centroids = centroids_gen.generate(RowCount::from(4), &mut rng).unwrap();
3625
3626        let centroid_values = centroids
3627            .as_fixed_size_list()
3628            .values()
3629            .as_primitive::<Float32Type>()
3630            .values()
3631            .to_vec();
3632
3633        let mut jitter_jen = array::jitter_centroids(centroids, 0.001);
3634        let jittered = jitter_jen.generate(RowCount::from(100), &mut rng).unwrap();
3635
3636        let values = jittered
3637            .as_fixed_size_list()
3638            .values()
3639            .as_primitive::<Float32Type>()
3640            .values()
3641            .to_vec();
3642
3643        for i in 0..100 {
3644            let centroid = i % 4;
3645            let centroid_x = centroid_values[centroid * 2];
3646            let centroid_y = centroid_values[centroid * 2 + 1];
3647            let value_x = values[i * 2];
3648            let value_y = values[i * 2 + 1];
3649
3650            let l2_dist = ((value_x - centroid_x).powi(2) + (value_y - centroid_y).powi(2)).sqrt();
3651            assert!(l2_dist < 0.001001);
3652            assert!(l2_dist > 0.000999);
3653        }
3654    }
3655
3656    #[test]
3657    fn test_rand_schema() {
3658        let schema = Schema::new(vec![
3659            Field::new("a", DataType::Int32, true),
3660            Field::new("b", DataType::Utf8, true),
3661            Field::new("c", DataType::Float32, true),
3662            Field::new("d", DataType::Int32, true),
3663            Field::new("e", DataType::Int32, true),
3664        ]);
3665        let rbr = rand(&schema)
3666            .into_reader_bytes(
3667                ByteCount::from(1024 * 1024),
3668                BatchCount::from(8),
3669                RoundingBehavior::ExactOrErr,
3670            )
3671            .unwrap();
3672        assert_eq!(*rbr.schema(), schema);
3673
3674        let batches = rbr.map(|val| val.unwrap()).collect::<Vec<_>>();
3675        assert_eq!(batches.len(), 8);
3676
3677        for batch in batches {
3678            assert_eq!(batch.num_rows(), 1024 * 1024 / 32);
3679            assert_eq!(batch.num_columns(), 5);
3680        }
3681    }
3682}