1use std::collections::HashMap;
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 FanoutWriter<B, I = DefaultInput, O = DefaultOutput>
43where
44 B: IcebergWriterBuilder<I, O>,
45 O: IntoIterator + FromIterator<<O as IntoIterator>::Item>,
46 <O as IntoIterator>::Item: Clone,
47{
48 inner_builder: B,
49 partition_writers: HashMap<Struct, B::R>,
50 output: Vec<<O as IntoIterator>::Item>,
51 _phantom: PhantomData<I>,
52}
53
54impl<B, I, O> FanoutWriter<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 partition_writers: HashMap::new(),
66 output: Vec::new(),
67 _phantom: PhantomData,
68 }
69 }
70
71 async fn get_or_create_writer(&mut self, partition_key: &PartitionKey) -> Result<&mut B::R> {
73 if !self.partition_writers.contains_key(partition_key.data()) {
74 let writer = self
75 .inner_builder
76 .build(Some(partition_key.clone()))
77 .await?;
78 self.partition_writers
79 .insert(partition_key.data().clone(), writer);
80 }
81
82 self.partition_writers
83 .get_mut(partition_key.data())
84 .ok_or_else(|| {
85 Error::new(
86 ErrorKind::Unexpected,
87 "Failed to get partition writer after creation",
88 )
89 })
90 }
91}
92
93#[async_trait]
94impl<B, I, O> PartitioningWriter<I, O> for FanoutWriter<B, I, O>
95where
96 B: IcebergWriterBuilder<I, O>,
97 I: Send + 'static,
98 O: IntoIterator + FromIterator<<O as IntoIterator>::Item> + Send + 'static,
99 <O as IntoIterator>::Item: Send + Clone,
100{
101 async fn write(&mut self, partition_key: PartitionKey, input: I) -> Result<()> {
102 let writer = self.get_or_create_writer(&partition_key).await?;
103 writer.write(input).await
104 }
105
106 async fn close(mut self) -> Result<O> {
107 for (_, mut writer) in self.partition_writers {
109 self.output.extend(writer.close().await?);
110 }
111
112 Ok(O::from_iter(self.output))
114 }
115}
116
117#[cfg(test)]
118mod tests {
119 use std::collections::HashMap;
120 use std::sync::Arc;
121
122 use arrow_array::{Float64Array, Int32Array, RecordBatch, StringArray};
123 use arrow_schema::{DataType, Field, Schema};
124 use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
125 use parquet::file::properties::WriterProperties;
126 use tempfile::TempDir;
127
128 use super::*;
129 use crate::arrow::schema_to_arrow_schema;
130 use crate::io::FileIO;
131 use crate::spec::{
132 DataFileFormat, Literal, NestedField, PartitionKey, PartitionSpec, PrimitiveType, Struct,
133 Transform, Type,
134 };
135 use crate::writer::base_writer::data_file_writer::DataFileWriterBuilder;
136 use crate::writer::file_writer::ParquetWriterBuilder;
137 use crate::writer::file_writer::location_generator::{
138 DefaultFileNameGenerator, DefaultLocationGenerator,
139 };
140 use crate::writer::file_writer::rolling_writer::RollingFileWriterBuilder;
141
142 #[tokio::test]
143 async fn test_fanout_writer_single_partition() -> Result<()> {
144 let temp_dir = TempDir::new()?;
145 let file_io = FileIO::new_with_fs();
146 let location_gen = DefaultLocationGenerator::with_data_location(
147 temp_dir.path().to_str().unwrap().to_string(),
148 );
149 let file_name_gen =
150 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
151
152 let schema = Arc::new(
154 crate::spec::Schema::builder()
155 .with_schema_id(1)
156 .with_fields(vec![
157 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
158 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
159 NestedField::required(3, "region", Type::Primitive(PrimitiveType::String))
160 .into(),
161 ])
162 .build()?,
163 );
164
165 let partition_spec = PartitionSpec::builder(schema.clone()).build()?;
167 let partition_value = Struct::from_iter([Some(Literal::string("US"))]);
168 let partition_key =
169 PartitionKey::new(partition_spec, schema.clone(), partition_value.clone());
170
171 let parquet_writer_builder =
173 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
174
175 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
177 parquet_writer_builder,
178 file_io.clone(),
179 location_gen,
180 file_name_gen,
181 );
182
183 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
185
186 let mut writer = FanoutWriter::new(data_file_writer_builder);
188
189 let arrow_schema = Schema::new(vec![
191 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
192 PARQUET_FIELD_ID_META_KEY.to_string(),
193 1.to_string(),
194 )])),
195 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
196 PARQUET_FIELD_ID_META_KEY.to_string(),
197 2.to_string(),
198 )])),
199 Field::new("region", DataType::Utf8, false).with_metadata(HashMap::from([(
200 PARQUET_FIELD_ID_META_KEY.to_string(),
201 3.to_string(),
202 )])),
203 ]);
204
205 let batch1 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
206 Arc::new(Int32Array::from(vec![1, 2])),
207 Arc::new(StringArray::from(vec!["Alice", "Bob"])),
208 Arc::new(StringArray::from(vec!["US", "US"])),
209 ])?;
210
211 let batch2 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
212 Arc::new(Int32Array::from(vec![3, 4])),
213 Arc::new(StringArray::from(vec!["Charlie", "Dave"])),
214 Arc::new(StringArray::from(vec!["US", "US"])),
215 ])?;
216
217 writer.write(partition_key.clone(), batch1).await?;
219 writer.write(partition_key.clone(), batch2).await?;
220
221 let data_files = writer.close().await?;
223
224 assert!(
226 !data_files.is_empty(),
227 "Expected at least one data file to be created"
228 );
229
230 for data_file in &data_files {
232 assert_eq!(data_file.partition, partition_value);
233 }
234
235 Ok(())
236 }
237
238 #[tokio::test]
239 async fn test_fanout_writer_multiple_partitions() -> Result<()> {
240 let temp_dir = TempDir::new()?;
241 let file_io = FileIO::new_with_fs();
242 let location_gen = DefaultLocationGenerator::with_data_location(
243 temp_dir.path().to_str().unwrap().to_string(),
244 );
245 let file_name_gen =
246 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
247
248 let schema = Arc::new(
250 crate::spec::Schema::builder()
251 .with_schema_id(1)
252 .with_fields(vec![
253 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
254 NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
255 NestedField::required(3, "region", Type::Primitive(PrimitiveType::String))
256 .into(),
257 ])
258 .build()?,
259 );
260
261 let partition_spec = PartitionSpec::builder(schema.clone()).build()?;
263
264 let partition_value_us = Struct::from_iter([Some(Literal::string("US"))]);
266 let partition_key_us = PartitionKey::new(
267 partition_spec.clone(),
268 schema.clone(),
269 partition_value_us.clone(),
270 );
271
272 let partition_value_eu = Struct::from_iter([Some(Literal::string("EU"))]);
273 let partition_key_eu = PartitionKey::new(
274 partition_spec.clone(),
275 schema.clone(),
276 partition_value_eu.clone(),
277 );
278
279 let partition_value_asia = Struct::from_iter([Some(Literal::string("ASIA"))]);
280 let partition_key_asia = PartitionKey::new(
281 partition_spec.clone(),
282 schema.clone(),
283 partition_value_asia.clone(),
284 );
285
286 let parquet_writer_builder =
288 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
289
290 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
292 parquet_writer_builder,
293 file_io.clone(),
294 location_gen,
295 file_name_gen,
296 );
297
298 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
300
301 let mut writer = FanoutWriter::new(data_file_writer_builder);
303
304 let arrow_schema = Schema::new(vec![
306 Field::new("id", DataType::Int32, false).with_metadata(HashMap::from([(
307 PARQUET_FIELD_ID_META_KEY.to_string(),
308 1.to_string(),
309 )])),
310 Field::new("name", DataType::Utf8, false).with_metadata(HashMap::from([(
311 PARQUET_FIELD_ID_META_KEY.to_string(),
312 2.to_string(),
313 )])),
314 Field::new("region", DataType::Utf8, false).with_metadata(HashMap::from([(
315 PARQUET_FIELD_ID_META_KEY.to_string(),
316 3.to_string(),
317 )])),
318 ]);
319
320 let batch_us1 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
322 Arc::new(Int32Array::from(vec![1, 2])),
323 Arc::new(StringArray::from(vec!["Alice", "Bob"])),
324 Arc::new(StringArray::from(vec!["US", "US"])),
325 ])?;
326
327 let batch_eu1 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
328 Arc::new(Int32Array::from(vec![3, 4])),
329 Arc::new(StringArray::from(vec!["Charlie", "Dave"])),
330 Arc::new(StringArray::from(vec!["EU", "EU"])),
331 ])?;
332
333 let batch_us2 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
334 Arc::new(Int32Array::from(vec![5])),
335 Arc::new(StringArray::from(vec!["Eve"])),
336 Arc::new(StringArray::from(vec!["US"])),
337 ])?;
338
339 let batch_asia1 = RecordBatch::try_new(Arc::new(arrow_schema.clone()), vec![
340 Arc::new(Int32Array::from(vec![6, 7])),
341 Arc::new(StringArray::from(vec!["Frank", "Grace"])),
342 Arc::new(StringArray::from(vec!["ASIA", "ASIA"])),
343 ])?;
344
345 writer.write(partition_key_us.clone(), batch_us1).await?;
348 writer.write(partition_key_eu.clone(), batch_eu1).await?;
349 writer.write(partition_key_us.clone(), batch_us2).await?; writer
351 .write(partition_key_asia.clone(), batch_asia1)
352 .await?;
353
354 let data_files = writer.close().await?;
356
357 assert!(
359 data_files.len() >= 3,
360 "Expected at least 3 data files (one per partition), got {}",
361 data_files.len()
362 );
363
364 let mut partitions_found = std::collections::HashSet::new();
366 for data_file in &data_files {
367 partitions_found.insert(data_file.partition.clone());
368 }
369
370 assert!(
371 partitions_found.contains(&partition_value_us),
372 "Missing US partition"
373 );
374 assert!(
375 partitions_found.contains(&partition_value_eu),
376 "Missing EU partition"
377 );
378 assert!(
379 partitions_found.contains(&partition_value_asia),
380 "Missing ASIA partition"
381 );
382
383 Ok(())
384 }
385
386 #[tokio::test]
387 async fn test_fanout_writer_signed_zero_partitions() -> Result<()> {
388 let temp_dir = TempDir::new()?;
389 let file_io = FileIO::new_with_fs();
390 let location_gen = DefaultLocationGenerator::with_data_location(
391 temp_dir.path().to_str().unwrap().to_string(),
392 );
393 let file_name_gen =
394 DefaultFileNameGenerator::new("test".to_string(), None, DataFileFormat::Parquet);
395
396 let schema = Arc::new(
397 crate::spec::Schema::builder()
398 .with_schema_id(1)
399 .with_fields(vec![
400 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
401 NestedField::required(2, "d", Type::Primitive(PrimitiveType::Double)).into(),
402 ])
403 .build()?,
404 );
405 let partition_spec = PartitionSpec::builder(schema.clone())
406 .add_partition_field("d", "d", Transform::Identity)?
407 .build()?;
408
409 let partition_value_neg_zero = Struct::from_iter([Some(Literal::double(-0.0))]);
410 let partition_key_neg_zero = PartitionKey::new(
411 partition_spec.clone(),
412 schema.clone(),
413 partition_value_neg_zero.clone(),
414 );
415
416 let partition_value_pos_zero = Struct::from_iter([Some(Literal::double(0.0))]);
417 let partition_key_pos_zero = PartitionKey::new(
418 partition_spec.clone(),
419 schema.clone(),
420 partition_value_pos_zero.clone(),
421 );
422
423 let parquet_writer_builder =
424 ParquetWriterBuilder::new(WriterProperties::builder().build(), schema.clone());
425 let rolling_writer_builder = RollingFileWriterBuilder::new_with_default_file_size(
426 parquet_writer_builder,
427 file_io.clone(),
428 location_gen,
429 file_name_gen,
430 );
431 let data_file_writer_builder = DataFileWriterBuilder::new(rolling_writer_builder);
432
433 let mut writer = FanoutWriter::new(data_file_writer_builder);
434
435 let arrow_schema = Arc::new(schema_to_arrow_schema(&schema)?);
436 let batch_neg_zero = RecordBatch::try_new(arrow_schema.clone(), vec![
437 Arc::new(Int32Array::from(vec![1])),
438 Arc::new(Float64Array::from(vec![-0.0])),
439 ])?;
440 let batch_pos_zero = RecordBatch::try_new(arrow_schema.clone(), vec![
441 Arc::new(Int32Array::from(vec![2])),
442 Arc::new(Float64Array::from(vec![0.0])),
443 ])?;
444
445 writer.write(partition_key_neg_zero, batch_neg_zero).await?;
446 writer.write(partition_key_pos_zero, batch_pos_zero).await?;
447
448 let data_files = writer.close().await?;
449
450 assert_eq!(data_files.len(), 2);
453 for data_file in &data_files {
454 assert_eq!(data_file.record_count, 1);
455 }
456 let partitions_found: std::collections::HashSet<Struct> = data_files
457 .iter()
458 .map(|data_file| data_file.partition.clone())
459 .collect();
460 assert_eq!(
461 partitions_found,
462 std::collections::HashSet::from([partition_value_neg_zero, partition_value_pos_zero])
463 );
464
465 Ok(())
466 }
467}