1use std::collections::HashSet;
21use std::marker::PhantomData;
22
23use async_trait::async_trait;
24
25use crate::spec::{PartitionKey, Struct};
26use crate::writer::partitioning::PartitioningWriter;
27use crate::writer::{DefaultInput, DefaultOutput, IcebergWriter, IcebergWriterBuilder};
28use crate::{Error, ErrorKind, Result};
29
30pub struct ClusteredWriter<B, I = DefaultInput, O = DefaultOutput>
41where
42 B: IcebergWriterBuilder<I, O>,
43 O: IntoIterator + FromIterator<<O as IntoIterator>::Item>,
44 <O as IntoIterator>::Item: Clone,
45{
46 inner_builder: B,
47 current_writer: Option<B::R>,
48 current_partition: Option<Struct>,
49 closed_partitions: HashSet<Struct>,
50 output: Vec<<O as IntoIterator>::Item>,
51 _phantom: PhantomData<I>,
52}
53
54impl<B, I, O> ClusteredWriter<B, I, O>
55where
56 B: IcebergWriterBuilder<I, O>,
57 I: Send + 'static,
58 O: IntoIterator + FromIterator<<O as IntoIterator>::Item>,
59 <O as IntoIterator>::Item: Send + Clone,
60{
61 pub fn new(inner_builder: B) -> Self {
63 Self {
64 inner_builder,
65 current_writer: None,
66 current_partition: None,
67 closed_partitions: HashSet::new(),
68 output: Vec::new(),
69 _phantom: PhantomData,
70 }
71 }
72
73 async fn close_current_writer(&mut self) -> Result<()> {
75 if let Some(mut writer) = self.current_writer.take() {
76 self.output.extend(writer.close().await?);
77
78 if let Some(current_partition) = self.current_partition.take() {
80 self.closed_partitions.insert(current_partition);
81 }
82 }
83
84 Ok(())
85 }
86}
87
88#[async_trait]
89impl<B, I, O> PartitioningWriter<I, O> for ClusteredWriter<B, I, O>
90where
91 B: IcebergWriterBuilder<I, O>,
92 I: Send + 'static,
93 O: IntoIterator + FromIterator<<O as IntoIterator>::Item> + Send + 'static,
94 <O as IntoIterator>::Item: Send + Clone,
95{
96 async fn write(&mut self, partition_key: PartitionKey, input: I) -> Result<()> {
97 let partition_value = partition_key.data();
98
99 if self.closed_partitions.contains(partition_value) {
101 return Err(Error::new(
102 ErrorKind::Unexpected,
103 format!(
104 "The input is not sorted! Cannot write to partition that was previously closed: {partition_key:?}"
105 ),
106 ));
107 }
108
109 let need_new_writer = match &self.current_partition {
111 Some(current) => current != partition_value,
112 None => true,
113 };
114
115 if need_new_writer {
116 self.close_current_writer().await?;
117
118 self.current_writer = Some(
120 self.inner_builder
121 .build(Some(partition_key.clone()))
122 .await?,
123 );
124 self.current_partition = Some(partition_value.clone());
125 }
126
127 self.current_writer
129 .as_mut()
130 .expect("Writer should be initialized")
131 .write(input)
132 .await
133 }
134
135 async fn close(mut self) -> Result<O> {
136 self.close_current_writer().await?;
137
138 Ok(O::from_iter(self.output))
140 }
141}
142
143#[cfg(test)]
144mod tests {
145 use std::collections::HashMap;
146 use std::sync::Arc;
147
148 use arrow_array::{Float64Array, Int32Array, RecordBatch, StringArray};
149 use arrow_schema::{DataType, Field, Schema};
150 use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
151 use parquet::file::properties::WriterProperties;
152 use tempfile::TempDir;
153
154 use super::*;
155 use crate::arrow::schema_to_arrow_schema;
156 use crate::io::FileIO;
157 use crate::spec::{DataFileFormat, NestedField, PrimitiveType, Type};
158 use crate::writer::base_writer::data_file_writer::DataFileWriterBuilder;
159 use crate::writer::file_writer::ParquetWriterBuilder;
160 use crate::writer::file_writer::location_generator::{
161 DefaultFileNameGenerator, DefaultLocationGenerator,
162 };
163 use crate::writer::file_writer::rolling_writer::RollingFileWriterBuilder;
164
165 #[tokio::test]
166 async fn test_clustered_writer_single_partition() -> Result<()> {
167 let temp_dir = TempDir::new()?;
168 let file_io = FileIO::new_with_fs();
169 let location_gen = DefaultLocationGenerator::with_data_location(
170 temp_dir.path().to_str().unwrap().to_string(),
171 );
172 let file_name_gen =
173 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
174
175 let schema = Arc::new(
177 crate::spec::Schema::builder()
178 .with_schema_id(1)
179 .with_fields(vec![
180 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
181 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
182 NestedField::required(3, "region", Type::Primitive(PrimitiveType::String))
183 .into(),
184 ])
185 .build()?,
186 );
187
188 let partition_spec = crate::spec::PartitionSpec::builder(schema.clone()).build()?;
190 let partition_value = Struct::from_iter([Some(crate::spec::Literal::string("US"))]);
191 let partition_key =
192 PartitionKey::new(partition_spec, schema.clone(), partition_value.clone());
193
194 let parquet_writer_builder =
196 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
197
198 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
200 parquet_writer_builder,
201 file_io.clone(),
202 location_gen,
203 file_name_gen,
204 );
205
206 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
208
209 let mut writer = ClusteredWriter::new(data_file_writer_builder);
211
212 let arrow_schema = Schema::new(vec![
214 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
215 PARQUET_FIELD_ID_META_KEY.to_string(),
216 1.to_string(),
217 )])),
218 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
219 PARQUET_FIELD_ID_META_KEY.to_string(),
220 2.to_string(),
221 )])),
222 Field::new("region", DataType::Utf8, false).with_metadata(HashMap::from([(
223 PARQUET_FIELD_ID_META_KEY.to_string(),
224 3.to_string(),
225 )])),
226 ]);
227
228 let batch1 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
229 Arc::new(Int32Array::from(vec![1, 2])),
230 Arc::new(StringArray::from(vec!["Alice", "Bob"])),
231 Arc::new(StringArray::from(vec!["US", "US"])),
232 ])?;
233
234 let batch2 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
235 Arc::new(Int32Array::from(vec![3, 4])),
236 Arc::new(StringArray::from(vec!["Charlie", "Dave"])),
237 Arc::new(StringArray::from(vec!["US", "US"])),
238 ])?;
239
240 writer.write(partition_key.clone(), batch1).await?;
242 writer.write(partition_key.clone(), batch2).await?;
243
244 let data_files = writer.close().await?;
246
247 assert!(
249 !data_files.is_empty(),
250 "Expected at least one data file to be created"
251 );
252
253 for data_file in &data_files {
255 assert_eq!(data_file.partition, partition_value);
256 }
257
258 Ok(())
259 }
260
261 #[tokio::test]
262 async fn test_clustered_writer_sorted_partitions() -> Result<()> {
263 let temp_dir = TempDir::new()?;
264 let file_io = FileIO::new_with_fs();
265 let location_gen = DefaultLocationGenerator::with_data_location(
266 temp_dir.path().to_str().unwrap().to_string(),
267 );
268 let file_name_gen =
269 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
270
271 let schema = Arc::new(
273 crate::spec::Schema::builder()
274 .with_schema_id(1)
275 .with_fields(vec![
276 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
277 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
278 NestedField::required(3, "region", Type::Primitive(PrimitiveType::String))
279 .into(),
280 ])
281 .build()?,
282 );
283
284 let partition_spec = crate::spec::PartitionSpec::builder(schema.clone()).build()?;
286
287 let partition_value_asia = Struct::from_iter([Some(crate::spec::Literal::string("ASIA"))]);
289 let partition_key_asia = PartitionKey::new(
290 partition_spec.clone(),
291 schema.clone(),
292 partition_value_asia.clone(),
293 );
294
295 let partition_value_eu = Struct::from_iter([Some(crate::spec::Literal::string("EU"))]);
296 let partition_key_eu = PartitionKey::new(
297 partition_spec.clone(),
298 schema.clone(),
299 partition_value_eu.clone(),
300 );
301
302 let partition_value_us = Struct::from_iter([Some(crate::spec::Literal::string("US"))]);
303 let partition_key_us = PartitionKey::new(
304 partition_spec.clone(),
305 schema.clone(),
306 partition_value_us.clone(),
307 );
308
309 let parquet_writer_builder =
311 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
312
313 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
315 parquet_writer_builder,
316 file_io.clone(),
317 location_gen,
318 file_name_gen,
319 );
320
321 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
323
324 let mut writer = ClusteredWriter::new(data_file_writer_builder);
326
327 let arrow_schema = Schema::new(vec![
329 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
330 PARQUET_FIELD_ID_META_KEY.to_string(),
331 1.to_string(),
332 )])),
333 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
334 PARQUET_FIELD_ID_META_KEY.to_string(),
335 2.to_string(),
336 )])),
337 Field::new("region", DataType::Utf8, false).with_metadata(HashMap::from([(
338 PARQUET_FIELD_ID_META_KEY.to_string(),
339 3.to_string(),
340 )])),
341 ]);
342
343 let batch_asia = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
345 Arc::new(Int32Array::from(vec![1, 2])),
346 Arc::new(StringArray::from(vec!["Alice", "Bob"])),
347 Arc::new(StringArray::from(vec!["ASIA", "ASIA"])),
348 ])?;
349
350 let batch_eu = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
351 Arc::new(Int32Array::from(vec![3, 4])),
352 Arc::new(StringArray::from(vec!["Charlie", "Dave"])),
353 Arc::new(StringArray::from(vec!["EU", "EU"])),
354 ])?;
355
356 let batch_us = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
357 Arc::new(Int32Array::from(vec![5, 6])),
358 Arc::new(StringArray::from(vec!["Eve", "Frank"])),
359 Arc::new(StringArray::from(vec!["US", "US"])),
360 ])?;
361
362 writer.write(partition_key_asia.clone(), batch_asia).await?;
364 writer.write(partition_key_eu.clone(), batch_eu).await?;
365 writer.write(partition_key_us.clone(), batch_us).await?;
366
367 let data_files = writer.close().await?;
369
370 assert!(
372 data_files.len() >= 3,
373 "Expected at least 3 data files (one per partition), got {}",
374 data_files.len()
375 );
376
377 let mut partitions_found = HashSet::new();
379 for data_file in &data_files {
380 partitions_found.insert(data_file.partition.clone());
381 }
382
383 assert!(
384 partitions_found.contains(&partition_value_asia),
385 "Missing ASIA partition"
386 );
387 assert!(
388 partitions_found.contains(&partition_value_eu),
389 "Missing EU partition"
390 );
391 assert!(
392 partitions_found.contains(&partition_value_us),
393 "Missing US partition"
394 );
395
396 Ok(())
397 }
398
399 #[tokio::test]
400 async fn test_clustered_writer_unsorted_partitions_error() -> Result<()> {
401 let temp_dir = TempDir::new()?;
402 let file_io = FileIO::new_with_fs();
403 let location_gen = DefaultLocationGenerator::with_data_location(
404 temp_dir.path().to_str().unwrap().to_string(),
405 );
406 let file_name_gen =
407 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
408
409 let schema = Arc::new(
411 crate::spec::Schema::builder()
412 .with_schema_id(1)
413 .with_fields(vec![
414 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
415 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
416 NestedField::required(3, "region", Type::Primitive(PrimitiveType::String))
417 .into(),
418 ])
419 .build()?,
420 );
421
422 let partition_spec = crate::spec::PartitionSpec::builder(schema.clone()).build()?;
424
425 let partition_value_us = Struct::from_iter([Some(crate::spec::Literal::string("US"))]);
427 let partition_key_us = PartitionKey::new(
428 partition_spec.clone(),
429 schema.clone(),
430 partition_value_us.clone(),
431 );
432
433 let partition_value_eu = Struct::from_iter([Some(crate::spec::Literal::string("EU"))]);
434 let partition_key_eu = PartitionKey::new(
435 partition_spec.clone(),
436 schema.clone(),
437 partition_value_eu.clone(),
438 );
439
440 let parquet_writer_builder =
442 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
443
444 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
446 parquet_writer_builder,
447 file_io.clone(),
448 location_gen,
449 file_name_gen,
450 );
451
452 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
454
455 let mut writer = ClusteredWriter::new(data_file_writer_builder);
457
458 let arrow_schema = Schema::new(vec![
460 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
461 PARQUET_FIELD_ID_META_KEY.to_string(),
462 1.to_string(),
463 )])),
464 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
465 PARQUET_FIELD_ID_META_KEY.to_string(),
466 2.to_string(),
467 )])),
468 Field::new("region", DataType::Utf8, false).with_metadata(HashMap::from([(
469 PARQUET_FIELD_ID_META_KEY.to_string(),
470 3.to_string(),
471 )])),
472 ]);
473
474 let batch_us = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
476 Arc::new(Int32Array::from(vec![1, 2])),
477 Arc::new(StringArray::from(vec!["Alice", "Bob"])),
478 Arc::new(StringArray::from(vec!["US", "US"])),
479 ])?;
480
481 let batch_eu = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
482 Arc::new(Int32Array::from(vec![3, 4])),
483 Arc::new(StringArray::from(vec!["Charlie", "Dave"])),
484 Arc::new(StringArray::from(vec!["EU", "EU"])),
485 ])?;
486
487 let batch_us2 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
488 Arc::new(Int32Array::from(vec![5])),
489 Arc::new(StringArray::from(vec!["Eve"])),
490 Arc::new(StringArray::from(vec!["US"])),
491 ])?;
492
493 writer.write(partition_key_us.clone(), batch_us).await?;
495
496 writer.write(partition_key_eu.clone(), batch_eu).await?;
498
499 let result = writer.write(partition_key_us.clone(), batch_us2).await;
501
502 assert!(result.is_err(), "Expected error when writing unsorted data");
503
504 let error = result.unwrap_err();
505 assert!(
506 error.to_string().contains("The input is not sorted"),
507 "Expected 'input is not sorted' error, got: {error}"
508 );
509
510 Ok(())
511 }
512
513 #[tokio::test]
514 async fn test_clustered_writer_signed_zero_partitions() -> Result<()> {
515 let temp_dir = TempDir::new()?;
516 let file_io = FileIO::new_with_fs();
517 let location_gen = DefaultLocationGenerator::with_data_location(
518 temp_dir.path().to_str().unwrap().to_string(),
519 );
520 let file_name_gen =
521 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
522
523 let schema = Arc::new(
524 crate::spec::Schema::builder()
525 .with_schema_id(1)
526 .with_fields(vec![
527 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
528 NestedField::required(2, "d", Type::Primitive(PrimitiveType::Double)).into(),
529 ])
530 .build()?,
531 );
532 let partition_spec = crate::spec::PartitionSpec::builder(schema.clone())
533 .add_partition_field("d", "d", crate::spec::Transform::Identity)?
534 .build()?;
535
536 let partition_value = |d: f64| Struct::from_iter([Some(crate::spec::Literal::double(d))]);
537 let partition_key =
538 |d: f64| PartitionKey::new(partition_spec.clone(), schema.clone(), partition_value(d));
539
540 let arrow_schema = Arc::new(schema_to_arrow_schema(&schema)?);
541 let batch = |id: i32, d: f64| {
542 RecordBatch::try_new(arrow_schema.clone(), vec![
543 Arc::new(Int32Array::from(vec![id])),
544 Arc::new(Float64Array::from(vec![d])),
545 ])
546 };
547
548 let parquet_writer_builder =
549 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
550 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
551 parquet_writer_builder,
552 file_io.clone(),
553 location_gen,
554 file_name_gen,
555 );
556
557 let mut writer =
560 ClusteredWriter::new(DataFileWriterBuilder::new(rolling_writer_builder.clone()));
561 writer.write(partition_key(-0.0), batch(1, -0.0)?).await?;
562 writer.write(partition_key(0.0), batch(2, 0.0)?).await?;
563 let partitions_written: Vec<Struct> = writer
564 .close()
565 .await?
566 .into_iter()
567 .map(|data_file| data_file.partition)
568 .collect();
569 assert_eq!(partitions_written, vec![
570 partition_value(-0.0),
571 partition_value(0.0)
572 ]);
573
574 let mut writer = ClusteredWriter::new(DataFileWriterBuilder::new(rolling_writer_builder));
576 writer.write(partition_key(-0.0), batch(1, -0.0)?).await?;
577 writer.write(partition_key(1.0), batch(3, 1.0)?).await?;
578 writer.write(partition_key(0.0), batch(2, 0.0)?).await?;
579 let partitions_written: Vec<Struct> = writer
580 .close()
581 .await?
582 .into_iter()
583 .map(|data_file| data_file.partition)
584 .collect();
585 assert_eq!(partitions_written, vec![
586 partition_value(-0.0),
587 partition_value(1.0),
588 partition_value(0.0)
589 ]);
590
591 Ok(())
592 }
593}