use ballista_core::client_pool::BallistaClientPool;
use ballista_core::execution_plans::sort_shuffle::SortShuffleWriterExec;
use ballista_core::execution_plans::{ShuffleReaderExec, ShuffleWriterExec};
use ballista_core::serde::protobuf::ShuffleWritePartition;
use ballista_core::{JobId, utils};
use datafusion::common::tree_node::{Transformed, TreeNode};
use datafusion::datasource::memory::MemorySourceConfig;
use datafusion::datasource::physical_plan::{
FileGroup, FileScanConfig, FileScanConfigBuilder,
};
use datafusion::datasource::source::DataSourceExec;
use datafusion::error::{DataFusionError, Result};
use datafusion::execution::context::TaskContext;
use datafusion::physical_plan::ExecutionPlan;
use datafusion::physical_plan::metrics::MetricsSet;
use datafusion::prelude::SessionConfig;
use log::warn;
use std::any::Any;
use std::fmt::{Debug, Display};
use std::sync::Arc;
pub trait ExecutionEngine: Sync + Send {
fn create_query_stage_exec(
&self,
job_id: JobId,
stage_id: usize,
partition_id: usize,
plan: Arc<dyn ExecutionPlan>,
work_dir: &str,
config: &SessionConfig,
) -> Result<Arc<dyn QueryStageExecutor>>;
}
#[async_trait::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>;
}
#[derive(Default)]
pub struct DefaultExecutionEngine {
client_pool: Option<Arc<dyn BallistaClientPool>>,
}
impl DefaultExecutionEngine {
pub fn new() -> Self {
Self { client_pool: None }
}
pub fn with_client_pool(client_pool: Arc<dyn BallistaClientPool>) -> Self {
Self {
client_pool: Some(client_pool),
}
}
}
fn restrict_scan_to_partition(
plan: &Arc<dyn ExecutionPlan>,
partition_id: usize,
) -> Option<Arc<dyn ExecutionPlan>> {
let exec = plan.downcast_ref::<DataSourceExec>()?;
let source: &dyn Any = exec.data_source().as_ref();
let Some(config) = source.downcast_ref::<FileScanConfig>() else {
if source.downcast_ref::<MemorySourceConfig>().is_none() {
warn!(
"restrict_scan_to_partition: unrecognized DataSourceExec source type \
left unrestricted; if it distributes work across partitions from a \
shared queue, a single-partition task could over-read"
);
}
return None;
};
if partition_id >= config.file_groups.len() {
return None;
}
let file_groups: Vec<FileGroup> = config
.file_groups
.iter()
.enumerate()
.map(|(i, group)| {
if i == partition_id {
group.clone()
} else {
FileGroup::new(vec![])
}
})
.collect();
let config = FileScanConfigBuilder::from(config.clone())
.with_file_groups(file_groups)
.build();
Some(DataSourceExec::from_data_source(config))
}
impl ExecutionEngine for DefaultExecutionEngine {
fn create_query_stage_exec(
&self,
job_id: JobId,
stage_id: usize,
partition_id: usize,
plan: Arc<dyn ExecutionPlan>,
work_dir: &str,
_config: &SessionConfig,
) -> Result<Arc<dyn QueryStageExecutor>> {
let plan = plan
.transform(|p| {
if let Some(reader) = p.downcast_ref::<ShuffleReaderExec>() {
match &self.client_pool {
Some(client_pool) => Ok(Transformed::yes(Arc::new(
reader
.with_work_dir(work_dir.to_string())
.with_client_pool(client_pool.clone()),
))),
None => Ok(Transformed::yes(Arc::new(
reader.with_work_dir(work_dir.to_string()),
))),
}
} else if let Some(rewritten) =
restrict_scan_to_partition(&p, partition_id)
{
Ok(Transformed::yes(rewritten))
} else {
Ok(Transformed::no(p))
}
})?
.data;
if let Some(shuffle_writer) = plan.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.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::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),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::datasource::listing::PartitionedFile;
use datafusion::datasource::physical_plan::ParquetSource;
use datafusion::execution::object_store::ObjectStoreUrl;
use datafusion::physical_plan::empty::EmptyExec;
fn scan_with_file_groups(n: usize) -> Arc<dyn ExecutionPlan> {
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
let source = Arc::new(ParquetSource::new(schema));
let mut builder =
FileScanConfigBuilder::new(ObjectStoreUrl::local_filesystem(), source);
for i in 0..n {
builder =
builder.with_file_group(FileGroup::new(vec![PartitionedFile::new(
format!("file{i}.parquet"),
100,
)]));
}
DataSourceExec::from_data_source(builder.build())
}
fn group_file_counts(plan: &Arc<dyn ExecutionPlan>) -> Vec<usize> {
let exec = plan.downcast_ref::<DataSourceExec>().unwrap();
let source: &dyn Any = exec.data_source().as_ref();
let config = source.downcast_ref::<FileScanConfig>().unwrap();
config.file_groups.iter().map(|g| g.len()).collect()
}
#[test]
fn restrict_scan_keeps_only_its_own_group() {
let plan = scan_with_file_groups(4);
let restricted = restrict_scan_to_partition(&plan, 2).expect("scan rewritten");
assert_eq!(group_file_counts(&restricted), vec![0, 0, 1, 0]);
}
#[test]
fn restrict_scan_partition_out_of_range_is_left_untouched() {
let plan = scan_with_file_groups(3);
assert!(restrict_scan_to_partition(&plan, 3).is_none());
}
#[test]
fn restrict_scan_ignores_non_file_scans() {
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));
let plan: Arc<dyn ExecutionPlan> = Arc::new(EmptyExec::new(schema));
assert!(restrict_scan_to_partition(&plan, 0).is_none());
}
}