Skip to main content

iceberg/writer/base_writer/
data_file_writer.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! This module provide `DataFileWriter`.
19
20use arrow_array::RecordBatch;
21
22use crate::error::invalid_data;
23use crate::spec::{DataContentType, DataFile, PartitionKey};
24use crate::writer::file_writer::FileWriterBuilder;
25use crate::writer::file_writer::location_generator::{FileNameGenerator, LocationGenerator};
26use crate::writer::file_writer::rolling_writer::{RollingFileWriter, RollingFileWriterBuilder};
27use crate::writer::{CurrentFileStatus, IcebergWriter, IcebergWriterBuilder};
28use crate::{Error, ErrorKind, Result};
29
30/// Builder for `DataFileWriter`.
31#[derive(Debug)]
32pub struct DataFileWriterBuilder<B: FileWriterBuilder, L: LocationGenerator, F: FileNameGenerator> {
33    inner: RollingFileWriterBuilder<B, L, F>,
34}
35
36impl<B, L, F> DataFileWriterBuilder<B, L, F>
37where
38    B: FileWriterBuilder,
39    L: LocationGenerator,
40    F: FileNameGenerator,
41{
42    /// Create a new `DataFileWriterBuilder` using a `RollingFileWriterBuilder`.
43    pub fn new(inner: RollingFileWriterBuilder<B, L, F>) -> Self {
44        Self { inner }
45    }
46}
47
48#[async_trait::async_trait]
49impl<B, L, F> IcebergWriterBuilder for DataFileWriterBuilder<B, L, F>
50where
51    B: FileWriterBuilder,
52    L: LocationGenerator,
53    F: FileNameGenerator,
54{
55    type R = DataFileWriter<B, L, F>;
56
57    async fn build(&self, partition_key: Option<PartitionKey>) -> Result<Self::R> {
58        Ok(DataFileWriter {
59            inner: Some(self.inner.build()),
60            partition_key,
61        })
62    }
63}
64
65/// A writer write data is within one spec/partition.
66#[derive(Debug)]
67pub struct DataFileWriter<B: FileWriterBuilder, L: LocationGenerator, F: FileNameGenerator> {
68    inner: Option<RollingFileWriter<B, L, F>>,
69    partition_key: Option<PartitionKey>,
70}
71
72#[async_trait::async_trait]
73impl<B, L, F> IcebergWriter for DataFileWriter<B, L, F>
74where
75    B: FileWriterBuilder,
76    L: LocationGenerator,
77    F: FileNameGenerator,
78{
79    async fn write(&mut self, batch: RecordBatch) -> Result<()> {
80        if let Some(writer) = self.inner.as_mut() {
81            writer.write(&self.partition_key, &batch).await
82        } else {
83            Err(Error::new(
84                ErrorKind::Unexpected,
85                "Writer is not initialized!",
86            ))
87        }
88    }
89
90    async fn close(&mut self) -> Result<Vec<DataFile>> {
91        if let Some(writer) = self.inner.take() {
92            writer
93                .close()
94                .await?
95                .into_iter()
96                .map(|mut res| {
97                    res.content(DataContentType::Data);
98                    if let Some(pk) = self.partition_key.as_ref() {
99                        res.partition(pk.data().clone());
100                        res.partition_spec_id(pk.spec().spec_id());
101                    }
102                    res.build()
103                        .map_err(|e| invalid_data!("Failed to build data file: {e}"))
104                })
105                .collect()
106        } else {
107            Err(Error::new(
108                ErrorKind::Unexpected,
109                "Data file writer has been closed.",
110            ))
111        }
112    }
113}
114
115impl<B, L, F> CurrentFileStatus for DataFileWriter<B, L, F>
116where
117    B: FileWriterBuilder,
118    L: LocationGenerator,
119    F: FileNameGenerator,
120{
121    fn current_file_path(&self) -> String {
122        self.inner.as_ref().unwrap().current_file_path()
123    }
124
125    fn current_row_num(&self) -> usize {
126        self.inner.as_ref().unwrap().current_row_num()
127    }
128
129    fn current_written_size(&self) -> usize {
130        self.inner.as_ref().unwrap().current_written_size()
131    }
132}
133
134#[cfg(test)]
135mod test {
136    use std::collections::HashMap;
137    use std::sync::Arc;
138
139    use arrow_array::{Int32Array, StringArray};
140    use arrow_schema::{DataType, Field};
141    use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
142    use parquet::arrow::arrow_reader::{ArrowReaderMetadata, ArrowReaderOptions};
143    use parquet::file::properties::WriterProperties;
144    use tempfile::TempDir;
145
146    use crate::Result;
147    use crate::io::FileIO;
148    use crate::spec::{
149        DataContentType, DataFileFormat, Literal, NestedField, PartitionKey, PartitionSpec,
150        PrimitiveType, Schema, Struct, Type,
151    };
152    use crate::writer::base_writer::data_file_writer::DataFileWriterBuilder;
153    use crate::writer::file_writer::ParquetWriterBuilder;
154    use crate::writer::file_writer::location_generator::{
155        DefaultFileNameGenerator, DefaultLocationGenerator,
156    };
157    use crate::writer::file_writer::rolling_writer::RollingFileWriterBuilder;
158    use crate::writer::{IcebergWriter, IcebergWriterBuilder, RecordBatch};
159
160    #[tokio::test]
161    async fn test_parquet_writer() -> Result<()> {
162        let temp_dir = TempDir::new().unwrap();
163        let file_io = FileIO::new_with_fs();
164        let location_gen = DefaultLocationGenerator::with_data_location(
165            temp_dir.path().to_str().unwrap().to_string(),
166        );
167        let file_name_gen =
168            DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
169
170        let schema = Schema::builder()
171            .with_schema_id(3)
172            .with_fields(vec![
173                NestedField::required(3, "foo", Type::Primitive(PrimitiveType::Int)).into(),
174                NestedField::required(4, "bar", Type::Primitive(PrimitiveType::String)).into(),
175            ])
176            .build()?;
177
178        let pw = ParquetWriterBuilder::new(WriterProperties::builder().build(), Arc::new(schema));
179
180        let rolling_file_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
181            pw,
182            file_io.clone(),
183            location_gen,
184            file_name_gen,
185        );
186
187        let mut data_file_writer = DataFileWriterBuilder::new(rolling_file_writer_builder)
188            .build(None)
189            .await
190            .unwrap();
191
192        let arrow_schema = arrow_schema::Schema::new(vec![
193            Field::new("foo", DataType::Int32, false).with_metadata(HashMap::from([(
194                PARQUET_FIELD_ID_META_KEY.to_string(),
195                3.to_string(),
196            )])),
197            Field::new("bar", DataType::Utf8, false).with_metadata(HashMap::from([(
198                PARQUET_FIELD_ID_META_KEY.to_string(),
199                4.to_string(),
200            )])),
201        ]);
202        let batch = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
203            Arc::new(Int32Array::from(vec![1, 2, 3])),
204            Arc::new(StringArray::from(vec!["Alice", "Bob", "Charlie"])),
205        ])?;
206        data_file_writer.write(batch).await?;
207
208        let data_files = data_file_writer.close().await.unwrap();
209        assert_eq!(data_files.len(), 1);
210
211        let data_file = &data_files[0];
212        assert_eq!(data_file.file_format, DataFileFormat::Parquet);
213        assert_eq!(data_file.content, DataContentType::Data);
214        assert_eq!(data_file.partition, Struct::empty());
215
216        let input_file = file_io.new_input(data_file.file_path.clone())?;
217        let input_content = input_file.read().await?;
218
219        let parquet_reader =
220            ArrowReaderMetadata::load(&input_content, ArrowReaderOptions::default())
221                .expect("Failed to load Parquet metadata");
222
223        let field_ids: Vec<i32> = parquet_reader
224            .parquet_schema()
225            .columns()
226            .iter()
227            .map(|col| col.self_type().get_basic_info().id())
228            .collect();
229
230        assert_eq!(field_ids, vec![3, 4]);
231        Ok(())
232    }
233
234    #[tokio::test]
235    async fn test_parquet_writer_with_partition() -> Result<()> {
236        let temp_dir = TempDir::new().unwrap();
237        let file_io = FileIO::new_with_fs();
238        let location_gen = DefaultLocationGenerator::with_data_location(
239            temp_dir.path().to_str().unwrap().to_string(),
240        );
241        let file_name_gen = DefaultFileNameGenerator::new(
242            "test_partitioned".to_string(),
243            None,
244            DataFileFormat::Parquet,
245        );
246
247        let schema = Schema::builder()
248            .with_schema_id(5)
249            .with_fields(vec![
250                NestedField::required(5, "id", Type::Primitive(PrimitiveType::Int)).into(),
251                NestedField::required(6, "name", Type::Primitive(PrimitiveType::String)).into(),
252            ])
253            .build()?;
254        let schema_ref = Arc::new(schema);
255
256        let partition_value = Struct::from_iter([Some(Literal::int(1))]);
257        let partition_key = PartitionKey::new(
258            PartitionSpec::builder(schema_ref.clone()).build()?,
259            schema_ref.clone(),
260            partition_value.clone(),
261        );
262
263        let parquet_writer_builder =
264            ParquetWriterBuilder::new(WriterProperties::builder().build(), schema_ref.clone());
265
266        let rolling_file_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
267            parquet_writer_builder,
268            file_io.clone(),
269            location_gen,
270            file_name_gen,
271        );
272
273        let mut data_file_writer = DataFileWriterBuilder::new(rolling_file_writer_builder)
274            .build(Some(partition_key))
275            .await?;
276
277        let arrow_schema = arrow_schema::Schema::new(vec![
278            Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
279                PARQUET_FIELD_ID_META_KEY.to_string(),
280                5.to_string(),
281            )])),
282            Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
283                PARQUET_FIELD_ID_META_KEY.to_string(),
284                6.to_string(),
285            )])),
286        ]);
287        let batch = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
288            Arc::new(Int32Array::from(vec![1, 2, 3])),
289            Arc::new(StringArray::from(vec!["Alice", "Bob", "Charlie"])),
290        ])?;
291        data_file_writer.write(batch).await?;
292
293        let data_files = data_file_writer.close().await.unwrap();
294        assert_eq!(data_files.len(), 1);
295
296        let data_file = &data_files[0];
297        assert_eq!(data_file.file_format, DataFileFormat::Parquet);
298        assert_eq!(data_file.content, DataContentType::Data);
299        assert_eq!(data_file.partition, partition_value);
300
301        let input_file = file_io.new_input(data_file.file_path.clone())?;
302        let input_content = input_file.read().await?;
303
304        let parquet_reader =
305            ArrowReaderMetadata::load(&input_content, ArrowReaderOptions::default())?;
306
307        let field_ids: Vec<i32> = parquet_reader
308            .parquet_schema()
309            .columns()
310            .iter()
311            .map(|col| col.self_type().get_basic_info().id())
312            .collect();
313        assert_eq!(field_ids, vec![5, 6]);
314
315        let field_names: Vec<&str> = parquet_reader
316            .parquet_schema()
317            .columns()
318            .iter()
319            .map(|col| col.name())
320            .collect();
321        assert_eq!(field_names, vec!["id", "name"]);
322
323        Ok(())
324    }
325}