diff --git a/datafusion/physical-plan/src/joins/asof_join.rs b/datafusion/physical-plan/src/joins/asof_join.rs index 3e67e929d0971..abf1b37f1f8e4 100644 --- a/datafusion/physical-plan/src/joins/asof_join.rs +++ b/datafusion/physical-plan/src/joins/asof_join.rs @@ -553,6 +553,180 @@ impl ExecutionPlan for AsOfJoinExec { column_statistics, })) } + #[cfg(feature = "proto")] + fn try_to_proto( + &self, + ctx: &crate::proto::ExecutionPlanEncodeCtx<'_>, + ) -> Result> { + use datafusion_proto_models::protobuf; + + // Destructure exhaustively (no `..`) so that a newly added field is a + // compile error here instead of being silently left out of the proto. + let Self { + left, + right, + on, + match_condition, + projection, + // derived from the children's schemas by `try_new` on decode + join_schema: _, + // derived from the children's schemas by `try_new` on decode + column_indices: _, + // runtime metrics, not part of the plan + metrics: _, + // recomputed from `on` and `match_condition.op` by `try_new` + left_ordering: _, + // recomputed from `on` and `match_condition.op` by `try_new` + right_ordering: _, + // right input collected at execution time, not part of the plan + right_fut: _, + // recomputed by `try_new` on decode + cache: _, + } = self; + + let left = ctx.encode_child(left)?; + let right = ctx.encode_child(right)?; + let on = on + .iter() + .map(|(left, right)| { + Ok(protobuf::JoinOn { + left: Some(ctx.encode_expr(left)?), + right: Some(ctx.encode_expr(right)?), + }) + }) + .collect::>>()?; + let match_operator = match match_condition.op { + Operator::Lt => protobuf::AsOfMatchOperator::Lt, + Operator::LtEq => protobuf::AsOfMatchOperator::LtEq, + Operator::Gt => protobuf::AsOfMatchOperator::Gt, + Operator::GtEq => protobuf::AsOfMatchOperator::GtEq, + op => { + return internal_err!( + "AsOfJoinExec cannot serialize unsupported match operator {op}" + ); + } + }; + + Ok(Some(protobuf::PhysicalPlanNode { + physical_plan_type: Some( + protobuf::physical_plan_node::PhysicalPlanType::AsOfJoin(Box::new( + protobuf::AsOfJoinExecNode { + left: Some(Box::new(left)), + right: Some(Box::new(right)), + on, + left_match_expr: Some(ctx.encode_expr(&match_condition.left)?), + right_match_expr: Some(ctx.encode_expr(&match_condition.right)?), + match_operator: match_operator.into(), + // Proto3 `repeated` cannot distinguish `None` from + // `Some(vec![])`; preserve the empty projection with + // the invalid column-index sentinel used by hash join. + projection: match projection.as_ref() { + None => Vec::new(), + Some(projection) if projection.is_empty() => vec![u32::MAX], + Some(projection) => { + projection.iter().map(|index| *index as u32).collect() + } + }, + }, + )), + ), + })) + } +} + +#[cfg(feature = "proto")] +impl AsOfJoinExec { + /// Reconstruct an [`AsOfJoinExec`] from its protobuf representation. + pub fn try_from_proto( + node: &datafusion_proto_models::protobuf::PhysicalPlanNode, + ctx: &crate::proto::ExecutionPlanDecodeCtx<'_>, + ) -> Result> { + use datafusion_proto_models::protobuf; + + let asof_join = crate::expect_plan_variant!( + node, + protobuf::physical_plan_node::PhysicalPlanType::AsOfJoin, + "AsOfJoinExec", + ); + // Destructure exhaustively (no `..`) so that a newly added proto field + // is a compile error here instead of being silently ignored. + let protobuf::AsOfJoinExecNode { + left, + right, + on, + left_match_expr, + right_match_expr, + match_operator, + projection, + } = &**asof_join; + + let left = ctx.decode_required_child(left.as_deref(), "AsOfJoinExec", "left")?; + let right = + ctx.decode_required_child(right.as_deref(), "AsOfJoinExec", "right")?; + let left_schema = left.schema(); + let right_schema = right.schema(); + let on = on + .iter() + .map(|pair| { + let left = ctx.decode_required_expr( + pair.left.as_ref(), + left_schema.as_ref(), + "AsOfJoinExec", + "on.left", + )?; + let right = ctx.decode_required_expr( + pair.right.as_ref(), + right_schema.as_ref(), + "AsOfJoinExec", + "on.right", + )?; + Ok((left, right)) + }) + .collect::>()?; + let left_match = ctx.decode_required_expr( + left_match_expr.as_ref(), + left_schema.as_ref(), + "AsOfJoinExec", + "left_match_expr", + )?; + let right_match = ctx.decode_required_expr( + right_match_expr.as_ref(), + right_schema.as_ref(), + "AsOfJoinExec", + "right_match_expr", + )?; + let match_operator = protobuf::AsOfMatchOperator::try_from(*match_operator) + .map_err(|_| { + datafusion_common::internal_datafusion_err!( + "AsOfJoinExec: unknown AsOfMatchOperator {}", + match_operator + ) + })?; + let op = match match_operator { + protobuf::AsOfMatchOperator::Lt => Operator::Lt, + protobuf::AsOfMatchOperator::LtEq => Operator::LtEq, + protobuf::AsOfMatchOperator::Gt => Operator::Gt, + protobuf::AsOfMatchOperator::GtEq => Operator::GtEq, + protobuf::AsOfMatchOperator::Unspecified => { + return internal_err!("AsOfJoinExec match operator must be specified"); + } + }; + + // Preserve the empty-projection sentinel written by `try_to_proto`. + let projection = match projection.as_slice() { + [] => None, + [u32::MAX] => Some(Vec::new()), + indices => Some(indices.iter().map(|index| *index as usize).collect()), + }; + + Ok(Arc::new(Self::try_new( + left, + right, + on, + AsOfMatchExpr::new(left_match, op, right_match), + projection, + )?)) + } } /// Materialized right input shared by every left output partition. diff --git a/datafusion/proto-models/proto/datafusion.proto b/datafusion/proto-models/proto/datafusion.proto index 9221c9e59c206..fac5ff27191cd 100644 --- a/datafusion/proto-models/proto/datafusion.proto +++ b/datafusion/proto-models/proto/datafusion.proto @@ -63,6 +63,7 @@ message LogicalPlanNode { CteWorkTableScanNode cte_work_table_scan = 32; DmlNode dml = 33; EmptyTableScanNode empty_table_scan = 34; + AsOfJoinNode as_of_join = 35; } } @@ -276,6 +277,25 @@ message JoinNode { bool null_aware = 9; } +enum AsOfMatchOperator { + AS_OF_MATCH_OPERATOR_UNSPECIFIED = 0; + AS_OF_MATCH_OPERATOR_LT = 1; + AS_OF_MATCH_OPERATOR_LT_EQ = 2; + AS_OF_MATCH_OPERATOR_GT = 3; + AS_OF_MATCH_OPERATOR_GT_EQ = 4; +} + +message AsOfJoinNode { + LogicalPlanNode left = 1; + LogicalPlanNode right = 2; + repeated LogicalExprNode left_join_key = 3; + repeated LogicalExprNode right_join_key = 4; + LogicalExprNode left_match_expr = 5; + LogicalExprNode right_match_expr = 6; + AsOfMatchOperator match_operator = 7; + datafusion_common.JoinConstraint join_constraint = 8; +} + message DistinctNode { LogicalPlanNode input = 1; } @@ -900,6 +920,7 @@ message PhysicalPlanNode { ArrowScanExecNode arrow_scan = 38; ScalarSubqueryExecNode scalar_subquery = 39; PiecewiseMergeJoinExecNode piecewise_merge_join = 40; + AsOfJoinExecNode as_of_join = 41; } } @@ -1725,6 +1746,16 @@ message PiecewiseMergeJoinExecNode { uint64 num_partitions = 7; } +message AsOfJoinExecNode { + PhysicalPlanNode left = 1; + PhysicalPlanNode right = 2; + repeated JoinOn on = 3; + PhysicalExprNode left_match_expr = 4; + PhysicalExprNode right_match_expr = 5; + AsOfMatchOperator match_operator = 6; + repeated uint32 projection = 7; +} + message AsyncFuncExecNode { PhysicalPlanNode input = 1; repeated PhysicalExprNode async_exprs = 2; diff --git a/datafusion/proto-models/src/generated/pbjson.rs b/datafusion/proto-models/src/generated/pbjson.rs index 61c5425585053..e06f5b7011504 100644 --- a/datafusion/proto-models/src/generated/pbjson.rs +++ b/datafusion/proto-models/src/generated/pbjson.rs @@ -1667,6 +1667,507 @@ impl<'de> serde::Deserialize<'de> for ArrowScanExecNode { deserializer.deserialize_struct("datafusion.ArrowScanExecNode", FIELDS, GeneratedVisitor) } } +impl serde::Serialize for AsOfJoinExecNode { + #[allow(deprecated)] + fn serialize(&self, serializer: S) -> std::result::Result + where + S: serde::Serializer, + { + use serde::ser::SerializeStruct; + let mut len = 0; + if self.left.is_some() { + len += 1; + } + if self.right.is_some() { + len += 1; + } + if !self.on.is_empty() { + len += 1; + } + if self.left_match_expr.is_some() { + len += 1; + } + if self.right_match_expr.is_some() { + len += 1; + } + if self.match_operator != 0 { + len += 1; + } + if !self.projection.is_empty() { + len += 1; + } + let mut struct_ser = serializer.serialize_struct("datafusion.AsOfJoinExecNode", len)?; + if let Some(v) = self.left.as_ref() { + struct_ser.serialize_field("left", v)?; + } + if let Some(v) = self.right.as_ref() { + struct_ser.serialize_field("right", v)?; + } + if !self.on.is_empty() { + struct_ser.serialize_field("on", &self.on)?; + } + if let Some(v) = self.left_match_expr.as_ref() { + struct_ser.serialize_field("leftMatchExpr", v)?; + } + if let Some(v) = self.right_match_expr.as_ref() { + struct_ser.serialize_field("rightMatchExpr", v)?; + } + if self.match_operator != 0 { + let v = AsOfMatchOperator::try_from(self.match_operator) + .map_err(|_| serde::ser::Error::custom(format!("Invalid variant {}", self.match_operator)))?; + struct_ser.serialize_field("matchOperator", &v)?; + } + if !self.projection.is_empty() { + struct_ser.serialize_field("projection", &self.projection)?; + } + struct_ser.end() + } +} +impl<'de> serde::Deserialize<'de> for AsOfJoinExecNode { + #[allow(deprecated)] + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + const FIELDS: &[&str] = &[ + "left", + "right", + "on", + "left_match_expr", + "leftMatchExpr", + "right_match_expr", + "rightMatchExpr", + "match_operator", + "matchOperator", + "projection", + ]; + + #[allow(clippy::enum_variant_names)] + enum GeneratedField { + Left, + Right, + On, + LeftMatchExpr, + RightMatchExpr, + MatchOperator, + Projection, + } + impl<'de> serde::Deserialize<'de> for GeneratedField { + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + struct GeneratedVisitor; + + impl serde::de::Visitor<'_> for GeneratedVisitor { + type Value = GeneratedField; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "expected one of: {:?}", &FIELDS) + } + + #[allow(unused_variables)] + fn visit_str(self, value: &str) -> std::result::Result + where + E: serde::de::Error, + { + match value { + "left" => Ok(GeneratedField::Left), + "right" => Ok(GeneratedField::Right), + "on" => Ok(GeneratedField::On), + "leftMatchExpr" | "left_match_expr" => Ok(GeneratedField::LeftMatchExpr), + "rightMatchExpr" | "right_match_expr" => Ok(GeneratedField::RightMatchExpr), + "matchOperator" | "match_operator" => Ok(GeneratedField::MatchOperator), + "projection" => Ok(GeneratedField::Projection), + _ => Err(serde::de::Error::unknown_field(value, FIELDS)), + } + } + } + deserializer.deserialize_identifier(GeneratedVisitor) + } + } + struct GeneratedVisitor; + impl<'de> serde::de::Visitor<'de> for GeneratedVisitor { + type Value = AsOfJoinExecNode; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("struct datafusion.AsOfJoinExecNode") + } + + fn visit_map(self, mut map_: V) -> std::result::Result + where + V: serde::de::MapAccess<'de>, + { + let mut left__ = None; + let mut right__ = None; + let mut on__ = None; + let mut left_match_expr__ = None; + let mut right_match_expr__ = None; + let mut match_operator__ = None; + let mut projection__ = None; + while let Some(k) = map_.next_key()? { + match k { + GeneratedField::Left => { + if left__.is_some() { + return Err(serde::de::Error::duplicate_field("left")); + } + left__ = map_.next_value()?; + } + GeneratedField::Right => { + if right__.is_some() { + return Err(serde::de::Error::duplicate_field("right")); + } + right__ = map_.next_value()?; + } + GeneratedField::On => { + if on__.is_some() { + return Err(serde::de::Error::duplicate_field("on")); + } + on__ = Some(map_.next_value()?); + } + GeneratedField::LeftMatchExpr => { + if left_match_expr__.is_some() { + return Err(serde::de::Error::duplicate_field("leftMatchExpr")); + } + left_match_expr__ = map_.next_value()?; + } + GeneratedField::RightMatchExpr => { + if right_match_expr__.is_some() { + return Err(serde::de::Error::duplicate_field("rightMatchExpr")); + } + right_match_expr__ = map_.next_value()?; + } + GeneratedField::MatchOperator => { + if match_operator__.is_some() { + return Err(serde::de::Error::duplicate_field("matchOperator")); + } + match_operator__ = Some(map_.next_value::()? as i32); + } + GeneratedField::Projection => { + if projection__.is_some() { + return Err(serde::de::Error::duplicate_field("projection")); + } + projection__ = + Some(map_.next_value::>>()? + .into_iter().map(|x| x.0).collect()) + ; + } + } + } + Ok(AsOfJoinExecNode { + left: left__, + right: right__, + on: on__.unwrap_or_default(), + left_match_expr: left_match_expr__, + right_match_expr: right_match_expr__, + match_operator: match_operator__.unwrap_or_default(), + projection: projection__.unwrap_or_default(), + }) + } + } + deserializer.deserialize_struct("datafusion.AsOfJoinExecNode", FIELDS, GeneratedVisitor) + } +} +impl serde::Serialize for AsOfJoinNode { + #[allow(deprecated)] + fn serialize(&self, serializer: S) -> std::result::Result + where + S: serde::Serializer, + { + use serde::ser::SerializeStruct; + let mut len = 0; + if self.left.is_some() { + len += 1; + } + if self.right.is_some() { + len += 1; + } + if !self.left_join_key.is_empty() { + len += 1; + } + if !self.right_join_key.is_empty() { + len += 1; + } + if self.left_match_expr.is_some() { + len += 1; + } + if self.right_match_expr.is_some() { + len += 1; + } + if self.match_operator != 0 { + len += 1; + } + if self.join_constraint != 0 { + len += 1; + } + let mut struct_ser = serializer.serialize_struct("datafusion.AsOfJoinNode", len)?; + if let Some(v) = self.left.as_ref() { + struct_ser.serialize_field("left", v)?; + } + if let Some(v) = self.right.as_ref() { + struct_ser.serialize_field("right", v)?; + } + if !self.left_join_key.is_empty() { + struct_ser.serialize_field("leftJoinKey", &self.left_join_key)?; + } + if !self.right_join_key.is_empty() { + struct_ser.serialize_field("rightJoinKey", &self.right_join_key)?; + } + if let Some(v) = self.left_match_expr.as_ref() { + struct_ser.serialize_field("leftMatchExpr", v)?; + } + if let Some(v) = self.right_match_expr.as_ref() { + struct_ser.serialize_field("rightMatchExpr", v)?; + } + if self.match_operator != 0 { + let v = AsOfMatchOperator::try_from(self.match_operator) + .map_err(|_| serde::ser::Error::custom(format!("Invalid variant {}", self.match_operator)))?; + struct_ser.serialize_field("matchOperator", &v)?; + } + if self.join_constraint != 0 { + let v = super::datafusion_common::JoinConstraint::try_from(self.join_constraint) + .map_err(|_| serde::ser::Error::custom(format!("Invalid variant {}", self.join_constraint)))?; + struct_ser.serialize_field("joinConstraint", &v)?; + } + struct_ser.end() + } +} +impl<'de> serde::Deserialize<'de> for AsOfJoinNode { + #[allow(deprecated)] + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + const FIELDS: &[&str] = &[ + "left", + "right", + "left_join_key", + "leftJoinKey", + "right_join_key", + "rightJoinKey", + "left_match_expr", + "leftMatchExpr", + "right_match_expr", + "rightMatchExpr", + "match_operator", + "matchOperator", + "join_constraint", + "joinConstraint", + ]; + + #[allow(clippy::enum_variant_names)] + enum GeneratedField { + Left, + Right, + LeftJoinKey, + RightJoinKey, + LeftMatchExpr, + RightMatchExpr, + MatchOperator, + JoinConstraint, + } + impl<'de> serde::Deserialize<'de> for GeneratedField { + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + struct GeneratedVisitor; + + impl serde::de::Visitor<'_> for GeneratedVisitor { + type Value = GeneratedField; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "expected one of: {:?}", &FIELDS) + } + + #[allow(unused_variables)] + fn visit_str(self, value: &str) -> std::result::Result + where + E: serde::de::Error, + { + match value { + "left" => Ok(GeneratedField::Left), + "right" => Ok(GeneratedField::Right), + "leftJoinKey" | "left_join_key" => Ok(GeneratedField::LeftJoinKey), + "rightJoinKey" | "right_join_key" => Ok(GeneratedField::RightJoinKey), + "leftMatchExpr" | "left_match_expr" => Ok(GeneratedField::LeftMatchExpr), + "rightMatchExpr" | "right_match_expr" => Ok(GeneratedField::RightMatchExpr), + "matchOperator" | "match_operator" => Ok(GeneratedField::MatchOperator), + "joinConstraint" | "join_constraint" => Ok(GeneratedField::JoinConstraint), + _ => Err(serde::de::Error::unknown_field(value, FIELDS)), + } + } + } + deserializer.deserialize_identifier(GeneratedVisitor) + } + } + struct GeneratedVisitor; + impl<'de> serde::de::Visitor<'de> for GeneratedVisitor { + type Value = AsOfJoinNode; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("struct datafusion.AsOfJoinNode") + } + + fn visit_map(self, mut map_: V) -> std::result::Result + where + V: serde::de::MapAccess<'de>, + { + let mut left__ = None; + let mut right__ = None; + let mut left_join_key__ = None; + let mut right_join_key__ = None; + let mut left_match_expr__ = None; + let mut right_match_expr__ = None; + let mut match_operator__ = None; + let mut join_constraint__ = None; + while let Some(k) = map_.next_key()? { + match k { + GeneratedField::Left => { + if left__.is_some() { + return Err(serde::de::Error::duplicate_field("left")); + } + left__ = map_.next_value()?; + } + GeneratedField::Right => { + if right__.is_some() { + return Err(serde::de::Error::duplicate_field("right")); + } + right__ = map_.next_value()?; + } + GeneratedField::LeftJoinKey => { + if left_join_key__.is_some() { + return Err(serde::de::Error::duplicate_field("leftJoinKey")); + } + left_join_key__ = Some(map_.next_value()?); + } + GeneratedField::RightJoinKey => { + if right_join_key__.is_some() { + return Err(serde::de::Error::duplicate_field("rightJoinKey")); + } + right_join_key__ = Some(map_.next_value()?); + } + GeneratedField::LeftMatchExpr => { + if left_match_expr__.is_some() { + return Err(serde::de::Error::duplicate_field("leftMatchExpr")); + } + left_match_expr__ = map_.next_value()?; + } + GeneratedField::RightMatchExpr => { + if right_match_expr__.is_some() { + return Err(serde::de::Error::duplicate_field("rightMatchExpr")); + } + right_match_expr__ = map_.next_value()?; + } + GeneratedField::MatchOperator => { + if match_operator__.is_some() { + return Err(serde::de::Error::duplicate_field("matchOperator")); + } + match_operator__ = Some(map_.next_value::()? as i32); + } + GeneratedField::JoinConstraint => { + if join_constraint__.is_some() { + return Err(serde::de::Error::duplicate_field("joinConstraint")); + } + join_constraint__ = Some(map_.next_value::()? as i32); + } + } + } + Ok(AsOfJoinNode { + left: left__, + right: right__, + left_join_key: left_join_key__.unwrap_or_default(), + right_join_key: right_join_key__.unwrap_or_default(), + left_match_expr: left_match_expr__, + right_match_expr: right_match_expr__, + match_operator: match_operator__.unwrap_or_default(), + join_constraint: join_constraint__.unwrap_or_default(), + }) + } + } + deserializer.deserialize_struct("datafusion.AsOfJoinNode", FIELDS, GeneratedVisitor) + } +} +impl serde::Serialize for AsOfMatchOperator { + #[allow(deprecated)] + fn serialize(&self, serializer: S) -> std::result::Result + where + S: serde::Serializer, + { + let variant = match self { + Self::Unspecified => "AS_OF_MATCH_OPERATOR_UNSPECIFIED", + Self::Lt => "AS_OF_MATCH_OPERATOR_LT", + Self::LtEq => "AS_OF_MATCH_OPERATOR_LT_EQ", + Self::Gt => "AS_OF_MATCH_OPERATOR_GT", + Self::GtEq => "AS_OF_MATCH_OPERATOR_GT_EQ", + }; + serializer.serialize_str(variant) + } +} +impl<'de> serde::Deserialize<'de> for AsOfMatchOperator { + #[allow(deprecated)] + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + const FIELDS: &[&str] = &[ + "AS_OF_MATCH_OPERATOR_UNSPECIFIED", + "AS_OF_MATCH_OPERATOR_LT", + "AS_OF_MATCH_OPERATOR_LT_EQ", + "AS_OF_MATCH_OPERATOR_GT", + "AS_OF_MATCH_OPERATOR_GT_EQ", + ]; + + struct GeneratedVisitor; + + impl serde::de::Visitor<'_> for GeneratedVisitor { + type Value = AsOfMatchOperator; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "expected one of: {:?}", &FIELDS) + } + + fn visit_i64(self, v: i64) -> std::result::Result + where + E: serde::de::Error, + { + i32::try_from(v) + .ok() + .and_then(|x| x.try_into().ok()) + .ok_or_else(|| { + serde::de::Error::invalid_value(serde::de::Unexpected::Signed(v), &self) + }) + } + + fn visit_u64(self, v: u64) -> std::result::Result + where + E: serde::de::Error, + { + i32::try_from(v) + .ok() + .and_then(|x| x.try_into().ok()) + .ok_or_else(|| { + serde::de::Error::invalid_value(serde::de::Unexpected::Unsigned(v), &self) + }) + } + + fn visit_str(self, value: &str) -> std::result::Result + where + E: serde::de::Error, + { + match value { + "AS_OF_MATCH_OPERATOR_UNSPECIFIED" => Ok(AsOfMatchOperator::Unspecified), + "AS_OF_MATCH_OPERATOR_LT" => Ok(AsOfMatchOperator::Lt), + "AS_OF_MATCH_OPERATOR_LT_EQ" => Ok(AsOfMatchOperator::LtEq), + "AS_OF_MATCH_OPERATOR_GT" => Ok(AsOfMatchOperator::Gt), + "AS_OF_MATCH_OPERATOR_GT_EQ" => Ok(AsOfMatchOperator::GtEq), + _ => Err(serde::de::Error::unknown_variant(value, FIELDS)), + } + } + } + deserializer.deserialize_any(GeneratedVisitor) + } +} impl serde::Serialize for AsyncFuncExecNode { #[allow(deprecated)] fn serialize(&self, serializer: S) -> std::result::Result @@ -13937,6 +14438,9 @@ impl serde::Serialize for LogicalPlanNode { logical_plan_node::LogicalPlanType::EmptyTableScan(v) => { struct_ser.serialize_field("emptyTableScan", v)?; } + logical_plan_node::LogicalPlanType::AsOfJoin(v) => { + struct_ser.serialize_field("asOfJoin", v)?; + } } } struct_ser.end() @@ -13998,6 +14502,8 @@ impl<'de> serde::Deserialize<'de> for LogicalPlanNode { "dml", "empty_table_scan", "emptyTableScan", + "as_of_join", + "asOfJoin", ]; #[allow(clippy::enum_variant_names)] @@ -14035,6 +14541,7 @@ impl<'de> serde::Deserialize<'de> for LogicalPlanNode { CteWorkTableScan, Dml, EmptyTableScan, + AsOfJoin, } impl<'de> serde::Deserialize<'de> for GeneratedField { fn deserialize(deserializer: D) -> std::result::Result @@ -14089,6 +14596,7 @@ impl<'de> serde::Deserialize<'de> for LogicalPlanNode { "cteWorkTableScan" | "cte_work_table_scan" => Ok(GeneratedField::CteWorkTableScan), "dml" => Ok(GeneratedField::Dml), "emptyTableScan" | "empty_table_scan" => Ok(GeneratedField::EmptyTableScan), + "asOfJoin" | "as_of_join" => Ok(GeneratedField::AsOfJoin), _ => Err(serde::de::Error::unknown_field(value, FIELDS)), } } @@ -14340,6 +14848,13 @@ impl<'de> serde::Deserialize<'de> for LogicalPlanNode { return Err(serde::de::Error::duplicate_field("emptyTableScan")); } logical_plan_type__ = map_.next_value::<::std::option::Option<_>>()?.map(logical_plan_node::LogicalPlanType::EmptyTableScan) +; + } + GeneratedField::AsOfJoin => { + if logical_plan_type__.is_some() { + return Err(serde::de::Error::duplicate_field("asOfJoin")); + } + logical_plan_type__ = map_.next_value::<::std::option::Option<_>>()?.map(logical_plan_node::LogicalPlanType::AsOfJoin) ; } } @@ -20679,6 +21194,9 @@ impl serde::Serialize for PhysicalPlanNode { physical_plan_node::PhysicalPlanType::PiecewiseMergeJoin(v) => { struct_ser.serialize_field("piecewiseMergeJoin", v)?; } + physical_plan_node::PhysicalPlanType::AsOfJoin(v) => { + struct_ser.serialize_field("asOfJoin", v)?; + } } } struct_ser.end() @@ -20753,6 +21271,8 @@ impl<'de> serde::Deserialize<'de> for PhysicalPlanNode { "scalarSubquery", "piecewise_merge_join", "piecewiseMergeJoin", + "as_of_join", + "asOfJoin", ]; #[allow(clippy::enum_variant_names)] @@ -20796,6 +21316,7 @@ impl<'de> serde::Deserialize<'de> for PhysicalPlanNode { ArrowScan, ScalarSubquery, PiecewiseMergeJoin, + AsOfJoin, } impl<'de> serde::Deserialize<'de> for GeneratedField { fn deserialize(deserializer: D) -> std::result::Result @@ -20856,6 +21377,7 @@ impl<'de> serde::Deserialize<'de> for PhysicalPlanNode { "arrowScan" | "arrow_scan" => Ok(GeneratedField::ArrowScan), "scalarSubquery" | "scalar_subquery" => Ok(GeneratedField::ScalarSubquery), "piecewiseMergeJoin" | "piecewise_merge_join" => Ok(GeneratedField::PiecewiseMergeJoin), + "asOfJoin" | "as_of_join" => Ok(GeneratedField::AsOfJoin), _ => Err(serde::de::Error::unknown_field(value, FIELDS)), } } @@ -21149,6 +21671,13 @@ impl<'de> serde::Deserialize<'de> for PhysicalPlanNode { return Err(serde::de::Error::duplicate_field("piecewiseMergeJoin")); } physical_plan_type__ = map_.next_value::<::std::option::Option<_>>()?.map(physical_plan_node::PhysicalPlanType::PiecewiseMergeJoin) +; + } + GeneratedField::AsOfJoin => { + if physical_plan_type__.is_some() { + return Err(serde::de::Error::duplicate_field("asOfJoin")); + } + physical_plan_type__ = map_.next_value::<::std::option::Option<_>>()?.map(physical_plan_node::PhysicalPlanType::AsOfJoin) ; } } diff --git a/datafusion/proto-models/src/generated/prost.rs b/datafusion/proto-models/src/generated/prost.rs index d057bfcfdbfe0..149839dacc967 100644 --- a/datafusion/proto-models/src/generated/prost.rs +++ b/datafusion/proto-models/src/generated/prost.rs @@ -5,7 +5,7 @@ pub struct LogicalPlanNode { #[prost( oneof = "logical_plan_node::LogicalPlanType", - tags = "1, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34" + tags = "1, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35" )] pub logical_plan_type: ::core::option::Option, } @@ -79,6 +79,8 @@ pub mod logical_plan_node { Dml(::prost::alloc::boxed::Box), #[prost(message, tag = "34")] EmptyTableScan(super::EmptyTableScanNode), + #[prost(message, tag = "35")] + AsOfJoin(::prost::alloc::boxed::Box), } } #[derive(Clone, PartialEq, ::prost::Message)] @@ -421,6 +423,29 @@ pub struct JoinNode { pub null_aware: bool, } #[derive(Clone, PartialEq, ::prost::Message)] +pub struct AsOfJoinNode { + #[prost(message, optional, boxed, tag = "1")] + pub left: ::core::option::Option<::prost::alloc::boxed::Box>, + #[prost(message, optional, boxed, tag = "2")] + pub right: ::core::option::Option<::prost::alloc::boxed::Box>, + #[prost(message, repeated, tag = "3")] + pub left_join_key: ::prost::alloc::vec::Vec, + #[prost(message, repeated, tag = "4")] + pub right_join_key: ::prost::alloc::vec::Vec, + #[prost(message, optional, boxed, tag = "5")] + pub left_match_expr: ::core::option::Option< + ::prost::alloc::boxed::Box, + >, + #[prost(message, optional, boxed, tag = "6")] + pub right_match_expr: ::core::option::Option< + ::prost::alloc::boxed::Box, + >, + #[prost(enumeration = "AsOfMatchOperator", tag = "7")] + pub match_operator: i32, + #[prost(enumeration = "super::datafusion_common::JoinConstraint", tag = "8")] + pub join_constraint: i32, +} +#[derive(Clone, PartialEq, ::prost::Message)] pub struct DistinctNode { #[prost(message, optional, boxed, tag = "1")] pub input: ::core::option::Option<::prost::alloc::boxed::Box>, @@ -1342,7 +1367,7 @@ pub mod table_reference { pub struct PhysicalPlanNode { #[prost( oneof = "physical_plan_node::PhysicalPlanType", - tags = "1, 2, 3, 4, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40" + tags = "1, 2, 3, 4, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41" )] pub physical_plan_type: ::core::option::Option, } @@ -1432,6 +1457,8 @@ pub mod physical_plan_node { PiecewiseMergeJoin( ::prost::alloc::boxed::Box, ), + #[prost(message, tag = "41")] + AsOfJoin(::prost::alloc::boxed::Box), } } #[derive(Clone, PartialEq, ::prost::Message)] @@ -2614,6 +2641,23 @@ pub struct PiecewiseMergeJoinExecNode { pub num_partitions: u64, } #[derive(Clone, PartialEq, ::prost::Message)] +pub struct AsOfJoinExecNode { + #[prost(message, optional, boxed, tag = "1")] + pub left: ::core::option::Option<::prost::alloc::boxed::Box>, + #[prost(message, optional, boxed, tag = "2")] + pub right: ::core::option::Option<::prost::alloc::boxed::Box>, + #[prost(message, repeated, tag = "3")] + pub on: ::prost::alloc::vec::Vec, + #[prost(message, optional, tag = "4")] + pub left_match_expr: ::core::option::Option, + #[prost(message, optional, tag = "5")] + pub right_match_expr: ::core::option::Option, + #[prost(enumeration = "AsOfMatchOperator", tag = "6")] + pub match_operator: i32, + #[prost(uint32, repeated, tag = "7")] + pub projection: ::prost::alloc::vec::Vec, +} +#[derive(Clone, PartialEq, ::prost::Message)] pub struct AsyncFuncExecNode { #[prost(message, optional, boxed, tag = "1")] pub input: ::core::option::Option<::prost::alloc::boxed::Box>, @@ -2651,6 +2695,41 @@ pub struct PhysicalScalarSubqueryExprNode { ::prost::alloc::string::String, >, } +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, ::prost::Enumeration)] +#[repr(i32)] +pub enum AsOfMatchOperator { + Unspecified = 0, + Lt = 1, + LtEq = 2, + Gt = 3, + GtEq = 4, +} +impl AsOfMatchOperator { + /// String value of the enum field names used in the ProtoBuf definition. + /// + /// The values are not transformed in any way and thus are considered stable + /// (if the ProtoBuf definition does not change) and safe for programmatic use. + pub fn as_str_name(&self) -> &'static str { + match self { + Self::Unspecified => "AS_OF_MATCH_OPERATOR_UNSPECIFIED", + Self::Lt => "AS_OF_MATCH_OPERATOR_LT", + Self::LtEq => "AS_OF_MATCH_OPERATOR_LT_EQ", + Self::Gt => "AS_OF_MATCH_OPERATOR_GT", + Self::GtEq => "AS_OF_MATCH_OPERATOR_GT_EQ", + } + } + /// Creates an enum from field names used in the ProtoBuf definition. + pub fn from_str_name(value: &str) -> ::core::option::Option { + match value { + "AS_OF_MATCH_OPERATOR_UNSPECIFIED" => Some(Self::Unspecified), + "AS_OF_MATCH_OPERATOR_LT" => Some(Self::Lt), + "AS_OF_MATCH_OPERATOR_LT_EQ" => Some(Self::LtEq), + "AS_OF_MATCH_OPERATOR_GT" => Some(Self::Gt), + "AS_OF_MATCH_OPERATOR_GT_EQ" => Some(Self::GtEq), + _ => None, + } + } +} /// Identifies a built-in file format supported by DataFusion. /// Used by DefaultLogicalExtensionCodec to serialize/deserialize /// FileFormatFactory instances (e.g. in CopyTo plans). diff --git a/datafusion/proto/src/logical_plan/mod.rs b/datafusion/proto/src/logical_plan/mod.rs index 647bceeea15cf..37b42fc9a3bab 100644 --- a/datafusion/proto/src/logical_plan/mod.rs +++ b/datafusion/proto/src/logical_plan/mod.rs @@ -59,17 +59,17 @@ use datafusion_datasource_json::file_format::{ use datafusion_datasource_parquet::file_format::{ParquetFormat, ParquetFormatFactory}; use datafusion_expr::dml::InsertOp; use datafusion_expr::{ - AggregateUDF, DmlStatement, FetchType, HigherOrderUDF, RangePartitioning, + AggregateUDF, DmlStatement, FetchType, HigherOrderUDF, Operator, RangePartitioning, RecursiveQuery, SkipType, TableSource, Unnest, WriteOp, }; use datafusion_expr::{ DistinctOn, DropView, Expr, JoinConstraint, LogicalPlan, LogicalPlanBuilder, ScalarUDF, SortExpr, Statement, WindowUDF, dml, logical_plan::{ - Aggregate, CreateCatalog, CreateCatalogSchema, CreateExternalTable, CreateView, - DdlStatement, Distinct, EmptyRelation, Extension, Join, Prepare, Projection, - Repartition, Sort, SubqueryAlias, TableScan, TableScanBuilder, Values, Window, - builder::project, + Aggregate, AsOfJoin, AsOfMatch, CreateCatalog, CreateCatalogSchema, + CreateExternalTable, CreateView, DdlStatement, Distinct, EmptyRelation, + Extension, Join, Prepare, Projection, Repartition, Sort, SubqueryAlias, + TableScan, TableScanBuilder, Values, Window, builder::project, }, }; use datafusion_proto_common::protobuf_common; @@ -1061,6 +1061,69 @@ impl AsLogicalPlan for LogicalPlanNode { join.null_aware, )?)) } + LogicalPlanType::AsOfJoin(join) => { + let left_keys = + from_proto::parse_exprs(&join.left_join_key, ctx, extension_codec)?; + let right_keys = + from_proto::parse_exprs(&join.right_join_key, ctx, extension_codec)?; + if left_keys.len() != right_keys.len() { + return Err(proto_error(format!( + "Received an AsOfJoinNode with left_join_key and right_join_key of different lengths: {} and {}", + left_keys.len(), + right_keys.len() + ))); + } + let left_match = from_proto::parse_expr( + join.left_match_expr.as_ref().ok_or_else(|| { + proto_error("AsOfJoinNode left_match_expr is missing") + })?, + ctx, + extension_codec, + )?; + let right_match = from_proto::parse_expr( + join.right_match_expr.as_ref().ok_or_else(|| { + proto_error("AsOfJoinNode right_match_expr is missing") + })?, + ctx, + extension_codec, + )?; + let match_operator = protobuf::AsOfMatchOperator::try_from( + join.match_operator, + ) + .map_err(|_| { + proto_error(format!( + "Unknown ASOF match operator {}", + join.match_operator + )) + })?; + let op = match match_operator { + protobuf::AsOfMatchOperator::Lt => Operator::Lt, + protobuf::AsOfMatchOperator::LtEq => Operator::LtEq, + protobuf::AsOfMatchOperator::Gt => Operator::Gt, + protobuf::AsOfMatchOperator::GtEq => Operator::GtEq, + protobuf::AsOfMatchOperator::Unspecified => { + return Err(proto_error("ASOF match operator must be specified")); + } + }; + let join_constraint = protobuf::JoinConstraint::try_from( + join.join_constraint, + ) + .map_err(|_| { + proto_error(format!( + "Unknown ASOF JoinConstraint {}", + join.join_constraint + )) + })?; + let left = into_logical_plan!(join.left, ctx, extension_codec)?; + let right = into_logical_plan!(join.right, ctx, extension_codec)?; + Ok(LogicalPlan::AsOfJoin(AsOfJoin::try_new( + Arc::new(left), + Arc::new(right), + left_keys.into_iter().zip(right_keys).collect(), + AsOfMatch::new(left_match, op, right_match), + JoinConstraint::from(join_constraint), + )?)) + } LogicalPlanType::Union(union) => { assert_or_internal_err!( union.inputs.len() >= 2, @@ -1718,6 +1781,68 @@ impl AsLogicalPlan for LogicalPlanNode { ))), }) } + LogicalPlan::AsOfJoin(AsOfJoin { + left, + right, + on, + match_condition, + join_constraint, + .. + }) => { + let left = LogicalPlanNode::try_from_logical_plan( + left.as_ref(), + extension_codec, + )?; + let right = LogicalPlanNode::try_from_logical_plan( + right.as_ref(), + extension_codec, + )?; + let (left_join_key, right_join_key) = on + .iter() + .map(|(left, right)| { + Ok(( + serialize_expr(left, extension_codec)?, + serialize_expr(right, extension_codec)?, + )) + }) + .collect::, ToProtoError>>()? + .into_iter() + .unzip(); + let match_operator = match match_condition.op { + Operator::Lt => protobuf::AsOfMatchOperator::Lt, + Operator::LtEq => protobuf::AsOfMatchOperator::LtEq, + Operator::Gt => protobuf::AsOfMatchOperator::Gt, + Operator::GtEq => protobuf::AsOfMatchOperator::GtEq, + op => { + return Err(proto_error(format!( + "Unsupported ASOF match operator {op}" + ))); + } + }; + Ok(LogicalPlanNode { + logical_plan_type: Some(LogicalPlanType::AsOfJoin(Box::new( + protobuf::AsOfJoinNode { + left: Some(Box::new(left)), + right: Some(Box::new(right)), + left_join_key, + right_join_key, + left_match_expr: Some(Box::new(serialize_expr( + &match_condition.left, + extension_codec, + )?)), + right_match_expr: Some(Box::new(serialize_expr( + &match_condition.right, + extension_codec, + )?)), + match_operator: match_operator.into(), + join_constraint: protobuf::JoinConstraint::from( + *join_constraint, + ) + .into(), + }, + ))), + }) + } LogicalPlan::Subquery(subquery) => { // Serialize the inner subquery plan directly — the // LogicalPlan::Subquery wrapper is reconstructed during @@ -2225,9 +2350,6 @@ impl AsLogicalPlan for LogicalPlanNode { LogicalPlan::DescribeTable(_) => Err(proto_error( "LogicalPlan serde is not yet implemented for DescribeTable", )), - LogicalPlan::AsOfJoin(_) => Err(proto_error( - "LogicalPlan serde is not yet implemented for AsOfJoin", - )), LogicalPlan::RecursiveQuery(recursive) => { let static_term = LogicalPlanNode::try_from_logical_plan( recursive.static_term.as_ref(), diff --git a/datafusion/proto/src/physical_plan/mod.rs b/datafusion/proto/src/physical_plan/mod.rs index 887337e291f43..940c732f81318 100644 --- a/datafusion/proto/src/physical_plan/mod.rs +++ b/datafusion/proto/src/physical_plan/mod.rs @@ -62,8 +62,8 @@ use datafusion_physical_plan::empty::EmptyExec; use datafusion_physical_plan::explain::ExplainExec; use datafusion_physical_plan::filter::FilterExec; use datafusion_physical_plan::joins::{ - CrossJoinExec, HashJoinExec, NestedLoopJoinExec, PiecewiseMergeJoinExec, - SortMergeJoinExec, SymmetricHashJoinExec, + AsOfJoinExec, CrossJoinExec, HashJoinExec, NestedLoopJoinExec, + PiecewiseMergeJoinExec, SortMergeJoinExec, SymmetricHashJoinExec, }; use datafusion_physical_plan::limit::{GlobalLimitExec, LocalLimitExec}; use datafusion_physical_plan::memory::LazyMemoryExec; @@ -1306,6 +1306,9 @@ pub trait PhysicalPlanNodeExt: Sized { PhysicalPlanType::SortMergeJoin(_) => { SortMergeJoinExec::try_from_proto(self.node(), &decode_ctx) } + PhysicalPlanType::AsOfJoin(_) => { + AsOfJoinExec::try_from_proto(self.node(), &decode_ctx) + } PhysicalPlanType::AsyncFunc(_) => { AsyncFuncExec::try_from_proto(self.node(), &decode_ctx) } diff --git a/datafusion/proto/tests/cases/plans/joins.rs b/datafusion/proto/tests/cases/plans/joins.rs index 3c0628f56cf55..da0c0bd5733a1 100644 --- a/datafusion/proto/tests/cases/plans/joins.rs +++ b/datafusion/proto/tests/cases/plans/joins.rs @@ -29,8 +29,9 @@ use datafusion::physical_plan::expressions::{ }; use datafusion::physical_plan::joins::utils::{ColumnIndex, JoinFilter}; use datafusion::physical_plan::joins::{ - HashJoinExec, NestedLoopJoinExec, PartitionMode, PiecewiseMergeJoinExec, - SortMergeJoinExec, StreamJoinPartitionMode, SymmetricHashJoinExec, + AsOfJoinExec, AsOfMatchExpr, HashJoinExec, NestedLoopJoinExec, PartitionMode, + PiecewiseMergeJoinExec, SortMergeJoinExec, StreamJoinPartitionMode, + SymmetricHashJoinExec, }; use datafusion::prelude::SessionContext; use datafusion_common::ScalarValue; @@ -89,6 +90,41 @@ fn roundtrip_hash_join() -> Result<()> { Ok(()) } +#[test] +fn roundtrip_asof_join() -> Result<()> { + let left_schema = Arc::new(Schema::new(vec![ + Field::new("symbol", DataType::Utf8, true), + Field::new("ts", DataType::Int64, true), + Field::new("id", DataType::Int32, false), + ])); + let right_schema = Arc::new(Schema::new(vec![ + Field::new("symbol", DataType::Utf8, true), + Field::new("ts", DataType::Int64, true), + Field::new("price", DataType::Int32, false), + ])); + let on = vec![( + Arc::new(Column::new("symbol", 0)) as _, + Arc::new(Column::new("symbol", 0)) as _, + )]; + + for projection in [None, Some(vec![]), Some(vec![0, 5])] { + for op in [Operator::Lt, Operator::LtEq, Operator::Gt, Operator::GtEq] { + roundtrip_test(Arc::new(AsOfJoinExec::try_new( + Arc::new(EmptyExec::new(Arc::clone(&left_schema))), + Arc::new(EmptyExec::new(Arc::clone(&right_schema))), + on.clone(), + AsOfMatchExpr::new( + Arc::new(Column::new("ts", 1)), + op, + Arc::new(Column::new("ts", 1)), + ), + projection.clone(), + )?))?; + } + } + Ok(()) +} + #[test] fn roundtrip_nested_loop_join() -> Result<()> { let field_a = Field::new("col", DataType::Int64, false); diff --git a/datafusion/proto/tests/cases/roundtrip_logical_plan.rs b/datafusion/proto/tests/cases/roundtrip_logical_plan.rs index 750b20323ad2e..b4c121ef30714 100644 --- a/datafusion/proto/tests/cases/roundtrip_logical_plan.rs +++ b/datafusion/proto/tests/cases/roundtrip_logical_plan.rs @@ -74,8 +74,9 @@ use datafusion_common::format::{ }; use datafusion_common::scalar::ScalarStructBuilder; use datafusion_common::{ - Constraints, DFSchema, DFSchemaRef, DataFusionError, Result, ScalarValue, SplitPoint, - TableReference, internal_datafusion_err, internal_err, not_impl_err, plan_err, + Column, Constraints, DFSchema, DFSchemaRef, DataFusionError, Result, ScalarValue, + SplitPoint, TableReference, internal_datafusion_err, internal_err, not_impl_err, + plan_err, }; use datafusion_execution::TaskContext; use datafusion_expr::dml::CopyTo; @@ -3853,6 +3854,39 @@ async fn roundtrip_join_null_equality() -> Result<()> { Ok(()) } +#[tokio::test] +async fn roundtrip_asof_join() -> Result<()> { + let ctx = SessionContext::new(); + let left_schema = Arc::new(Schema::new(vec![ + Field::new("symbol", DataType::Utf8, true), + Field::new("ts", DataType::Int64, true), + Field::new("id", DataType::Int32, false), + ])); + let right_schema = Arc::new(Schema::new(vec![ + Field::new("symbol", DataType::Utf8, true), + Field::new("ts", DataType::Int64, true), + Field::new("price", DataType::Int32, false), + ])); + ctx.register_table("trades", Arc::new(EmptyTable::new(left_schema)))?; + ctx.register_table("prices", Arc::new(EmptyTable::new(right_schema)))?; + + let left = ctx.table("trades").await?.into_optimized_plan()?; + let right = ctx.table("prices").await?.into_optimized_plan()?; + for op in [Operator::Lt, Operator::LtEq, Operator::Gt, Operator::GtEq] { + let plan = LogicalPlanBuilder::from(left.clone()) + .asof_join_using( + right.clone(), + vec![Column::from_name("symbol")], + binary_expr(col("trades.ts"), op, col("prices.ts")), + )? + .build()?; + let bytes = logical_plan_to_bytes(&plan)?; + let round_trip = logical_plan_from_bytes(&bytes, &ctx.task_ctx())?; + assert_eq!(format!("{plan:?}"), format!("{round_trip:?}")); + } + Ok(()) +} + // Single column, single split point range partitioning #[tokio::test] async fn roundtrip_range_partitioning_single_col() -> Result<()> {