Skip to main content

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