1use std::fmt::{Debug, Formatter};
19
20use arrow_array::RecordBatch;
21
22use crate::io::{FileIO, OutputFile};
23use crate::spec::{DataFileBuilder, PartitionKey, TableProperties};
24use crate::writer::CurrentFileStatus;
25use crate::writer::file_writer::location_generator::{FileNameGenerator, LocationGenerator};
26use crate::writer::file_writer::{FileWriter, FileWriterBuilder};
27use crate::{Error, ErrorKind, Result};
28
29#[derive(Clone, Debug)]
31pub struct RollingFileWriterBuilder<
32 B: FileWriterBuilder,
33 L: LocationGenerator,
34 F: FileNameGenerator,
35> {
36 inner_builder: B,
37 target_file_size: usize,
38 file_io: FileIO,
39 location_generator: L,
40 file_name_generator: F,
41}
42
43impl<B, L, F> RollingFileWriterBuilder<B, L, F>
44where
45 B: FileWriterBuilder,
46 L: LocationGenerator,
47 F: FileNameGenerator,
48{
49 pub fn new(
63 inner_builder: B,
64 target_file_size: usize,
65 file_io: FileIO,
66 location_generator: L,
67 file_name_generator: F,
68 ) -> Self {
69 Self {
70 inner_builder,
71 target_file_size,
72 file_io,
73 location_generator,
74 file_name_generator,
75 }
76 }
77
78 pub fn new_with_default_file_size(
91 inner_builder: B,
92 file_io: FileIO,
93 location_generator: L,
94 file_name_generator: F,
95 ) -> Self {
96 Self {
97 inner_builder,
98 target_file_size: TableProperties::PROPERTY_WRITE_TARGET_FILE_SIZE_BYTES_DEFAULT,
99 file_io,
100 location_generator,
101 file_name_generator,
102 }
103 }
104
105 pub fn build(&self) -> RollingFileWriter<B, L, F> {
107 RollingFileWriter {
108 inner: None,
109 inner_builder: self.inner_builder.clone(),
110 target_file_size: self.target_file_size,
111 data_file_builders: vec![],
112 file_io: self.file_io.clone(),
113 location_generator: self.location_generator.clone(),
114 file_name_generator: self.file_name_generator.clone(),
115 }
116 }
117}
118
119pub struct RollingFileWriter<B: FileWriterBuilder, L: LocationGenerator, F: FileNameGenerator> {
126 inner: Option<B::R>,
127 inner_builder: B,
128 target_file_size: usize,
129 data_file_builders: Vec<DataFileBuilder>,
130 file_io: FileIO,
131 location_generator: L,
132 file_name_generator: F,
133}
134
135impl<B, L, F> Debug for RollingFileWriter<B, L, F>
136where
137 B: FileWriterBuilder,
138 L: LocationGenerator,
139 F: FileNameGenerator,
140{
141 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
142 f.debug_struct("RollingFileWriter")
143 .field("target_file_size", &self.target_file_size)
144 .field("file_io", &self.file_io)
145 .finish()
146 }
147}
148
149impl<B, L, F> RollingFileWriter<B, L, F>
150where
151 B: FileWriterBuilder,
152 L: LocationGenerator,
153 F: FileNameGenerator,
154{
155 fn should_roll(&self) -> bool {
161 self.current_written_size() > self.target_file_size
162 }
163
164 fn new_output_file(&self, partition_key: &Option<PartitionKey>) -> Result<OutputFile> {
165 self.file_io
166 .new_output(self.location_generator.generate_location(
167 partition_key.as_ref(),
168 &self.file_name_generator.generate_file_name(),
169 ))
170 }
171
172 pub async fn write(
187 &mut self,
188 partition_key: &Option<PartitionKey>,
189 input: &RecordBatch,
190 ) -> Result<()> {
191 if self.inner.is_none() {
192 self.inner = Some(
194 self.inner_builder
195 .build(self.new_output_file(partition_key)?)
196 .await?,
197 );
198 }
199
200 if self.should_roll()
201 && let Some(inner) = self.inner.take()
202 {
203 self.data_file_builders.extend(inner.close().await?);
205
206 self.inner = Some(
208 self.inner_builder
209 .build(self.new_output_file(partition_key)?)
210 .await?,
211 );
212 }
213
214 if let Some(writer) = self.inner.as_mut() {
216 Ok(writer.write(input).await?)
217 } else {
218 Err(Error::new(
219 ErrorKind::Unexpected,
220 "Writer is not initialized!",
221 ))
222 }
223 }
224
225 pub async fn close(mut self) -> Result<Vec<DataFileBuilder>> {
232 if let Some(current_writer) = self.inner {
234 self.data_file_builders
235 .extend(current_writer.close().await?);
236 }
237
238 Ok(self.data_file_builders)
239 }
240}
241
242impl<B: FileWriterBuilder, L: LocationGenerator, F: FileNameGenerator> CurrentFileStatus
243 for RollingFileWriter<B, L, F>
244{
245 fn current_file_path(&self) -> String {
246 self.inner.as_ref().unwrap().current_file_path()
247 }
248
249 fn current_row_num(&self) -> usize {
250 self.inner.as_ref().unwrap().current_row_num()
251 }
252
253 fn current_written_size(&self) -> usize {
254 self.inner.as_ref().unwrap().current_written_size()
255 }
256}
257
258#[cfg(test)]
259mod tests {
260 use std::collections::{HashMap, HashSet};
261 use std::sync::Arc;
262
263 use arrow_array::{ArrayRef, Int32Array, StringArray};
264 use arrow_schema::{DataType, Field, Schema as ArrowSchema};
265 use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
266 use parquet::file::properties::WriterProperties;
267 use rand::prelude::IteratorRandom;
268 use tempfile::TempDir;
269
270 use super::*;
271 use crate::arrow::test_utils::read_encrypted_parquet;
272 use crate::encryption::StandardKeyMetadata;
273 use crate::io::FileIO;
274 use crate::spec::{DataFileFormat, NestedField, PrimitiveType, Schema, Type};
275 use crate::test_utils::make_encryption_manager;
276 use crate::writer::base_writer::data_file_writer::DataFileWriterBuilder;
277 use crate::writer::file_writer::ParquetWriterBuilder;
278 use crate::writer::file_writer::location_generator::{
279 DefaultFileNameGenerator, DefaultLocationGenerator,
280 };
281 use crate::writer::tests::check_parquet_data_file;
282 use crate::writer::{IcebergWriter, IcebergWriterBuilder, RecordBatch};
283
284 fn make_test_schema() -> Result<Schema> {
285 Schema::builder()
286 .with_schema_id(1)
287 .with_fields(vec![
288 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
289 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
290 ])
291 .build()
292 }
293
294 fn make_test_arrow_schema() -> ArrowSchema {
295 ArrowSchema::new(vec![
296 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
297 PARQUET_FIELD_ID_META_KEY.to_string(),
298 1.to_string(),
299 )])),
300 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
301 PARQUET_FIELD_ID_META_KEY.to_string(),
302 2.to_string(),
303 )])),
304 ])
305 }
306
307 #[tokio::test]
308 async fn test_rolling_writer_basic() -> Result<()> {
309 let temp_dir = TempDir::new()?;
310 let file_io = FileIO::new_with_fs();
311 let location_gen = DefaultLocationGenerator::with_data_location(
312 temp_dir.path().to_str().unwrap().to_string(),
313 );
314 let file_name_gen =
315 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
316
317 let schema = make_test_schema()?;
319
320 let parquet_writer_builder =
322 ParquetWriterBuilder::new(WriterProperties::builder().build(), Arc::new(schema));
323
324 let rolling_file_writer_builder = RollingFileWriterBuilder::new(
326 parquet_writer_builder,
327 1024 * 1024,
328 file_io.clone(),
329 location_gen,
330 file_name_gen,
331 );
332
333 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_file_writer_builder);
334
335 let mut writer = data_file_writer_builder.build(None).await?;
337
338 let arrow_schema = make_test_arrow_schema();
340
341 let batch = RecordBatch::try_new(Arc::new(arrow_schema), vec![
342 Arc::new(Int32Array::from(vec![1, 2, 3])),
343 Arc::new(StringArray::from(vec!["Alice", "Bob", "Charlie"])),
344 ])?;
345
346 writer.write(batch.clone()).await?;
348
349 let data_files = writer.close().await?;
351
352 assert_eq!(
354 data_files.len(),
355 1,
356 "Expected only one data file to be created"
357 );
358
359 check_parquet_data_file(&file_io, &data_files[0], &batch).await;
361
362 Ok(())
363 }
364
365 #[tokio::test]
366 async fn test_rolling_writer_with_rolling() -> Result<()> {
367 let temp_dir = TempDir::new()?;
368 let file_io = FileIO::new_with_fs();
369 let location_gen = DefaultLocationGenerator::with_data_location(
370 temp_dir.path().to_str().unwrap().to_string(),
371 );
372 let file_name_gen =
373 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
374
375 let schema = make_test_schema()?;
377
378 let parquet_writer_builder =
380 ParquetWriterBuilder::new(WriterProperties::builder().build(), Arc::new(schema));
381
382 let rolling_writer_builder = RollingFileWriterBuilder::new(
384 parquet_writer_builder,
385 1024,
386 file_io,
387 location_gen,
388 file_name_gen,
389 );
390
391 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
392
393 let mut writer = data_file_writer_builder.build(None).await?;
395
396 let arrow_schema = make_test_arrow_schema();
398 let arrow_schema_ref = Arc::new(arrow_schema.clone());
399
400 let names = vec![
401 "Alice", "Bob", "Charlie", "Dave", "Eve", "Frank", "Grace", "Heidi", "Ivan", "Judy",
402 "Kelly", "Larry", "Mallory", "Shawn",
403 ];
404
405 let mut rng = rand::rng();
406 let batch_num = 10;
407 let batch_rows = 100;
408 let expected_rows = batch_num * batch_rows;
409
410 for i in 0..batch_num {
411 let int_values: Vec<i32> = (0..batch_rows).map(|row| i * batch_rows + row).collect();
412 let str_values: Vec<&str> = (0..batch_rows)
413 .map(|_| *names.iter().choose(&mut rng).unwrap())
414 .collect();
415
416 let int_array = Arc::new(Int32Array::from(int_values)) as ArrayRef;
417 let str_array = Arc::new(StringArray::from(str_values)) as ArrayRef;
418
419 let batch =
420 RecordBatch::try_new(Arc::clone(&arrow_schema_ref), vec![int_array, str_array])
421 .expect("Failed to create RecordBatch");
422
423 writer.write(batch).await?;
424 }
425
426 let data_files = writer.close().await?;
428
429 assert!(
431 data_files.len() > 4,
432 "Expected at least 4 data files to be created, but got {}",
433 data_files.len()
434 );
435
436 let total_records: u64 = data_files.iter().map(|file| file.record_count).sum();
438 assert_eq!(
439 total_records, expected_rows as u64,
440 "Expected {expected_rows} total records across all files"
441 );
442
443 Ok(())
444 }
445
446 #[tokio::test]
448 async fn test_rolling_writer_encrypted() -> Result<()> {
449 let temp_dir = TempDir::new()?;
450 let file_io = FileIO::new_with_fs();
451 let location_gen = DefaultLocationGenerator::with_data_location(
452 temp_dir.path().to_str().unwrap().to_string(),
453 );
454 let file_name_gen =
455 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
456
457 let schema = make_test_schema()?;
458
459 let raw_properties = HashMap::from([(
460 TableProperties::PROPERTY_ENCRYPTION_KEY_ID.to_string(),
461 "test-key".to_string(),
462 )]);
463 let table_properties = TableProperties::new(&raw_properties);
464 let parquet_writer_builder =
465 ParquetWriterBuilder::from_table_properties(&table_properties, Arc::new(schema))?
466 .with_encryption_manager(make_encryption_manager("test-key"));
467
468 let rolling_writer_builder = RollingFileWriterBuilder::new(
470 parquet_writer_builder,
471 1024,
472 file_io.clone(),
473 location_gen,
474 file_name_gen,
475 );
476
477 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
478
479 let mut writer = data_file_writer_builder.build(None).await?;
481
482 let arrow_schema = make_test_arrow_schema();
484 let arrow_schema_ref = Arc::new(arrow_schema.clone());
485
486 let names = vec![
487 "Alice", "Bob", "Charlie", "Dave", "Eve", "Frank", "Grace", "Heidi", "Ivan", "Judy",
488 "Kelly", "Larry", "Mallory", "Shawn",
489 ];
490
491 let mut rng = rand::rng();
492 let batch_num = 10;
493 let batch_rows = 100;
494 let expected_rows = batch_num * batch_rows;
495
496 for i in 0..batch_num {
497 let int_values: Vec<i32> = (0..batch_rows).map(|row| i * batch_rows + row).collect();
498 let str_values: Vec<&str> = (0..batch_rows)
499 .map(|_| *names.iter().choose(&mut rng).unwrap())
500 .collect();
501
502 let int_array = Arc::new(Int32Array::from(int_values)) as ArrayRef;
503 let str_array = Arc::new(StringArray::from(str_values)) as ArrayRef;
504
505 let batch =
506 RecordBatch::try_new(Arc::clone(&arrow_schema_ref), vec![int_array, str_array])
507 .expect("Failed to create RecordBatch");
508
509 writer.write(batch).await?;
510 }
511
512 let data_files = writer.close().await?;
513
514 assert!(
515 data_files.len() > 4,
516 "Expected at least 4 data files to be created, but got {}",
517 data_files.len()
518 );
519
520 let total_records: u64 = data_files.iter().map(|file| file.record_count).sum();
522 assert_eq!(
523 total_records, expected_rows as u64,
524 "Expected {expected_rows} total records across all files"
525 );
526
527 let mut distinct_keys = HashSet::new();
528 for data_file in &data_files {
529 let key_metadata = StandardKeyMetadata::decode(
530 data_file
531 .key_metadata()
532 .expect("each rolled file must carry key metadata"),
533 )?;
534 assert!(
535 distinct_keys.insert(key_metadata.encryption_key().as_bytes().to_vec()),
536 "each rolled file must use a distinct DEK"
537 );
538 }
539
540 let data_file = &data_files[0];
541 let key_metadata = StandardKeyMetadata::decode(data_file.key_metadata().unwrap())?;
542 let batches = read_encrypted_parquet(
543 &data_file.file_path,
544 key_metadata.encryption_key().as_bytes(),
545 key_metadata.aad_prefix(),
546 );
547 let read_rows: u64 = batches.iter().map(|b| b.num_rows() as u64).sum();
548 assert_eq!(read_rows, data_file.record_count);
549
550 Ok(())
551 }
552}