Skip to main content

lance_file/previous/writer/
mod.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4mod statistics;
5
6use std::collections::HashMap;
7use std::marker::PhantomData;
8
9use arrow_array::builder::{ArrayBuilder, PrimitiveBuilder};
10use arrow_array::cast::{as_large_list_array, as_list_array, as_struct_array};
11use arrow_array::types::{Int32Type, Int64Type};
12use arrow_array::{Array, ArrayRef, RecordBatch, StructArray};
13use arrow_buffer::ArrowNativeType;
14use arrow_data::ArrayData;
15use arrow_schema::DataType;
16use async_recursion::async_recursion;
17use async_trait::async_trait;
18use lance_arrow::*;
19use lance_core::datatypes::{Encoding, Field, NullabilityComparison, Schema, SchemaCompareOptions};
20use lance_core::{Error, Result};
21use lance_io::encodings::{
22    Encoder, binary::BinaryEncoder, dictionary::DictionaryEncoder, plain::PlainEncoder,
23};
24use lance_io::object_store::ObjectStore;
25use lance_io::traits::{WriteExt, Writer};
26use object_store::path::Path;
27use tokio::io::AsyncWriteExt;
28
29use crate::format::{MAGIC, MAJOR_VERSION, MINOR_VERSION};
30use crate::previous::format::metadata::{Metadata, StatisticsMetadata};
31use crate::previous::page_table::{PageInfo, PageTable};
32use crate::writer::FileWriteSummary;
33
34/// The file format currently includes a "manifest" where it stores the schema for
35/// self-describing files.  Historically this has been a table format manifest that
36/// is empty except for the schema field.
37///
38/// Since this crate is not aware of the table format we need this to be provided
39/// externally.  You should always use lance_table::io::manifest::ManifestDescribing
40/// for this today.
41#[async_trait]
42pub trait ManifestProvider {
43    /// Store the schema in the file
44    ///
45    /// This should just require writing the schema (or a manifest wrapper) as a proto struct
46    ///
47    /// Note: the dictionaries have already been written by this point and the schema should
48    /// be populated with the dictionary lengths/offsets
49    async fn store_schema(object_writer: &mut dyn Writer, schema: &Schema)
50    -> Result<Option<usize>>;
51}
52
53/// Implementation of ManifestProvider that does not store the schema
54#[cfg(test)]
55pub(crate) struct NotSelfDescribing {}
56
57#[cfg(test)]
58#[async_trait]
59impl ManifestProvider for NotSelfDescribing {
60    async fn store_schema(_: &mut dyn Writer, _: &Schema) -> Result<Option<usize>> {
61        Ok(None)
62    }
63}
64
65/// [FileWriter] writes Arrow [RecordBatch] to one Lance file.
66///
67/// ```ignored
68/// use lance::io::FileWriter;
69/// use futures::stream::Stream;
70///
71/// let mut file_writer = FileWriter::new(object_store, &path, &schema);
72/// while let Ok(batch) = stream.next().await {
73///     file_writer.write(&batch).unwrap();
74/// }
75/// // Need to close file writer to flush buffer and footer.
76/// file_writer.shutdown();
77/// ```
78pub struct FileWriter<M: ManifestProvider + Send + Sync> {
79    pub object_writer: Box<dyn Writer>,
80    schema: Schema,
81    batch_id: i32,
82    page_table: PageTable,
83    metadata: Metadata,
84    stats_collector: Option<statistics::StatisticsCollector>,
85    manifest_provider: PhantomData<M>,
86}
87
88#[derive(Debug, Clone, Default)]
89pub struct FileWriterOptions {
90    /// The field ids to collect statistics for.
91    ///
92    /// If None, will collect for all fields in the schema (that support stats).
93    /// If an empty vector, will not collect any statistics.
94    pub collect_stats_for_fields: Option<Vec<i32>>,
95}
96
97impl<M: ManifestProvider + Send + Sync> FileWriter<M> {
98    pub async fn try_new(
99        object_store: &ObjectStore,
100        path: &Path,
101        schema: Schema,
102        options: &FileWriterOptions,
103    ) -> Result<Self> {
104        let object_writer = object_store.create(path).await?;
105        Self::with_object_writer(object_writer, schema, options)
106    }
107
108    pub fn with_object_writer(
109        object_writer: Box<dyn Writer>,
110        schema: Schema,
111        options: &FileWriterOptions,
112    ) -> Result<Self> {
113        let collect_stats_for_fields = if let Some(stats_fields) = &options.collect_stats_for_fields
114        {
115            stats_fields.clone()
116        } else {
117            schema.field_ids()
118        };
119
120        let stats_collector = if !collect_stats_for_fields.is_empty() {
121            let stats_schema = schema.project_by_ids(&collect_stats_for_fields, true);
122            statistics::StatisticsCollector::try_new(&stats_schema)
123        } else {
124            None
125        };
126
127        Ok(Self {
128            object_writer,
129            schema,
130            batch_id: 0,
131            page_table: PageTable::default(),
132            metadata: Metadata::default(),
133            stats_collector,
134            manifest_provider: PhantomData,
135        })
136    }
137
138    /// Return the schema of the file writer.
139    pub fn schema(&self) -> &Schema {
140        &self.schema
141    }
142
143    fn verify_field_nullability(arr: &ArrayData, field: &Field) -> Result<()> {
144        if !field.nullable && arr.null_count() > 0 {
145            return Err(Error::invalid_input(format!(
146                "The field `{}` contained null values even though the field is marked non-null in the schema",
147                field.name
148            )));
149        }
150
151        for (child_field, child_arr) in field.children.iter().zip(arr.child_data()) {
152            Self::verify_field_nullability(child_arr, child_field)?;
153        }
154
155        Ok(())
156    }
157
158    fn verify_nullability_constraints(&self, batch: &RecordBatch) -> Result<()> {
159        for (col, field) in batch.columns().iter().zip(self.schema.fields.iter()) {
160            Self::verify_field_nullability(&col.to_data(), field)?;
161        }
162        Ok(())
163    }
164
165    /// Write a [RecordBatch] to the open file.
166    /// All RecordBatch will be treated as one RecordBatch on disk
167    ///
168    /// Returns [Err] if the schema does not match with the batch.
169    pub async fn write(&mut self, batches: &[RecordBatch]) -> Result<()> {
170        if batches.is_empty() {
171            return Ok(());
172        }
173
174        for batch in batches {
175            // Compare, ignore metadata and dictionary
176            //   dictionary should have been checked earlier and could be an expensive check
177            let schema = Schema::try_from(batch.schema().as_ref())?;
178            schema.check_compatible(
179                &self.schema,
180                &SchemaCompareOptions {
181                    compare_nullability: NullabilityComparison::Ignore,
182                    ..Default::default()
183                },
184            )?;
185            self.verify_nullability_constraints(batch)?;
186        }
187
188        // If we are collecting stats for this column, collect them.
189        // Statistics need to traverse nested arrays, so it's a separate loop
190        // from writing which is done on top-level arrays.
191        if let Some(stats_collector) = &mut self.stats_collector {
192            for (field, arrays) in fields_in_batches(batches, &self.schema) {
193                if let Some(stats_builder) = stats_collector.get_builder(field.id) {
194                    let stats_row = statistics::collect_statistics(&arrays);
195                    stats_builder.append(stats_row);
196                }
197            }
198        }
199
200        // Copy a list of fields to avoid borrow checker error.
201        let fields = self.schema.fields.clone();
202        for field in fields.iter() {
203            let arrs = batches
204                .iter()
205                .map(|batch| {
206                    batch.column_by_name(&field.name).ok_or_else(|| {
207                        Error::invalid_input(format!(
208                            "FileWriter::write: Field '{}' not found",
209                            field.name
210                        ))
211                    })
212                })
213                .collect::<Result<Vec<_>>>()?;
214
215            Self::write_array(
216                self.object_writer.as_mut(),
217                field,
218                &arrs,
219                self.batch_id,
220                &mut self.page_table,
221            )
222            .await?;
223        }
224        let batch_length = batches.iter().map(|b| b.num_rows() as i32).sum();
225        self.metadata.push_batch_length(batch_length);
226
227        // It's imperative we complete any in-flight requests, since we are
228        // returning control to the caller. If the caller takes a long time to
229        // write the next batch, the in-flight requests will not be polled and
230        // may time out.
231        self.object_writer.flush().await?;
232
233        self.batch_id += 1;
234        Ok(())
235    }
236
237    /// Add schema metadata, as (key, value) pair to the file.
238    pub fn add_metadata(&mut self, key: &str, value: &str) {
239        self.schema
240            .metadata
241            .insert(key.to_string(), value.to_string());
242    }
243
244    pub async fn finish_with_metadata(
245        &mut self,
246        metadata: &HashMap<String, String>,
247    ) -> Result<FileWriteSummary> {
248        self.schema
249            .metadata
250            .extend(metadata.iter().map(|(k, y)| (k.clone(), y.clone())));
251        self.finish().await
252    }
253
254    pub async fn finish(&mut self) -> Result<FileWriteSummary> {
255        self.write_footer().await?;
256        // `shutdown` flushes the footer and reports the authoritative on-disk
257        // byte count, so the size is sourced here instead of via a later
258        // `tell()` that would rely on the cursor surviving shutdown.
259        let write_result = Writer::shutdown(self.object_writer.as_mut()).await?;
260        let num_rows = self
261            .metadata
262            .batch_offsets
263            .last()
264            .cloned()
265            .unwrap_or_default();
266        Ok(FileWriteSummary {
267            num_rows: num_rows as u64,
268            size_bytes: write_result.size as u64,
269        })
270    }
271
272    /// Total records written in this file.
273    pub fn len(&self) -> usize {
274        self.metadata.len()
275    }
276
277    /// Total bytes written so far
278    pub async fn tell(&mut self) -> Result<usize> {
279        self.object_writer.tell().await
280    }
281
282    /// Return the id of the next batch to be written.
283    pub fn next_batch_id(&self) -> i32 {
284        self.batch_id
285    }
286
287    pub fn is_empty(&self) -> bool {
288        self.len() == 0
289    }
290
291    #[async_recursion]
292    async fn write_array(
293        object_writer: &mut dyn Writer,
294        field: &Field,
295        arrs: &[&ArrayRef],
296        batch_id: i32,
297        page_table: &mut PageTable,
298    ) -> Result<()> {
299        assert!(!arrs.is_empty());
300        let data_type = arrs[0].data_type();
301        let arrs_ref = arrs.iter().map(|a| a.as_ref()).collect::<Vec<_>>();
302
303        match data_type {
304            DataType::Null => {
305                Self::write_null_array(
306                    object_writer,
307                    field,
308                    arrs_ref.as_slice(),
309                    batch_id,
310                    page_table,
311                )
312                .await
313            }
314            dt if dt.is_fixed_stride() => {
315                Self::write_fixed_stride_array(
316                    object_writer,
317                    field,
318                    arrs_ref.as_slice(),
319                    batch_id,
320                    page_table,
321                )
322                .await
323            }
324            dt if dt.is_binary_like() => {
325                Self::write_binary_array(
326                    object_writer,
327                    field,
328                    arrs_ref.as_slice(),
329                    batch_id,
330                    page_table,
331                )
332                .await
333            }
334            DataType::Dictionary(key_type, _) => {
335                Self::write_dictionary_arr(
336                    object_writer,
337                    field,
338                    arrs_ref.as_slice(),
339                    key_type,
340                    batch_id,
341                    page_table,
342                )
343                .await
344            }
345            dt if dt.is_struct() => {
346                let struct_arrays = arrs.iter().map(|a| as_struct_array(a)).collect::<Vec<_>>();
347                Self::write_struct_array(
348                    object_writer,
349                    field,
350                    struct_arrays.as_slice(),
351                    batch_id,
352                    page_table,
353                )
354                .await
355            }
356            DataType::FixedSizeList(_, _) | DataType::FixedSizeBinary(_) => {
357                Self::write_fixed_stride_array(
358                    object_writer,
359                    field,
360                    arrs_ref.as_slice(),
361                    batch_id,
362                    page_table,
363                )
364                .await
365            }
366            DataType::List(_) => {
367                Self::write_list_array(
368                    object_writer,
369                    field,
370                    arrs_ref.as_slice(),
371                    batch_id,
372                    page_table,
373                )
374                .await
375            }
376            DataType::LargeList(_) => {
377                Self::write_large_list_array(
378                    object_writer,
379                    field,
380                    arrs_ref.as_slice(),
381                    batch_id,
382                    page_table,
383                )
384                .await
385            }
386            _ => Err(Error::schema(format!(
387                "FileWriter::write: unsupported data type: {data_type}"
388            ))),
389        }
390    }
391
392    async fn write_null_array(
393        object_writer: &mut dyn Writer,
394        field: &Field,
395        arrs: &[&dyn Array],
396        batch_id: i32,
397        page_table: &mut PageTable,
398    ) -> Result<()> {
399        let arrs_length: i32 = arrs.iter().map(|a| a.len() as i32).sum();
400        let page_info = PageInfo::new(object_writer.tell().await?, arrs_length as usize);
401        page_table.set(field.id, batch_id, page_info);
402        Ok(())
403    }
404
405    /// Write fixed size array, including, primtiives, fixed size binary, and fixed size list.
406    async fn write_fixed_stride_array(
407        object_writer: &mut dyn Writer,
408        field: &Field,
409        arrs: &[&dyn Array],
410        batch_id: i32,
411        page_table: &mut PageTable,
412    ) -> Result<()> {
413        assert_eq!(field.encoding, Some(Encoding::Plain));
414        assert!(!arrs.is_empty());
415        let data_type = arrs[0].data_type();
416
417        let mut encoder = PlainEncoder::new(object_writer, data_type);
418        let pos = encoder.encode(arrs).await?;
419        let arrs_length: i32 = arrs.iter().map(|a| a.len() as i32).sum();
420        let page_info = PageInfo::new(pos, arrs_length as usize);
421        page_table.set(field.id, batch_id, page_info);
422        Ok(())
423    }
424
425    /// Write var-length binary arrays.
426    async fn write_binary_array(
427        object_writer: &mut dyn Writer,
428        field: &Field,
429        arrs: &[&dyn Array],
430        batch_id: i32,
431        page_table: &mut PageTable,
432    ) -> Result<()> {
433        assert_eq!(field.encoding, Some(Encoding::VarBinary));
434        let mut encoder = BinaryEncoder::new(object_writer);
435        let pos = encoder.encode(arrs).await?;
436        let arrs_length: i32 = arrs.iter().map(|a| a.len() as i32).sum();
437        let page_info = PageInfo::new(pos, arrs_length as usize);
438        page_table.set(field.id, batch_id, page_info);
439        Ok(())
440    }
441
442    async fn write_dictionary_arr(
443        object_writer: &mut dyn Writer,
444        field: &Field,
445        arrs: &[&dyn Array],
446        key_type: &DataType,
447        batch_id: i32,
448        page_table: &mut PageTable,
449    ) -> Result<()> {
450        assert_eq!(field.encoding, Some(Encoding::Dictionary));
451
452        // Write the dictionary keys.
453        let mut encoder = DictionaryEncoder::new(object_writer, key_type);
454        let pos = encoder.encode(arrs).await?;
455        let arrs_length: i32 = arrs.iter().map(|a| a.len() as i32).sum();
456        let page_info = PageInfo::new(pos, arrs_length as usize);
457        page_table.set(field.id, batch_id, page_info);
458        Ok(())
459    }
460
461    #[async_recursion]
462    async fn write_struct_array(
463        object_writer: &mut dyn Writer,
464        field: &Field,
465        arrays: &[&StructArray],
466        batch_id: i32,
467        page_table: &mut PageTable,
468    ) -> Result<()> {
469        arrays
470            .iter()
471            .for_each(|a| assert_eq!(a.num_columns(), field.children.len()));
472
473        for child in &field.children {
474            let mut arrs: Vec<&ArrayRef> = Vec::new();
475            for struct_array in arrays {
476                let arr = struct_array
477                    .column_by_name(&child.name)
478                    .ok_or(Error::schema(format!(
479                        "FileWriter: schema mismatch: column {} does not exist in array: {:?}",
480                        child.name,
481                        struct_array.data_type()
482                    )))?;
483                arrs.push(arr);
484            }
485            Self::write_array(object_writer, child, arrs.as_slice(), batch_id, page_table).await?;
486        }
487        Ok(())
488    }
489
490    async fn write_list_array(
491        object_writer: &mut dyn Writer,
492        field: &Field,
493        arrs: &[&dyn Array],
494        batch_id: i32,
495        page_table: &mut PageTable,
496    ) -> Result<()> {
497        let capacity: usize = arrs.iter().map(|a| a.len()).sum();
498        let mut list_arrs: Vec<ArrayRef> = Vec::new();
499        let mut pos_builder: PrimitiveBuilder<Int32Type> =
500            PrimitiveBuilder::with_capacity(capacity);
501
502        let mut last_offset: usize = 0;
503        pos_builder.append_value(last_offset as i32);
504        for array in arrs.iter() {
505            let list_arr = as_list_array(*array);
506            let offsets = list_arr.value_offsets();
507
508            assert!(!offsets.is_empty());
509            let start_offset = offsets[0].as_usize();
510            let end_offset = offsets[offsets.len() - 1].as_usize();
511
512            let list_values = list_arr.values();
513            let sliced_values = list_values.slice(start_offset, end_offset - start_offset);
514            list_arrs.push(sliced_values);
515
516            offsets
517                .iter()
518                .skip(1)
519                .map(|b| b.as_usize() - start_offset + last_offset)
520                .for_each(|o| pos_builder.append_value(o as i32));
521            last_offset = pos_builder.values_slice()[pos_builder.len() - 1_usize] as usize;
522        }
523
524        let positions: &dyn Array = &pos_builder.finish();
525        Self::write_fixed_stride_array(object_writer, field, &[positions], batch_id, page_table)
526            .await?;
527        let arrs = list_arrs.iter().collect::<Vec<_>>();
528        Self::write_array(
529            object_writer,
530            &field.children[0],
531            arrs.as_slice(),
532            batch_id,
533            page_table,
534        )
535        .await
536    }
537
538    async fn write_large_list_array(
539        object_writer: &mut dyn Writer,
540        field: &Field,
541        arrs: &[&dyn Array],
542        batch_id: i32,
543        page_table: &mut PageTable,
544    ) -> Result<()> {
545        let capacity: usize = arrs.iter().map(|a| a.len()).sum();
546        let mut list_arrs: Vec<ArrayRef> = Vec::new();
547        let mut pos_builder: PrimitiveBuilder<Int64Type> =
548            PrimitiveBuilder::with_capacity(capacity);
549
550        let mut last_offset: usize = 0;
551        pos_builder.append_value(last_offset as i64);
552        for array in arrs.iter() {
553            let list_arr = as_large_list_array(*array);
554            let offsets = list_arr.value_offsets();
555
556            assert!(!offsets.is_empty());
557            let start_offset = offsets[0].as_usize();
558            let end_offset = offsets[offsets.len() - 1].as_usize();
559
560            let sliced_values = list_arr
561                .values()
562                .slice(start_offset, end_offset - start_offset);
563            list_arrs.push(sliced_values);
564
565            offsets
566                .iter()
567                .skip(1)
568                .map(|b| b.as_usize() - start_offset + last_offset)
569                .for_each(|o| pos_builder.append_value(o as i64));
570            last_offset = pos_builder.values_slice()[pos_builder.len() - 1_usize] as usize;
571        }
572
573        let positions: &dyn Array = &pos_builder.finish();
574        Self::write_fixed_stride_array(object_writer, field, &[positions], batch_id, page_table)
575            .await?;
576        let arrs = list_arrs.iter().collect::<Vec<_>>();
577        Self::write_array(
578            object_writer,
579            &field.children[0],
580            arrs.as_slice(),
581            batch_id,
582            page_table,
583        )
584        .await
585    }
586
587    async fn write_statistics(&mut self) -> Result<Option<StatisticsMetadata>> {
588        let statistics = self
589            .stats_collector
590            .as_mut()
591            .map(|collector| collector.finish());
592
593        match statistics {
594            Some(Ok(stats_batch)) if stats_batch.num_rows() > 0 => {
595                debug_assert_eq!(self.next_batch_id() as usize, stats_batch.num_rows());
596                let schema = Schema::try_from(stats_batch.schema().as_ref())?;
597                let leaf_field_ids = schema.field_ids();
598
599                let mut stats_page_table = PageTable::default();
600                for (i, field) in schema.fields.iter().enumerate() {
601                    Self::write_array(
602                        self.object_writer.as_mut(),
603                        field,
604                        &[stats_batch.column(i)],
605                        0, // Only one batch for statistics.
606                        &mut stats_page_table,
607                    )
608                    .await?;
609                }
610
611                let page_table_position = stats_page_table
612                    .write(self.object_writer.as_mut(), 0)
613                    .await?;
614
615                Ok(Some(StatisticsMetadata {
616                    schema,
617                    leaf_field_ids,
618                    page_table_position,
619                }))
620            }
621            Some(Err(e)) => Err(e),
622            _ => Ok(None),
623        }
624    }
625
626    /// Writes the dictionaries (using plain/binary encoding) into the file
627    ///
628    /// The offsets and lengths of the written buffers are stored in the given
629    /// schema so that the dictionaries can be loaded in the future.
630    async fn write_dictionaries(writer: &mut dyn Writer, schema: &mut Schema) -> Result<()> {
631        // Write dictionary values.
632        let max_field_id = schema.max_field_id().unwrap_or(-1);
633        for field_id in 0..max_field_id + 1 {
634            if let Some(field) = schema.mut_field_by_id(field_id)
635                && field.data_type().is_dictionary()
636            {
637                let dict_info = field.dictionary.as_mut().ok_or_else(|| {
638                    // and wrap it in here.
639                    Error::io(format!("Lance field {} misses dictionary info", field.name))
640                })?;
641
642                let value_arr = dict_info.values.as_ref().ok_or_else(|| {
643                    Error::invalid_input(format!(
644                        "Lance field {} is dictionary type, but misses the dictionary value array",
645                        field.name
646                    ))
647                })?;
648
649                let data_type = value_arr.data_type();
650                let pos = match data_type {
651                    dt if dt.is_numeric() => {
652                        let mut encoder = PlainEncoder::new(writer, dt);
653                        encoder.encode(&[value_arr]).await?
654                    }
655                    dt if dt.is_binary_like() => {
656                        let mut encoder = BinaryEncoder::new(writer);
657                        encoder.encode(&[value_arr]).await?
658                    }
659                    _ => {
660                        return Err(Error::schema(format!(
661                            "Does not support {} as dictionary value type",
662                            value_arr.data_type()
663                        )));
664                    }
665                };
666                dict_info.offset = pos;
667                dict_info.length = value_arr.len();
668            }
669        }
670        Ok(())
671    }
672
673    async fn write_footer(&mut self) -> Result<()> {
674        // Step 1. Write page table.
675        let field_id_offset = *self.schema.field_ids().iter().min().unwrap();
676        let pos = self
677            .page_table
678            .write(self.object_writer.as_mut(), field_id_offset)
679            .await?;
680        self.metadata.page_table_position = pos;
681
682        // Step 2. Write statistics.
683        self.metadata.stats_metadata = self.write_statistics().await?;
684
685        // Step 3. Write manifest and dictionary values.
686        Self::write_dictionaries(self.object_writer.as_mut(), &mut self.schema).await?;
687        let pos = M::store_schema(self.object_writer.as_mut(), &self.schema).await?;
688
689        // Step 4. Write metadata.
690        self.metadata.manifest_position = pos;
691        let pos = self.object_writer.write_struct(&self.metadata).await?;
692
693        // Step 5. Write magics.
694        self.object_writer
695            .write_magics(pos, MAJOR_VERSION, MINOR_VERSION, MAGIC)
696            .await
697    }
698}
699
700/// Walk through the schema and return arrays with their Lance field.
701///
702/// This skips over nested arrays and fields within list arrays. It does walk
703/// over the children of structs.
704fn fields_in_batches<'a>(
705    batches: &'a [RecordBatch],
706    schema: &'a Schema,
707) -> impl Iterator<Item = (&'a Field, Vec<&'a ArrayRef>)> {
708    let num_columns = batches[0].num_columns();
709    let array_iters = (0..num_columns).map(|col_i| {
710        batches
711            .iter()
712            .map(|batch| batch.column(col_i))
713            .collect::<Vec<_>>()
714    });
715    let mut to_visit: Vec<(&'a Field, Vec<&'a ArrayRef>)> =
716        schema.fields.iter().zip(array_iters).collect();
717
718    std::iter::from_fn(move || {
719        loop {
720            let (field, arrays): (_, Vec<&'a ArrayRef>) = to_visit.pop()?;
721            match field.data_type() {
722                DataType::Struct(_) => {
723                    for (i, child_field) in field.children.iter().enumerate() {
724                        let child_arrays = arrays
725                            .iter()
726                            .map(|arr| as_struct_array(*arr).column(i))
727                            .collect::<Vec<&'a ArrayRef>>();
728                        to_visit.push((child_field, child_arrays));
729                    }
730                    continue;
731                }
732                // We only walk structs right now.
733                _ if field.data_type().is_nested() => continue,
734                _ => return Some((field, arrays)),
735            }
736        }
737    })
738}
739
740#[cfg(test)]
741mod tests {
742    use super::*;
743
744    use std::sync::Arc;
745
746    use arrow_array::{
747        BooleanArray, Decimal128Array, Decimal256Array, DictionaryArray, DurationMicrosecondArray,
748        DurationMillisecondArray, DurationNanosecondArray, DurationSecondArray,
749        FixedSizeBinaryArray, FixedSizeListArray, Float32Array, Int32Array, Int64Array, ListArray,
750        NullArray, StringArray, TimestampMicrosecondArray, TimestampSecondArray, UInt8Array,
751        types::UInt32Type,
752    };
753    use arrow_buffer::i256;
754    use arrow_schema::{
755        Field as ArrowField, Fields as ArrowFields, Schema as ArrowSchema, TimeUnit,
756    };
757    use arrow_select::concat::concat_batches;
758
759    use crate::previous::reader::FileReader;
760
761    #[tokio::test]
762    async fn test_write_file() {
763        let arrow_schema = ArrowSchema::new(vec![
764            ArrowField::new("null", DataType::Null, true),
765            ArrowField::new("bool", DataType::Boolean, true),
766            ArrowField::new("i", DataType::Int64, true),
767            ArrowField::new("f", DataType::Float32, false),
768            ArrowField::new("b", DataType::Utf8, true),
769            ArrowField::new("decimal128", DataType::Decimal128(7, 3), false),
770            ArrowField::new("decimal256", DataType::Decimal256(7, 3), false),
771            ArrowField::new("duration_sec", DataType::Duration(TimeUnit::Second), false),
772            ArrowField::new(
773                "duration_msec",
774                DataType::Duration(TimeUnit::Millisecond),
775                false,
776            ),
777            ArrowField::new(
778                "duration_usec",
779                DataType::Duration(TimeUnit::Microsecond),
780                false,
781            ),
782            ArrowField::new(
783                "duration_nsec",
784                DataType::Duration(TimeUnit::Nanosecond),
785                false,
786            ),
787            ArrowField::new(
788                "d",
789                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
790                true,
791            ),
792            ArrowField::new(
793                "fixed_size_list",
794                DataType::FixedSizeList(
795                    Arc::new(ArrowField::new("item", DataType::Float32, true)),
796                    16,
797                ),
798                true,
799            ),
800            ArrowField::new("fixed_size_binary", DataType::FixedSizeBinary(8), true),
801            ArrowField::new(
802                "l",
803                DataType::List(Arc::new(ArrowField::new("item", DataType::Utf8, true))),
804                true,
805            ),
806            ArrowField::new(
807                "large_l",
808                DataType::LargeList(Arc::new(ArrowField::new("item", DataType::Utf8, true))),
809                true,
810            ),
811            ArrowField::new(
812                "l_dict",
813                DataType::List(Arc::new(ArrowField::new(
814                    "item",
815                    DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
816                    true,
817                ))),
818                true,
819            ),
820            ArrowField::new(
821                "large_l_dict",
822                DataType::LargeList(Arc::new(ArrowField::new(
823                    "item",
824                    DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
825                    true,
826                ))),
827                true,
828            ),
829            ArrowField::new(
830                "s",
831                DataType::Struct(ArrowFields::from(vec![
832                    ArrowField::new("si", DataType::Int64, true),
833                    ArrowField::new("sb", DataType::Utf8, true),
834                ])),
835                true,
836            ),
837        ]);
838        let mut schema = Schema::try_from(&arrow_schema).unwrap();
839
840        let dict_vec = (0..100).map(|n| ["a", "b", "c"][n % 3]).collect::<Vec<_>>();
841        let dict_arr: DictionaryArray<UInt32Type> = dict_vec.into_iter().collect();
842
843        let fixed_size_list_arr = FixedSizeListArray::try_new_from_values(
844            Float32Array::from_iter((0..1600).map(|n| n as f32).collect::<Vec<_>>()),
845            16,
846        )
847        .unwrap();
848
849        let binary_data: [u8; 800] = [123; 800];
850        let fixed_size_binary_arr =
851            FixedSizeBinaryArray::try_new_from_values(&UInt8Array::from_iter(binary_data), 8)
852                .unwrap();
853
854        let list_offsets: Int32Array = (0..202).step_by(2).collect();
855        let list_values =
856            StringArray::from((0..200).map(|n| format!("str-{}", n)).collect::<Vec<_>>());
857        let list_arr: arrow_array::GenericListArray<i32> =
858            try_new_generic_list_array(list_values, &list_offsets).unwrap();
859
860        let large_list_offsets: Int64Array = (0..202).step_by(2).collect();
861        let large_list_values =
862            StringArray::from((0..200).map(|n| format!("str-{}", n)).collect::<Vec<_>>());
863        let large_list_arr: arrow_array::GenericListArray<i64> =
864            try_new_generic_list_array(large_list_values, &large_list_offsets).unwrap();
865
866        let list_dict_offsets: Int32Array = (0..202).step_by(2).collect();
867        let list_dict_vec = (0..200).map(|n| ["a", "b", "c"][n % 3]).collect::<Vec<_>>();
868        let list_dict_arr: DictionaryArray<UInt32Type> = list_dict_vec.into_iter().collect();
869        let list_dict_arr: arrow_array::GenericListArray<i32> =
870            try_new_generic_list_array(list_dict_arr, &list_dict_offsets).unwrap();
871
872        let large_list_dict_offsets: Int64Array = (0..202).step_by(2).collect();
873        let large_list_dict_vec = (0..200).map(|n| ["a", "b", "c"][n % 3]).collect::<Vec<_>>();
874        let large_list_dict_arr: DictionaryArray<UInt32Type> =
875            large_list_dict_vec.into_iter().collect();
876        let large_list_dict_arr: arrow_array::GenericListArray<i64> =
877            try_new_generic_list_array(large_list_dict_arr, &large_list_dict_offsets).unwrap();
878
879        let columns: Vec<ArrayRef> = vec![
880            Arc::new(NullArray::new(100)),
881            Arc::new(BooleanArray::from_iter(
882                (0..100).map(|f| Some(f % 3 == 0)).collect::<Vec<_>>(),
883            )),
884            Arc::new(Int64Array::from_iter((0..100).collect::<Vec<_>>())),
885            Arc::new(Float32Array::from_iter(
886                (0..100).map(|n| n as f32).collect::<Vec<_>>(),
887            )),
888            Arc::new(StringArray::from(
889                (0..100).map(|n| n.to_string()).collect::<Vec<_>>(),
890            )),
891            Arc::new(
892                Decimal128Array::from_iter_values(0..100)
893                    .with_precision_and_scale(7, 3)
894                    .unwrap(),
895            ),
896            Arc::new(
897                Decimal256Array::from_iter_values((0..100).map(|v| i256::from_i128(v as i128)))
898                    .with_precision_and_scale(7, 3)
899                    .unwrap(),
900            ),
901            Arc::new(DurationSecondArray::from_iter_values(0..100)),
902            Arc::new(DurationMillisecondArray::from_iter_values(0..100)),
903            Arc::new(DurationMicrosecondArray::from_iter_values(0..100)),
904            Arc::new(DurationNanosecondArray::from_iter_values(0..100)),
905            Arc::new(dict_arr),
906            Arc::new(fixed_size_list_arr),
907            Arc::new(fixed_size_binary_arr),
908            Arc::new(list_arr),
909            Arc::new(large_list_arr),
910            Arc::new(list_dict_arr),
911            Arc::new(large_list_dict_arr),
912            Arc::new(StructArray::from(vec![
913                (
914                    Arc::new(ArrowField::new("si", DataType::Int64, true)),
915                    Arc::new(Int64Array::from_iter((100..200).collect::<Vec<_>>())) as ArrayRef,
916                ),
917                (
918                    Arc::new(ArrowField::new("sb", DataType::Utf8, true)),
919                    Arc::new(StringArray::from(
920                        (0..100).map(|n| n.to_string()).collect::<Vec<_>>(),
921                    )) as ArrayRef,
922                ),
923            ])),
924        ];
925        let batch = RecordBatch::try_new(Arc::new(arrow_schema), columns).unwrap();
926        schema.set_dictionary(&batch).unwrap();
927
928        let store = ObjectStore::memory();
929        let path = Path::from("/foo");
930        let mut file_writer = FileWriter::<NotSelfDescribing>::try_new(
931            &store,
932            &path,
933            schema.clone(),
934            &Default::default(),
935        )
936        .await
937        .unwrap();
938        file_writer
939            .write(std::slice::from_ref(&batch))
940            .await
941            .unwrap();
942        file_writer.finish().await.unwrap();
943
944        let reader = FileReader::try_new(&store, &path, schema).await.unwrap();
945        let actual = reader.read_batch(0, .., reader.schema()).await.unwrap();
946        assert_eq!(actual, batch);
947    }
948
949    #[tokio::test]
950    async fn test_dictionary_first_element_file() {
951        let arrow_schema = ArrowSchema::new(vec![ArrowField::new(
952            "d",
953            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
954            true,
955        )]);
956        let mut schema = Schema::try_from(&arrow_schema).unwrap();
957
958        let dict_vec = (0..100).map(|n| ["a", "b", "c"][n % 3]).collect::<Vec<_>>();
959        let dict_arr: DictionaryArray<UInt32Type> = dict_vec.into_iter().collect();
960
961        let columns: Vec<ArrayRef> = vec![Arc::new(dict_arr)];
962        let batch = RecordBatch::try_new(Arc::new(arrow_schema), columns).unwrap();
963        schema.set_dictionary(&batch).unwrap();
964
965        let store = ObjectStore::memory();
966        let path = Path::from("/foo");
967        let mut file_writer = FileWriter::<NotSelfDescribing>::try_new(
968            &store,
969            &path,
970            schema.clone(),
971            &Default::default(),
972        )
973        .await
974        .unwrap();
975        file_writer
976            .write(std::slice::from_ref(&batch))
977            .await
978            .unwrap();
979        file_writer.finish().await.unwrap();
980
981        let reader = FileReader::try_new(&store, &path, schema).await.unwrap();
982        let actual = reader.read_batch(0, .., reader.schema()).await.unwrap();
983        assert_eq!(actual, batch);
984    }
985
986    #[tokio::test]
987    async fn test_write_temporal_types() {
988        let arrow_schema = Arc::new(ArrowSchema::new(vec![
989            ArrowField::new(
990                "ts_notz",
991                DataType::Timestamp(TimeUnit::Second, None),
992                false,
993            ),
994            ArrowField::new(
995                "ts_tz",
996                DataType::Timestamp(TimeUnit::Microsecond, Some("America/Los_Angeles".into())),
997                false,
998            ),
999        ]));
1000        let columns: Vec<ArrayRef> = vec![
1001            Arc::new(TimestampSecondArray::from(vec![11111111, 22222222])),
1002            Arc::new(
1003                TimestampMicrosecondArray::from(vec![3333333, 4444444])
1004                    .with_timezone("America/Los_Angeles"),
1005            ),
1006        ];
1007        let batch = RecordBatch::try_new(arrow_schema.clone(), columns).unwrap();
1008
1009        let schema = Schema::try_from(arrow_schema.as_ref()).unwrap();
1010        let store = ObjectStore::memory();
1011        let path = Path::from("/foo");
1012        let mut file_writer = FileWriter::<NotSelfDescribing>::try_new(
1013            &store,
1014            &path,
1015            schema.clone(),
1016            &Default::default(),
1017        )
1018        .await
1019        .unwrap();
1020        file_writer
1021            .write(std::slice::from_ref(&batch))
1022            .await
1023            .unwrap();
1024        file_writer.finish().await.unwrap();
1025
1026        let reader = FileReader::try_new(&store, &path, schema).await.unwrap();
1027        let actual = reader.read_batch(0, .., reader.schema()).await.unwrap();
1028        assert_eq!(actual, batch);
1029    }
1030
1031    #[tokio::test]
1032    async fn test_collect_stats() {
1033        // Validate:
1034        // Only collects stats for requested columns
1035        // Can collect stats in nested structs
1036        // Won't collect stats for list columns (for now)
1037
1038        let arrow_schema = ArrowSchema::new(vec![
1039            ArrowField::new("i", DataType::Int64, true),
1040            ArrowField::new("i2", DataType::Int64, true),
1041            ArrowField::new(
1042                "l",
1043                DataType::List(Arc::new(ArrowField::new("item", DataType::Int32, true))),
1044                true,
1045            ),
1046            ArrowField::new(
1047                "s",
1048                DataType::Struct(ArrowFields::from(vec![
1049                    ArrowField::new("si", DataType::Int64, true),
1050                    ArrowField::new("sb", DataType::Utf8, true),
1051                ])),
1052                true,
1053            ),
1054        ]);
1055
1056        let schema = Schema::try_from(&arrow_schema).unwrap();
1057
1058        let store = ObjectStore::memory();
1059        let path = Path::from("/foo");
1060
1061        let options = FileWriterOptions {
1062            collect_stats_for_fields: Some(vec![0, 1, 5, 6]),
1063        };
1064        let mut file_writer =
1065            FileWriter::<NotSelfDescribing>::try_new(&store, &path, schema.clone(), &options)
1066                .await
1067                .unwrap();
1068
1069        let batch1 = RecordBatch::try_new(
1070            Arc::new(arrow_schema.clone()),
1071            vec![
1072                Arc::new(Int64Array::from(vec![1, 2, 3])),
1073                Arc::new(Int64Array::from(vec![4, 5, 6])),
1074                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
1075                    Some(vec![Some(1i32), Some(2), Some(3)]),
1076                    Some(vec![Some(4), Some(5)]),
1077                    Some(vec![]),
1078                ])),
1079                Arc::new(StructArray::from(vec![
1080                    (
1081                        Arc::new(ArrowField::new("si", DataType::Int64, true)),
1082                        Arc::new(Int64Array::from(vec![1, 2, 3])) as ArrayRef,
1083                    ),
1084                    (
1085                        Arc::new(ArrowField::new("sb", DataType::Utf8, true)),
1086                        Arc::new(StringArray::from(vec!["a", "b", "c"])) as ArrayRef,
1087                    ),
1088                ])),
1089            ],
1090        )
1091        .unwrap();
1092        file_writer.write(&[batch1]).await.unwrap();
1093
1094        let batch2 = RecordBatch::try_new(
1095            Arc::new(arrow_schema.clone()),
1096            vec![
1097                Arc::new(Int64Array::from(vec![5, 6])),
1098                Arc::new(Int64Array::from(vec![10, 11])),
1099                Arc::new(ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
1100                    Some(vec![Some(1i32), Some(2), Some(3)]),
1101                    Some(vec![]),
1102                ])),
1103                Arc::new(StructArray::from(vec![
1104                    (
1105                        Arc::new(ArrowField::new("si", DataType::Int64, true)),
1106                        Arc::new(Int64Array::from(vec![4, 5])) as ArrayRef,
1107                    ),
1108                    (
1109                        Arc::new(ArrowField::new("sb", DataType::Utf8, true)),
1110                        Arc::new(StringArray::from(vec!["d", "e"])) as ArrayRef,
1111                    ),
1112                ])),
1113            ],
1114        )
1115        .unwrap();
1116        file_writer.write(&[batch2]).await.unwrap();
1117
1118        file_writer.finish().await.unwrap();
1119
1120        let reader = FileReader::try_new(&store, &path, schema).await.unwrap();
1121
1122        let read_stats = reader.read_page_stats(&[0, 1, 5, 6]).await.unwrap();
1123        assert!(read_stats.is_some());
1124        let read_stats = read_stats.unwrap();
1125
1126        let expected_stats_schema = stats_schema([
1127            (0, DataType::Int64),
1128            (1, DataType::Int64),
1129            (5, DataType::Int64),
1130            (6, DataType::Utf8),
1131        ]);
1132
1133        assert_eq!(read_stats.schema().as_ref(), &expected_stats_schema);
1134
1135        let expected_stats = stats_batch(&[
1136            Stats {
1137                field_id: 0,
1138                null_counts: vec![0, 0],
1139                min_values: Arc::new(Int64Array::from(vec![1, 5])),
1140                max_values: Arc::new(Int64Array::from(vec![3, 6])),
1141            },
1142            Stats {
1143                field_id: 1,
1144                null_counts: vec![0, 0],
1145                min_values: Arc::new(Int64Array::from(vec![4, 10])),
1146                max_values: Arc::new(Int64Array::from(vec![6, 11])),
1147            },
1148            Stats {
1149                field_id: 5,
1150                null_counts: vec![0, 0],
1151                min_values: Arc::new(Int64Array::from(vec![1, 4])),
1152                max_values: Arc::new(Int64Array::from(vec![3, 5])),
1153            },
1154            // FIXME: these max values shouldn't be incremented
1155            // https://github.com/lancedb/lance/issues/1517
1156            Stats {
1157                field_id: 6,
1158                null_counts: vec![0, 0],
1159                min_values: Arc::new(StringArray::from(vec!["a", "d"])),
1160                max_values: Arc::new(StringArray::from(vec!["c", "e"])),
1161            },
1162        ]);
1163
1164        assert_eq!(read_stats, expected_stats);
1165    }
1166
1167    fn stats_schema(data_fields: impl IntoIterator<Item = (i32, DataType)>) -> ArrowSchema {
1168        let fields = data_fields
1169            .into_iter()
1170            .map(|(field_id, data_type)| {
1171                Arc::new(ArrowField::new(
1172                    format!("{}", field_id),
1173                    DataType::Struct(
1174                        vec![
1175                            Arc::new(ArrowField::new("null_count", DataType::Int64, false)),
1176                            Arc::new(ArrowField::new("min_value", data_type.clone(), true)),
1177                            Arc::new(ArrowField::new("max_value", data_type, true)),
1178                        ]
1179                        .into(),
1180                    ),
1181                    false,
1182                ))
1183            })
1184            .collect::<Vec<_>>();
1185        ArrowSchema::new(fields)
1186    }
1187
1188    struct Stats {
1189        field_id: i32,
1190        null_counts: Vec<i64>,
1191        min_values: ArrayRef,
1192        max_values: ArrayRef,
1193    }
1194
1195    fn stats_batch(stats: &[Stats]) -> RecordBatch {
1196        let schema = stats_schema(
1197            stats
1198                .iter()
1199                .map(|s| (s.field_id, s.min_values.data_type().clone())),
1200        );
1201
1202        let columns = stats
1203            .iter()
1204            .map(|s| {
1205                let data_type = s.min_values.data_type().clone();
1206                let fields = vec![
1207                    Arc::new(ArrowField::new("null_count", DataType::Int64, false)),
1208                    Arc::new(ArrowField::new("min_value", data_type.clone(), true)),
1209                    Arc::new(ArrowField::new("max_value", data_type, true)),
1210                ];
1211                let arrays = vec![
1212                    Arc::new(Int64Array::from(s.null_counts.clone())),
1213                    s.min_values.clone(),
1214                    s.max_values.clone(),
1215                ];
1216                Arc::new(StructArray::new(fields.into(), arrays, None)) as ArrayRef
1217            })
1218            .collect();
1219
1220        RecordBatch::try_new(Arc::new(schema), columns).unwrap()
1221    }
1222
1223    async fn read_file_as_one_batch(
1224        object_store: &ObjectStore,
1225        path: &Path,
1226        schema: Schema,
1227    ) -> RecordBatch {
1228        let reader = FileReader::try_new(object_store, path, schema)
1229            .await
1230            .unwrap();
1231        let mut batches = vec![];
1232        for i in 0..reader.num_batches() {
1233            batches.push(
1234                reader
1235                    .read_batch(i as i32, .., reader.schema())
1236                    .await
1237                    .unwrap(),
1238            );
1239        }
1240        let arrow_schema = Arc::new(reader.schema().into());
1241        concat_batches(&arrow_schema, &batches).unwrap()
1242    }
1243
1244    /// Test encoding arrays that share the same underneath buffer.
1245    #[tokio::test]
1246    async fn test_encode_slice() {
1247        let store = ObjectStore::memory();
1248        let path = Path::from("/shared_slice");
1249
1250        let arrow_schema = Arc::new(ArrowSchema::new(vec![ArrowField::new(
1251            "i",
1252            DataType::Int32,
1253            false,
1254        )]));
1255        let schema = Schema::try_from(arrow_schema.as_ref()).unwrap();
1256        let mut file_writer = FileWriter::<NotSelfDescribing>::try_new(
1257            &store,
1258            &path,
1259            schema.clone(),
1260            &Default::default(),
1261        )
1262        .await
1263        .unwrap();
1264
1265        let array = Int32Array::from_iter_values(0..1000);
1266
1267        for i in (0..1000).step_by(4) {
1268            let data = array.slice(i, 4);
1269            file_writer
1270                .write(&[RecordBatch::try_new(arrow_schema.clone(), vec![Arc::new(data)]).unwrap()])
1271                .await
1272                .unwrap();
1273        }
1274        file_writer.finish().await.unwrap();
1275        assert!(store.size(&path).await.unwrap() < 2 * 8 * 1000);
1276
1277        let batch = read_file_as_one_batch(&store, &path, schema).await;
1278        assert_eq!(batch.column_by_name("i").unwrap().as_ref(), &array);
1279    }
1280
1281    #[tokio::test]
1282    async fn test_write_schema_with_holes() {
1283        let store = ObjectStore::memory();
1284        let path = Path::from("test");
1285
1286        let mut field0 = Field::try_from(&ArrowField::new("a", DataType::Int32, true)).unwrap();
1287        field0.set_id(-1, &mut 0);
1288        assert_eq!(field0.id, 0);
1289        let mut field2 = Field::try_from(&ArrowField::new("b", DataType::Int32, true)).unwrap();
1290        field2.set_id(-1, &mut 2);
1291        assert_eq!(field2.id, 2);
1292        // There is a hole at field id 1.
1293        let schema = Schema {
1294            fields: vec![field0, field2],
1295            metadata: Default::default(),
1296        };
1297
1298        let arrow_schema = Arc::new(ArrowSchema::new(vec![
1299            ArrowField::new("a", DataType::Int32, true),
1300            ArrowField::new("b", DataType::Int32, true),
1301        ]));
1302        let data = RecordBatch::try_new(
1303            arrow_schema.clone(),
1304            vec![
1305                Arc::new(Int32Array::from_iter_values(0..10)),
1306                Arc::new(Int32Array::from_iter_values(10..20)),
1307            ],
1308        )
1309        .unwrap();
1310
1311        let mut file_writer = FileWriter::<NotSelfDescribing>::try_new(
1312            &store,
1313            &path,
1314            schema.clone(),
1315            &Default::default(),
1316        )
1317        .await
1318        .unwrap();
1319        file_writer.write(&[data]).await.unwrap();
1320        file_writer.finish().await.unwrap();
1321
1322        let page_table = file_writer.page_table;
1323        assert!(page_table.get(0, 0).is_some());
1324        assert!(page_table.get(2, 0).is_some());
1325    }
1326}