use super::expressions::PhysicalSortExpr;
use super::{
DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, SendableRecordBatchStream,
Statistics,
};
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use arrow_array::{ArrayRef, UInt64Array};
use arrow_schema::{DataType, Field, Schema};
use async_trait::async_trait;
use core::fmt;
use datafusion_common::Result;
use datafusion_physical_expr::PhysicalSortRequirement;
use futures::StreamExt;
use std::any::Any;
use std::fmt::Debug;
use std::sync::Arc;
use crate::physical_plan::stream::RecordBatchStreamAdapter;
use datafusion_common::{exec_err, internal_err, DataFusionError};
use datafusion_execution::TaskContext;
#[async_trait]
pub trait DataSink: DisplayAs + Debug + Send + Sync {
async fn write_all(
&self,
data: Vec<SendableRecordBatchStream>,
context: &Arc<TaskContext>,
) -> Result<u64>;
}
pub struct FileSinkExec {
input: Arc<dyn ExecutionPlan>,
sink: Arc<dyn DataSink>,
sink_schema: SchemaRef,
count_schema: SchemaRef,
}
impl fmt::Debug for FileSinkExec {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "FileSinkExec schema: {:?}", self.count_schema)
}
}
impl FileSinkExec {
pub fn new(
input: Arc<dyn ExecutionPlan>,
sink: Arc<dyn DataSink>,
sink_schema: SchemaRef,
) -> Self {
Self {
input,
sink,
sink_schema,
count_schema: make_count_schema(),
}
}
fn execute_input_stream(
&self,
partition: usize,
context: Arc<TaskContext>,
) -> Result<SendableRecordBatchStream> {
let input_stream = self.input.execute(partition, context)?;
debug_assert_eq!(
self.sink_schema.fields().len(),
self.input.schema().fields().len()
);
let risky_columns: Vec<_> = self
.sink_schema
.fields()
.iter()
.zip(self.input.schema().fields().iter())
.enumerate()
.filter_map(|(i, (sink_field, input_field))| {
if !sink_field.is_nullable() && input_field.is_nullable() {
Some(i)
} else {
None
}
})
.collect();
if risky_columns.is_empty() {
Ok(input_stream)
} else {
Ok(Box::pin(RecordBatchStreamAdapter::new(
self.sink_schema.clone(),
input_stream
.map(move |batch| check_not_null_contraits(batch?, &risky_columns)),
)))
}
}
fn execute_all_input_streams(
&self,
context: Arc<TaskContext>,
) -> Result<Vec<SendableRecordBatchStream>> {
let n_input_parts = self.input.output_partitioning().partition_count();
let mut streams = Vec::with_capacity(n_input_parts);
for part in 0..n_input_parts {
streams.push(self.execute_input_stream(part, context.clone())?);
}
Ok(streams)
}
}
impl DisplayAs for FileSinkExec {
fn fmt_as(
&self,
t: DisplayFormatType,
f: &mut std::fmt::Formatter,
) -> std::fmt::Result {
match t {
DisplayFormatType::Default | DisplayFormatType::Verbose => {
write!(f, "InsertExec: sink=")?;
self.sink.fmt_as(t, f)
}
}
}
}
impl ExecutionPlan for FileSinkExec {
fn as_any(&self) -> &dyn Any {
self
}
fn schema(&self) -> SchemaRef {
self.count_schema.clone()
}
fn output_partitioning(&self) -> Partitioning {
Partitioning::UnknownPartitioning(1)
}
fn output_ordering(&self) -> Option<&[PhysicalSortExpr]> {
None
}
fn benefits_from_input_partitioning(&self) -> Vec<bool> {
vec![false]
}
fn required_input_ordering(&self) -> Vec<Option<Vec<PhysicalSortRequirement>>> {
vec![self
.input
.output_ordering()
.map(PhysicalSortRequirement::from_sort_exprs)]
}
fn maintains_input_order(&self) -> Vec<bool> {
vec![false]
}
fn children(&self) -> Vec<Arc<dyn ExecutionPlan>> {
vec![self.input.clone()]
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> Result<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(Self {
input: children[0].clone(),
sink: self.sink.clone(),
sink_schema: self.sink_schema.clone(),
count_schema: self.count_schema.clone(),
}))
}
fn execute(
&self,
partition: usize,
context: Arc<TaskContext>,
) -> Result<SendableRecordBatchStream> {
if partition != 0 {
return internal_err!("FileSinkExec can only be called on partition 0!");
}
let data = self.execute_all_input_streams(context.clone())?;
let count_schema = self.count_schema.clone();
let sink = self.sink.clone();
let stream = futures::stream::once(async move {
sink.write_all(data, &context).await.map(make_count_batch)
})
.boxed();
Ok(Box::pin(RecordBatchStreamAdapter::new(
count_schema,
stream,
)))
}
fn statistics(&self) -> Statistics {
Statistics::default()
}
}
fn make_count_batch(count: u64) -> RecordBatch {
let array = Arc::new(UInt64Array::from(vec![count])) as ArrayRef;
RecordBatch::try_from_iter_with_nullable(vec![("count", array, false)]).unwrap()
}
fn make_count_schema() -> SchemaRef {
Arc::new(Schema::new(vec![Field::new(
"count",
DataType::UInt64,
false,
)]))
}
fn check_not_null_contraits(
batch: RecordBatch,
column_indices: &Vec<usize>,
) -> Result<RecordBatch> {
for &index in column_indices {
if batch.num_columns() <= index {
return exec_err!(
"Invalid batch column count {} expected > {}",
batch.num_columns(),
index
);
}
if batch.column(index).null_count() > 0 {
return exec_err!(
"Invalid batch column at '{}' has null but schema specifies non-nullable",
index
);
}
}
Ok(batch)
}