datafusion_datasource/write/
demux.rs1use std::borrow::Cow;
22use std::collections::HashMap;
23use std::sync::Arc;
24
25use crate::url::ListingTableUrl;
26use crate::write::FileSinkConfig;
27use datafusion_common::error::Result;
28use datafusion_physical_plan::SendableRecordBatchStream;
29
30use arrow::array::{
31 ArrayAccessor, RecordBatch, StringArray, StructArray, builder::UInt64Builder,
32 cast::AsArray, downcast_dictionary_array,
33};
34use arrow::datatypes::{DataType, Schema};
35use datafusion_common::cast::{
36 as_boolean_array, as_date32_array, as_date64_array, as_float16_array,
37 as_float32_array, as_float64_array, as_int8_array, as_int16_array, as_int32_array,
38 as_int64_array, as_large_string_array, as_string_array, as_string_view_array,
39 as_uint8_array, as_uint16_array, as_uint32_array, as_uint64_array,
40};
41use datafusion_common::{exec_datafusion_err, internal_datafusion_err, not_impl_err};
42use datafusion_common_runtime::SpawnedTask;
43
44use chrono::NaiveDate;
45use datafusion_execution::TaskContext;
46use futures::StreamExt;
47use object_store::path::Path;
48use rand::distr::SampleString;
49use tokio::sync::mpsc::{self, Receiver, Sender, UnboundedReceiver, UnboundedSender};
50
51type RecordBatchReceiver = Receiver<RecordBatch>;
52pub type DemuxedStreamReceiver = UnboundedReceiver<(Path, RecordBatchReceiver)>;
53
54pub(crate) fn start_demuxer_task(
100 config: &FileSinkConfig,
101 data: SendableRecordBatchStream,
102 context: &Arc<TaskContext>,
103) -> (SpawnedTask<Result<()>>, DemuxedStreamReceiver) {
104 let (tx, rx) = mpsc::unbounded_channel();
105 let context = Arc::clone(context);
106 let file_extension = config.file_extension.clone();
107 let base_output_path = config.table_paths[0].clone();
108 let task = if config.table_partition_cols.is_empty() {
109 let single_file_output = config
110 .file_output_mode
111 .single_file_output(&base_output_path);
112 SpawnedTask::spawn(async move {
113 row_count_demuxer(
114 tx,
115 data,
116 context,
117 base_output_path,
118 file_extension,
119 single_file_output,
120 )
121 .await
122 })
123 } else {
124 let partition_by = config.table_partition_cols.clone();
127 let keep_partition_by_columns = config.keep_partition_by_columns;
128 SpawnedTask::spawn(async move {
129 hive_style_partitions_demuxer(
130 tx,
131 data,
132 context,
133 partition_by,
134 base_output_path,
135 file_extension,
136 keep_partition_by_columns,
137 )
138 .await
139 })
140 };
141
142 (task, rx)
143}
144
145async fn row_count_demuxer(
147 mut tx: UnboundedSender<(Path, Receiver<RecordBatch>)>,
148 mut input: SendableRecordBatchStream,
149 context: Arc<TaskContext>,
150 base_output_path: ListingTableUrl,
151 file_extension: String,
152 single_file_output: bool,
153) -> Result<()> {
154 let exec_options = &context.session_config().options().execution;
155
156 let max_rows_per_file = exec_options.soft_max_rows_per_output_file.get();
157 let max_buffered_batches = exec_options.max_buffered_batches_per_output_file.get();
158 let minimum_parallel_files = exec_options.minimum_parallel_output_files.get();
159 let mut part_idx = 0;
160 let write_id = rand::distr::Alphanumeric.sample_string(&mut rand::rng(), 16);
161
162 let mut open_file_streams = Vec::with_capacity(minimum_parallel_files);
163
164 let mut next_send_steam = 0;
165 let mut row_counts = Vec::with_capacity(minimum_parallel_files);
166
167 let minimum_parallel_files = if single_file_output {
169 1
170 } else {
171 minimum_parallel_files
172 };
173
174 let max_rows_per_file = if single_file_output {
175 usize::MAX
176 } else {
177 max_rows_per_file
178 };
179
180 if single_file_output {
181 open_file_streams.push(create_new_file_stream(
183 &base_output_path,
184 &write_id,
185 part_idx,
186 &file_extension,
187 single_file_output,
188 max_buffered_batches,
189 &mut tx,
190 )?);
191 row_counts.push(0);
192 part_idx += 1;
193 }
194
195 let schema = input.schema();
196 let mut is_batch_received = false;
197
198 while let Some(rb) = input.next().await.transpose()? {
199 is_batch_received = true;
200 if open_file_streams.len() < minimum_parallel_files {
202 open_file_streams.push(create_new_file_stream(
203 &base_output_path,
204 &write_id,
205 part_idx,
206 &file_extension,
207 single_file_output,
208 max_buffered_batches,
209 &mut tx,
210 )?);
211 row_counts.push(0);
212 part_idx += 1;
213 } else if row_counts[next_send_steam] >= max_rows_per_file {
214 row_counts[next_send_steam] = 0;
215 open_file_streams[next_send_steam] = create_new_file_stream(
216 &base_output_path,
217 &write_id,
218 part_idx,
219 &file_extension,
220 single_file_output,
221 max_buffered_batches,
222 &mut tx,
223 )?;
224 part_idx += 1;
225 }
226 row_counts[next_send_steam] += rb.num_rows();
227 open_file_streams[next_send_steam]
228 .send(rb)
229 .await
230 .map_err(|_| {
231 exec_datafusion_err!("Error sending RecordBatch to file stream!")
232 })?;
233
234 next_send_steam = (next_send_steam + 1) % minimum_parallel_files;
235 }
236
237 if single_file_output && !is_batch_received {
239 open_file_streams
240 .first_mut()
241 .ok_or_else(|| internal_datafusion_err!("Expected a single output file"))?
242 .send(RecordBatch::new_empty(schema))
243 .await
244 .map_err(|_| {
245 exec_datafusion_err!("Error sending empty RecordBatch to file stream!")
246 })?;
247 }
248
249 Ok(())
250}
251
252fn generate_file_path(
254 base_output_path: &ListingTableUrl,
255 write_id: &str,
256 part_idx: usize,
257 file_extension: &str,
258 single_file_output: bool,
259) -> Path {
260 if !single_file_output {
261 base_output_path
262 .prefix()
263 .clone()
264 .join(format!("{write_id}_{part_idx}.{file_extension}"))
265 } else {
266 base_output_path.prefix().to_owned()
267 }
268}
269
270fn create_new_file_stream(
272 base_output_path: &ListingTableUrl,
273 write_id: &str,
274 part_idx: usize,
275 file_extension: &str,
276 single_file_output: bool,
277 max_buffered_batches: usize,
278 tx: &mut UnboundedSender<(Path, Receiver<RecordBatch>)>,
279) -> Result<Sender<RecordBatch>> {
280 let file_path = generate_file_path(
281 base_output_path,
282 write_id,
283 part_idx,
284 file_extension,
285 single_file_output,
286 );
287 let (tx_file, rx_file) = mpsc::channel(max_buffered_batches / 2);
288 tx.send((file_path, rx_file))
289 .map_err(|_| exec_datafusion_err!("Error sending RecordBatch to file stream!"))?;
290 Ok(tx_file)
291}
292
293async fn hive_style_partitions_demuxer(
297 tx: UnboundedSender<(Path, Receiver<RecordBatch>)>,
298 mut input: SendableRecordBatchStream,
299 context: Arc<TaskContext>,
300 partition_by: Vec<(String, DataType)>,
301 base_output_path: ListingTableUrl,
302 file_extension: String,
303 keep_partition_by_columns: bool,
304) -> Result<()> {
305 let write_id = rand::distr::Alphanumeric.sample_string(&mut rand::rng(), 16);
306
307 let exec_options = &context.session_config().options().execution;
308 let max_buffered_recordbatches =
309 exec_options.max_buffered_batches_per_output_file.get();
310
311 let mut value_map: HashMap<Vec<String>, Sender<RecordBatch>> = HashMap::new();
313
314 while let Some(rb) = input.next().await.transpose()? {
315 let all_partition_values = compute_partition_keys_by_row(&rb, &partition_by)?;
317
318 let take_map = compute_take_arrays(&rb, &all_partition_values);
320
321 for (part_key, mut builder) in take_map.into_iter() {
323 let take_indices = builder.finish();
326 let struct_array: StructArray = rb.clone().into();
327 let parted_batch = RecordBatch::from(
328 arrow::compute::take(&struct_array, &take_indices, None)?.as_struct(),
329 );
330
331 let part_tx = match value_map.get_mut(&part_key) {
333 Some(part_tx) => part_tx,
334 None => {
335 let (part_tx, part_rx) =
337 mpsc::channel::<RecordBatch>(max_buffered_recordbatches);
338 let file_path = compute_hive_style_file_path(
339 &part_key,
340 &partition_by,
341 &write_id,
342 &file_extension,
343 &base_output_path,
344 );
345
346 tx.send((file_path, part_rx)).map_err(|_| {
347 exec_datafusion_err!("Error sending new file stream!")
348 })?;
349
350 value_map.insert(part_key.clone(), part_tx);
351 value_map.get_mut(&part_key).ok_or_else(|| {
352 exec_datafusion_err!("Key must exist since it was just inserted!")
353 })?
354 }
355 };
356
357 let final_batch_to_send = if keep_partition_by_columns {
358 parted_batch
359 } else {
360 remove_partition_by_columns(&parted_batch, &partition_by)?
361 };
362
363 part_tx.send(final_batch_to_send).await.map_err(|_| {
365 internal_datafusion_err!("Unexpected error sending parted batch!")
366 })?;
367 }
368 }
369
370 Ok(())
371}
372
373fn compute_partition_keys_by_row<'a>(
374 rb: &'a RecordBatch,
375 partition_by: &'a [(String, DataType)],
376) -> Result<Vec<Vec<Cow<'a, str>>>> {
377 let mut all_partition_values = vec![];
378
379 const EPOCH_DAYS_FROM_CE: i32 = 719_163;
380
381 let schema = rb.schema();
387 for (col, _) in partition_by.iter() {
388 let mut partition_values = vec![];
389
390 let dtype = schema.field_with_name(col)?.data_type();
391 let col_array = rb.column_by_name(col).ok_or(exec_datafusion_err!(
392 "PartitionBy Column {} does not exist in source data! Got schema {schema}.",
393 col
394 ))?;
395
396 match dtype {
397 DataType::Utf8 => {
398 let array = as_string_array(col_array)?;
399 for i in 0..rb.num_rows() {
400 partition_values.push(Cow::from(array.value(i)));
401 }
402 }
403 DataType::LargeUtf8 => {
404 let array = as_large_string_array(col_array)?;
405 for i in 0..rb.num_rows() {
406 partition_values.push(Cow::from(array.value(i)));
407 }
408 }
409 DataType::Utf8View => {
410 let array = as_string_view_array(col_array)?;
411 for i in 0..rb.num_rows() {
412 partition_values.push(Cow::from(array.value(i)));
413 }
414 }
415 DataType::Boolean => {
416 let array = as_boolean_array(col_array)?;
417 for i in 0..rb.num_rows() {
418 partition_values.push(Cow::from(array.value(i).to_string()));
419 }
420 }
421 DataType::Date32 => {
422 let array = as_date32_array(col_array)?;
423 let format = "%Y-%m-%d";
425 for i in 0..rb.num_rows() {
426 let date = NaiveDate::from_num_days_from_ce_opt(
427 EPOCH_DAYS_FROM_CE + array.value(i),
428 )
429 .unwrap()
430 .format(format)
431 .to_string();
432 partition_values.push(Cow::from(date));
433 }
434 }
435 DataType::Date64 => {
436 let array = as_date64_array(col_array)?;
437 let format = "%Y-%m-%d";
439 for i in 0..rb.num_rows() {
440 let date = NaiveDate::from_num_days_from_ce_opt(
441 EPOCH_DAYS_FROM_CE + (array.value(i) / 86_400_000) as i32,
442 )
443 .unwrap()
444 .format(format)
445 .to_string();
446 partition_values.push(Cow::from(date));
447 }
448 }
449 DataType::Int8 => {
450 let array = as_int8_array(col_array)?;
451 for i in 0..rb.num_rows() {
452 partition_values.push(Cow::from(array.value(i).to_string()));
453 }
454 }
455 DataType::Int16 => {
456 let array = as_int16_array(col_array)?;
457 for i in 0..rb.num_rows() {
458 partition_values.push(Cow::from(array.value(i).to_string()));
459 }
460 }
461 DataType::Int32 => {
462 let array = as_int32_array(col_array)?;
463 for i in 0..rb.num_rows() {
464 partition_values.push(Cow::from(array.value(i).to_string()));
465 }
466 }
467 DataType::Int64 => {
468 let array = as_int64_array(col_array)?;
469 for i in 0..rb.num_rows() {
470 partition_values.push(Cow::from(array.value(i).to_string()));
471 }
472 }
473 DataType::UInt8 => {
474 let array = as_uint8_array(col_array)?;
475 for i in 0..rb.num_rows() {
476 partition_values.push(Cow::from(array.value(i).to_string()));
477 }
478 }
479 DataType::UInt16 => {
480 let array = as_uint16_array(col_array)?;
481 for i in 0..rb.num_rows() {
482 partition_values.push(Cow::from(array.value(i).to_string()));
483 }
484 }
485 DataType::UInt32 => {
486 let array = as_uint32_array(col_array)?;
487 for i in 0..rb.num_rows() {
488 partition_values.push(Cow::from(array.value(i).to_string()));
489 }
490 }
491 DataType::UInt64 => {
492 let array = as_uint64_array(col_array)?;
493 for i in 0..rb.num_rows() {
494 partition_values.push(Cow::from(array.value(i).to_string()));
495 }
496 }
497 DataType::Float16 => {
498 let array = as_float16_array(col_array)?;
499 for i in 0..rb.num_rows() {
500 partition_values.push(Cow::from(array.value(i).to_string()));
501 }
502 }
503 DataType::Float32 => {
504 let array = as_float32_array(col_array)?;
505 for i in 0..rb.num_rows() {
506 partition_values.push(Cow::from(array.value(i).to_string()));
507 }
508 }
509 DataType::Float64 => {
510 let array = as_float64_array(col_array)?;
511 for i in 0..rb.num_rows() {
512 partition_values.push(Cow::from(array.value(i).to_string()));
513 }
514 }
515 DataType::Dictionary(_, _) => {
516 downcast_dictionary_array!(
517 col_array => {
518 let array = col_array.downcast_dict::<StringArray>()
519 .ok_or(exec_datafusion_err!("it is not yet supported to write to hive partitions with datatype {}",
520 dtype))?;
521
522 for i in 0..rb.num_rows() {
523 partition_values.push(Cow::from(array.value(i)));
524 }
525 },
526 _ => unreachable!(),
527 )
528 }
529 _ => {
530 return not_impl_err!(
531 "it is not yet supported to write to hive partitions with datatype {}",
532 dtype
533 );
534 }
535 }
536
537 all_partition_values.push(partition_values);
538 }
539
540 Ok(all_partition_values)
541}
542
543fn compute_take_arrays(
544 rb: &RecordBatch,
545 all_partition_values: &[Vec<Cow<str>>],
546) -> HashMap<Vec<String>, UInt64Builder> {
547 let mut take_map = HashMap::new();
548 for i in 0..rb.num_rows() {
549 let mut part_key = vec![];
550 for vals in all_partition_values.iter() {
551 part_key.push(vals[i].clone().into());
552 }
553 let builder = take_map.entry(part_key).or_insert_with(UInt64Builder::new);
554 builder.append_value(i as u64);
555 }
556 take_map
557}
558
559fn remove_partition_by_columns(
560 parted_batch: &RecordBatch,
561 partition_by: &[(String, DataType)],
562) -> Result<RecordBatch> {
563 let partition_names: Vec<_> = partition_by.iter().map(|(s, _)| s).collect();
564 let (non_part_cols, non_part_fields): (Vec<_>, Vec<_>) = parted_batch
565 .columns()
566 .iter()
567 .zip(parted_batch.schema().fields())
568 .filter_map(|(a, f)| {
569 if !partition_names.contains(&f.name()) {
570 Some((Arc::clone(a), (**f).clone()))
571 } else {
572 None
573 }
574 })
575 .unzip();
576
577 let non_part_schema = Schema::new(non_part_fields);
578 let final_batch_to_send =
579 RecordBatch::try_new(Arc::new(non_part_schema), non_part_cols)?;
580
581 Ok(final_batch_to_send)
582}
583
584fn compute_hive_style_file_path(
585 part_key: &[String],
586 partition_by: &[(String, DataType)],
587 write_id: &str,
588 file_extension: &str,
589 base_output_path: &ListingTableUrl,
590) -> Path {
591 let mut file_path = base_output_path.prefix().clone();
592 for j in 0..part_key.len() {
593 file_path = file_path.join(format!("{}={}", partition_by[j].0, part_key[j]));
594 }
595
596 file_path.join(format!("{write_id}.{file_extension}"))
597}