use async_trait::async_trait;
use ballista_core::execution_plans::ShuffleWriterExec;
use ballista_core::execution_plans::sort_shuffle::SortShuffleWriterExec;
use ballista_core::serde::protobuf::ShuffleWritePartition;
use ballista_core::utils;
use datafusion::error::{DataFusionError, Result};
use datafusion::execution::context::TaskContext;
use datafusion::physical_plan::ExecutionPlan;
use datafusion::physical_plan::metrics::MetricsSet;
use std::fmt::{Debug, Display};
use std::sync::Arc;
pub trait ExecutionEngine: Sync + Send {
fn create_query_stage_exec(
&self,
job_id: String,
stage_id: usize,
plan: Arc<dyn ExecutionPlan>,
work_dir: &str,
) -> Result<Arc<dyn QueryStageExecutor>>;
}
#[async_trait]
pub trait QueryStageExecutor: Sync + Send + Debug + Display {
async fn execute_query_stage(
&self,
input_partition: usize,
context: Arc<TaskContext>,
) -> Result<Vec<ShuffleWritePartition>>;
fn collect_plan_metrics(&self) -> Vec<MetricsSet>;
}
pub struct DefaultExecutionEngine {}
impl ExecutionEngine for DefaultExecutionEngine {
fn create_query_stage_exec(
&self,
job_id: String,
stage_id: usize,
plan: Arc<dyn ExecutionPlan>,
work_dir: &str,
) -> Result<Arc<dyn QueryStageExecutor>> {
if let Some(shuffle_writer) = plan.as_any().downcast_ref::<ShuffleWriterExec>() {
let exec = ShuffleWriterExec::try_new(
job_id,
stage_id,
plan.children()[0].clone(),
work_dir.to_string(),
shuffle_writer.shuffle_output_partitioning().cloned(),
)?;
Ok(Arc::new(DefaultQueryStageExec::new(
ShuffleWriterVariant::Hash(exec),
)))
} else if let Some(sort_shuffle_writer) =
plan.as_any().downcast_ref::<SortShuffleWriterExec>()
{
let exec = SortShuffleWriterExec::try_new(
job_id,
stage_id,
plan.children()[0].clone(),
work_dir.to_string(),
sort_shuffle_writer.shuffle_output_partitioning().clone(),
sort_shuffle_writer.config().clone(),
)?;
Ok(Arc::new(DefaultQueryStageExec::new(
ShuffleWriterVariant::Sort(exec),
)))
} else {
Err(DataFusionError::Internal(
"Plan passed to new_query_stage_exec is not a ShuffleWriterExec or SortShuffleWriterExec"
.to_string(),
))
}
}
}
#[derive(Debug, Clone)]
pub enum ShuffleWriterVariant {
Hash(ShuffleWriterExec),
Sort(SortShuffleWriterExec),
}
#[derive(Debug)]
pub struct DefaultQueryStageExec {
shuffle_writer: ShuffleWriterVariant,
}
impl DefaultQueryStageExec {
pub fn new(shuffle_writer: ShuffleWriterVariant) -> Self {
Self { shuffle_writer }
}
}
impl Display for DefaultQueryStageExec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &self.shuffle_writer {
ShuffleWriterVariant::Hash(writer) => {
let stage_metrics: Vec<String> = writer
.metrics()
.unwrap_or_default()
.iter()
.map(|m| m.to_string())
.collect();
write!(
f,
"DefaultQueryStageExec(Hash): ({})\n{}",
stage_metrics.join(", "),
writer
)
}
ShuffleWriterVariant::Sort(writer) => {
let stage_metrics: Vec<String> = writer
.metrics()
.unwrap_or_default()
.iter()
.map(|m| m.to_string())
.collect();
write!(
f,
"DefaultQueryStageExec(Sort): ({})\n{}",
stage_metrics.join(", "),
writer
)
}
}
}
}
#[async_trait]
impl QueryStageExecutor for DefaultQueryStageExec {
async fn execute_query_stage(
&self,
input_partition: usize,
context: Arc<TaskContext>,
) -> Result<Vec<ShuffleWritePartition>> {
match &self.shuffle_writer {
ShuffleWriterVariant::Hash(writer) => {
writer
.clone()
.execute_shuffle_write(input_partition, context)
.await
}
ShuffleWriterVariant::Sort(writer) => {
writer
.clone()
.execute_shuffle_write(input_partition, context)
.await
}
}
}
fn collect_plan_metrics(&self) -> Vec<MetricsSet> {
match &self.shuffle_writer {
ShuffleWriterVariant::Hash(writer) => utils::collect_plan_metrics(writer),
ShuffleWriterVariant::Sort(writer) => utils::collect_plan_metrics(writer),
}
}
}