Skip to main content

iceberg/writer/partitioning/
clustered_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 provides the `ClusteredWriter` implementation.
19
20use 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
30/// A writer that writes data to a single partition at a time.
31///
32/// This writer expects input data to be sorted by partition key. It maintains only one
33/// active writer at a time, making it memory efficient for sorted data.
34///
35/// # Type Parameters
36///
37/// * `B` - The inner writer builder type
38/// * `I` - Input type (defaults to `RecordBatch`)
39/// * `O` - Output collection type (defaults to `Vec<DataFile>`)
40pub 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    /// Create a new `ClusteredWriter`.
62    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    /// Closes the current writer if it exists, flushes the written data to output, and record closed partition.
74    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            // Add the current partition to the set of closed partitions
79            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        // Check if this partition has been closed already
100        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        // Check if we need to switch to a new partition
110        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            // Create a new writer for the new partition
119            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        // do write
128        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        // Collect all output items into the output collection type
139        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        // Create schema with partition field
176        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        // Create partition spec and key
189        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        // Create writer builder
195        let parquet_writer_builder =
196            ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
197
198        // Create rolling file writer builder
199        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        // Create data file writer builder
207        let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
208
209        // Create clustered writer
210        let mut writer = ClusteredWriter::new(data_file_writer_builder);
211
212        // Create test data with proper field ID metadata
213        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        // Write data to the same partition (this should work)
241        writer.write(partition_key.clone(), batch1).await?;
242        writer.write(partition_key.clone(), batch2).await?;
243
244        // Close writer and get data files
245        let data_files = writer.close().await?;
246
247        // Verify at least one file was created
248        assert!(
249            !data_files.is_empty(),
250            "Expected at least one data file to be created"
251        );
252
253        // Verify that all data files have the correct partition value
254        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        // Create schema with partition field
272        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        // Create partition spec
285        let partition_spec = crate::spec::PartitionSpec::builder(schema.clone()).build()?;
286
287        // Create partition keys for different regions (in sorted order)
288        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        // Create writer builder
310        let parquet_writer_builder =
311            ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
312
313        // Create rolling file writer builder
314        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        // Create data file writer builder
322        let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
323
324        // Create clustered writer
325        let mut writer = ClusteredWriter::new(data_file_writer_builder);
326
327        // Create test data with proper field ID metadata
328        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        // Create batches for different partitions (in sorted order)
344        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        // Write data in sorted partition order (this should work)
363        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        // Close writer and get data files
368        let data_files = writer.close().await?;
369
370        // Verify files were created for all partitions
371        assert!(
372            data_files.len() >= 3,
373            "Expected at least 3 data files (one per partition), got {}",
374            data_files.len()
375        );
376
377        // Verify that we have files for each partition
378        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        // Create schema with partition field
410        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        // Create partition spec
423        let partition_spec = crate::spec::PartitionSpec::builder(schema.clone()).build()?;
424
425        // Create partition keys for different regions
426        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        // Create writer builder
441        let parquet_writer_builder =
442            ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
443
444        // Create rolling file writer builder
445        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        // Create data file writer builder
453        let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
454
455        // Create clustered writer
456        let mut writer = ClusteredWriter::new(data_file_writer_builder);
457
458        // Create test data with proper field ID metadata
459        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        // Create batches for different partitions
475        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        // Write data to US partition first
494        writer.write(partition_key_us.clone(), batch_us).await?;
495
496        // Write data to EU partition (this closes US partition)
497        writer.write(partition_key_eu.clone(), batch_eu).await?;
498
499        // Try to write to US partition again - this should fail because data is not sorted
500        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        // -0.0 and 0.0 are different partition values, as in iceberg-java, so 0.0 right
558        // after -0.0 starts a new data file.
559        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        // 0.0 is not the closed -0.0 partition, so this input is still sorted.
575        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}