iceberg/writer/base_writer/
data_file_writer.rs1use 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#[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 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#[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}