use std::sync::Arc;
use arrow::datatypes::Schema;
use datafusion_common::{Result, internal_datafusion_err};
use datafusion_execution::TaskContext;
use datafusion_expr::physical_planning_context::ScalarSubqueryResults;
use datafusion_expr::{AggregateUDF, ScalarUDF, WindowUDF};
use datafusion_physical_expr::PhysicalExpr;
use datafusion_physical_expr_common::physical_expr::proto_decode::{
PhysicalExprDecode, PhysicalExprDecodeCtx,
};
use datafusion_physical_expr_common::physical_expr::proto_encode::{
PhysicalExprEncode, PhysicalExprEncodeCtx,
};
use datafusion_proto_models::protobuf::{PhysicalExprNode, PhysicalPlanNode};
use crate::ExecutionPlan;
#[doc(hidden)]
pub trait ExecutionPlanEncode {
fn encode_plan(&self, plan: &Arc<dyn ExecutionPlan>) -> Result<PhysicalPlanNode>;
fn encode_expr(&self, expr: &Arc<dyn PhysicalExpr>) -> Result<PhysicalExprNode>;
fn encode_udf(&self, udf: &ScalarUDF) -> Result<Option<Vec<u8>>>;
fn encode_udaf(&self, udaf: &AggregateUDF) -> Result<Option<Vec<u8>>>;
fn encode_udwf(&self, udwf: &WindowUDF) -> Result<Option<Vec<u8>>>;
}
#[doc(hidden)]
pub trait ExecutionPlanDecode {
fn decode_plan(&self, node: &PhysicalPlanNode) -> Result<Arc<dyn ExecutionPlan>>;
fn decode_plan_with_scalar_subquery_results(
&self,
node: &PhysicalPlanNode,
results: ScalarSubqueryResults,
) -> Result<Arc<dyn ExecutionPlan>>;
fn decode_expr(
&self,
node: &PhysicalExprNode,
input_schema: &Schema,
) -> Result<Arc<dyn PhysicalExpr>>;
fn task_ctx(&self) -> &TaskContext;
fn decode_udf(&self, name: &str, payload: Option<&[u8]>) -> Result<Arc<ScalarUDF>>;
fn decode_udaf(
&self,
name: &str,
payload: Option<&[u8]>,
) -> Result<Arc<AggregateUDF>>;
fn decode_udwf(&self, name: &str, payload: Option<&[u8]>) -> Result<Arc<WindowUDF>>;
}
pub struct ExecutionPlanEncodeCtx<'a> {
encoder: &'a dyn ExecutionPlanEncode,
}
impl<'a> ExecutionPlanEncodeCtx<'a> {
pub fn new(encoder: &'a dyn ExecutionPlanEncode) -> Self {
Self { encoder }
}
pub fn encode_child(
&self,
plan: &Arc<dyn ExecutionPlan>,
) -> Result<PhysicalPlanNode> {
self.encoder.encode_plan(plan)
}
pub fn encode_children<'b, I>(&self, plans: I) -> Result<Vec<PhysicalPlanNode>>
where
I: IntoIterator<Item = &'b Arc<dyn ExecutionPlan>>,
{
plans.into_iter().map(|p| self.encode_child(p)).collect()
}
pub fn encode_expr(&self, expr: &Arc<dyn PhysicalExpr>) -> Result<PhysicalExprNode> {
self.encoder.encode_expr(expr)
}
pub fn encode_expressions<'b, I>(&self, exprs: I) -> Result<Vec<PhysicalExprNode>>
where
I: IntoIterator<Item = &'b Arc<dyn PhysicalExpr>>,
{
exprs.into_iter().map(|e| self.encode_expr(e)).collect()
}
pub fn encode_udf(&self, udf: &ScalarUDF) -> Result<Option<Vec<u8>>> {
self.encoder.encode_udf(udf)
}
pub fn encode_udaf(&self, udaf: &AggregateUDF) -> Result<Option<Vec<u8>>> {
self.encoder.encode_udaf(udaf)
}
pub fn encode_udwf(&self, udwf: &WindowUDF) -> Result<Option<Vec<u8>>> {
self.encoder.encode_udwf(udwf)
}
pub fn expr_ctx(&self) -> PhysicalExprEncodeCtx<'_> {
PhysicalExprEncodeCtx::new(self)
}
}
impl PhysicalExprEncode for ExecutionPlanEncodeCtx<'_> {
fn encode(&self, expr: &Arc<dyn PhysicalExpr>) -> Result<PhysicalExprNode> {
self.encode_expr(expr)
}
}
pub struct ExecutionPlanDecodeCtx<'a> {
decoder: &'a dyn ExecutionPlanDecode,
}
impl<'a> ExecutionPlanDecodeCtx<'a> {
pub fn new(decoder: &'a dyn ExecutionPlanDecode) -> Self {
Self { decoder }
}
pub fn decode_child(
&self,
node: &PhysicalPlanNode,
) -> Result<Arc<dyn ExecutionPlan>> {
self.decoder.decode_plan(node)
}
pub fn decode_child_with_scalar_subquery_results(
&self,
node: &PhysicalPlanNode,
results: ScalarSubqueryResults,
) -> Result<Arc<dyn ExecutionPlan>> {
self.decoder
.decode_plan_with_scalar_subquery_results(node, results)
}
pub fn decode_required_child(
&self,
node: Option<&PhysicalPlanNode>,
plan_name: &str,
field: &str,
) -> Result<Arc<dyn ExecutionPlan>> {
let node = node.ok_or_else(|| {
internal_datafusion_err!("{plan_name} is missing required field '{field}'")
})?;
self.decode_child(node)
}
pub fn decode_expr(
&self,
node: &PhysicalExprNode,
input_schema: &Schema,
) -> Result<Arc<dyn PhysicalExpr>> {
self.decoder.decode_expr(node, input_schema)
}
pub fn decode_required_expr(
&self,
node: Option<&PhysicalExprNode>,
input_schema: &Schema,
plan_name: &str,
field: &str,
) -> Result<Arc<dyn PhysicalExpr>> {
let node = node.ok_or_else(|| {
internal_datafusion_err!("{plan_name} is missing required field '{field}'")
})?;
self.decode_expr(node, input_schema)
}
pub fn task_ctx(&self) -> &TaskContext {
self.decoder.task_ctx()
}
pub fn decode_udf(
&self,
name: &str,
payload: Option<&[u8]>,
) -> Result<Arc<ScalarUDF>> {
self.decoder.decode_udf(name, payload)
}
pub fn decode_udaf(
&self,
name: &str,
payload: Option<&[u8]>,
) -> Result<Arc<AggregateUDF>> {
self.decoder.decode_udaf(name, payload)
}
pub fn decode_udwf(
&self,
name: &str,
payload: Option<&[u8]>,
) -> Result<Arc<WindowUDF>> {
self.decoder.decode_udwf(name, payload)
}
pub fn expr_ctx<'s>(&'s self, input_schema: &'s Schema) -> PhysicalExprDecodeCtx<'s> {
PhysicalExprDecodeCtx::new(input_schema, self)
}
}
impl PhysicalExprDecode for ExecutionPlanDecodeCtx<'_> {
fn decode(
&self,
node: &PhysicalExprNode,
schema: &Schema,
) -> Result<Arc<dyn PhysicalExpr>> {
self.decode_expr(node, schema)
}
}
#[macro_export]
macro_rules! expect_plan_variant {
($node:expr, $variant:path, $plan_name:literal $(,)?) => {{
match &$node.physical_plan_type {
Some($variant(inner)) => inner,
_ => {
return ::datafusion_common::internal_err!(concat!(
"PhysicalPlanNode is not a ",
$plan_name
));
}
}
}};
}