From 875a829ee54a00c1fad597fde697bebcf19426c6 Mon Sep 17 00:00:00 2001 From: LiaCastaneda Date: Tue, 5 Aug 2025 15:23:46 +0200 Subject: [PATCH 1/3] Fix serialization bug --- src/plan/codec.rs | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/src/plan/codec.rs b/src/plan/codec.rs index 314a73f6..c58748b1 100644 --- a/src/plan/codec.rs +++ b/src/plan/codec.rs @@ -84,22 +84,30 @@ impl PhysicalExtensionCodec for DistributedCodec { buf: &mut Vec, ) -> datafusion::common::Result<()> { if let Some(node) = node.as_any().downcast_ref::() { - ArrowFlightReadExecProto { + let inner = ArrowFlightReadExecProto { schema: Some(node.schema().try_into()?), partitioning: Some(serialize_partitioning( node.properties().output_partitioning(), &DistributedCodec {}, )?), stage_num: node.stage_num as u64, - } - .encode(buf) - .map_err(|err| proto_error(format!("{err}"))) + }; + + let wrapper = DistributedExecProto { + node: Some(DistributedExecNode::ArrowFlightReadExec(inner)), + }; + + wrapper.encode(buf).map_err(|e| proto_error(format!("{e}"))) } else if let Some(node) = node.as_any().downcast_ref::() { - PartitionIsolatorExecProto { + let inner = PartitionIsolatorExecProto { partition_count: node.partition_count as u64, - } - .encode(buf) - .map_err(|err| proto_error(format!("{err}"))) + }; + + let wrapper = DistributedExecProto { + node: Some(DistributedExecNode::PartitionIsolatorExec(inner)), + }; + + wrapper.encode(buf).map_err(|e| proto_error(format!("{e}"))) } else { Err(proto_error(format!("Unexpected plan {}", node.name()))) } From 73af6352903ca88d27e699ff0af3e72ec4f13cfb Mon Sep 17 00:00:00 2001 From: LiaCastaneda Date: Wed, 6 Aug 2025 09:45:23 +0200 Subject: [PATCH 2/3] Add roundtrip tests --- src/plan/codec.rs | 95 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 95 insertions(+) diff --git a/src/plan/codec.rs b/src/plan/codec.rs index c58748b1..2a39c472 100644 --- a/src/plan/codec.rs +++ b/src/plan/codec.rs @@ -146,3 +146,98 @@ pub struct ArrowFlightReadExecProto { #[prost(uint64, tag = "3")] stage_num: u64, } + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::arrow::datatypes::{DataType, Field}; + use datafusion::{ + execution::registry::MemoryFunctionRegistry, + physical_expr::{expressions::col, expressions::Column, Partitioning, PhysicalSortExpr}, + physical_plan::{displayable, sorts::sort::SortExec, union::UnionExec, ExecutionPlan}, + }; + + type TestCase = ( + &'static str, + Arc, + Vec>, + ); + + fn schema_i32(name: &str) -> Arc { + Arc::new(Schema::new(vec![Field::new(name, DataType::Int32, false)])) + } + + fn repr(plan: &Arc) -> String { + displayable(plan.as_ref()).indent(true).to_string() + } + + #[test] + fn distributed_codec_roundtrips() -> datafusion::common::Result<()> { + let codec = DistributedCodec; + let registry = MemoryFunctionRegistry::new(); + + let mut cases: Vec = Vec::new(); + + // ArrowFlightReadExec + let schema = schema_i32("a"); + let part = Partitioning::Hash(vec![Arc::new(Column::new("a", 0))], 4); + let plan: Arc = Arc::new(ArrowFlightReadExec::new(part, schema, 0)); + cases.push(("single_flight", plan, vec![])); + + // PartitionIsolatorExec -> ArrowFlightReadExec + let schema = schema_i32("b"); + let flight = Arc::new(ArrowFlightReadExec::new( + Partitioning::UnknownPartitioning(1), + schema, + 0, + )); + let plan: Arc = Arc::new(PartitionIsolatorExec::new(flight.clone(), 3)); + cases.push(("isolator_flight", plan, vec![flight])); + + // PartitionIsolatorExec -> UnionExec(ArrowFlightReadExec) + let schema = schema_i32("c"); + let left = Arc::new(ArrowFlightReadExec::new( + Partitioning::RoundRobinBatch(2), + schema.clone(), + 0, + )); + let right = Arc::new(ArrowFlightReadExec::new( + Partitioning::RoundRobinBatch(2), + schema.clone(), + 1, + )); + let union = Arc::new(UnionExec::new(vec![left.clone(), right.clone()])); + let plan: Arc = Arc::new(PartitionIsolatorExec::new(union.clone(), 5)); + cases.push(("isolator_union", plan, vec![union])); + + // PartitionIsolatorExec -> SortExec -> ArrowFlightReadExec + let schema = schema_i32("d"); + let flight = Arc::new(ArrowFlightReadExec::new( + Partitioning::UnknownPartitioning(1), + schema.clone(), + 0, + )); + let sort_expr = PhysicalSortExpr { + expr: col("d", &schema)?, + options: Default::default(), + }; + let sort = Arc::new(SortExec::new(vec![sort_expr].into(), flight.clone())); + let plan: Arc = Arc::new(PartitionIsolatorExec::new(sort.clone(), 2)); + cases.push(("isolator_sort_flight", plan, vec![sort])); + + // Test each case + for (name, original, inputs) in cases { + let mut buf = Vec::new(); + codec.try_encode(original.clone(), &mut buf)?; + + let decoded = codec.try_decode(&buf, &inputs, ®istry)?; + + assert_eq!( + repr(&original), + repr(&decoded), + "mismatch after round-trip for {name}" + ); + } + Ok(()) + } +} From bfb00ad6effa652a62b5df75438ef23933d60cfb Mon Sep 17 00:00:00 2001 From: LiaCastaneda Date: Wed, 6 Aug 2025 10:03:38 +0200 Subject: [PATCH 3/3] Separate tests --- src/plan/codec.rs | 79 ++++++++++++++++++++++++++++++----------------- 1 file changed, 51 insertions(+), 28 deletions(-) diff --git a/src/plan/codec.rs b/src/plan/codec.rs index 2a39c472..f4b38719 100644 --- a/src/plan/codec.rs +++ b/src/plan/codec.rs @@ -157,12 +157,6 @@ mod tests { physical_plan::{displayable, sorts::sort::SortExec, union::UnionExec, ExecutionPlan}, }; - type TestCase = ( - &'static str, - Arc, - Vec>, - ); - fn schema_i32(name: &str) -> Arc { Arc::new(Schema::new(vec![Field::new(name, DataType::Int32, false)])) } @@ -172,29 +166,51 @@ mod tests { } #[test] - fn distributed_codec_roundtrips() -> datafusion::common::Result<()> { + fn test_roundtrip_single_flight() -> datafusion::common::Result<()> { let codec = DistributedCodec; let registry = MemoryFunctionRegistry::new(); - let mut cases: Vec = Vec::new(); - - // ArrowFlightReadExec let schema = schema_i32("a"); let part = Partitioning::Hash(vec![Arc::new(Column::new("a", 0))], 4); let plan: Arc = Arc::new(ArrowFlightReadExec::new(part, schema, 0)); - cases.push(("single_flight", plan, vec![])); - // PartitionIsolatorExec -> ArrowFlightReadExec + let mut buf = Vec::new(); + codec.try_encode(plan.clone(), &mut buf)?; + + let decoded = codec.try_decode(&buf, &[], ®istry)?; + assert_eq!(repr(&plan), repr(&decoded)); + + Ok(()) + } + + #[test] + fn test_roundtrip_isolator_flight() -> datafusion::common::Result<()> { + let codec = DistributedCodec; + let registry = MemoryFunctionRegistry::new(); + let schema = schema_i32("b"); let flight = Arc::new(ArrowFlightReadExec::new( Partitioning::UnknownPartitioning(1), schema, 0, )); + let plan: Arc = Arc::new(PartitionIsolatorExec::new(flight.clone(), 3)); - cases.push(("isolator_flight", plan, vec![flight])); - // PartitionIsolatorExec -> UnionExec(ArrowFlightReadExec) + let mut buf = Vec::new(); + codec.try_encode(plan.clone(), &mut buf)?; + + let decoded = codec.try_decode(&buf, &[flight], ®istry)?; + assert_eq!(repr(&plan), repr(&decoded)); + + Ok(()) + } + + #[test] + fn test_roundtrip_isolator_union() -> datafusion::common::Result<()> { + let codec = DistributedCodec; + let registry = MemoryFunctionRegistry::new(); + let schema = schema_i32("c"); let left = Arc::new(ArrowFlightReadExec::new( Partitioning::RoundRobinBatch(2), @@ -206,38 +222,45 @@ mod tests { schema.clone(), 1, )); + let union = Arc::new(UnionExec::new(vec![left.clone(), right.clone()])); let plan: Arc = Arc::new(PartitionIsolatorExec::new(union.clone(), 5)); - cases.push(("isolator_union", plan, vec![union])); - // PartitionIsolatorExec -> SortExec -> ArrowFlightReadExec + let mut buf = Vec::new(); + codec.try_encode(plan.clone(), &mut buf)?; + + let decoded = codec.try_decode(&buf, &[union], ®istry)?; + assert_eq!(repr(&plan), repr(&decoded)); + + Ok(()) + } + + #[test] + fn test_roundtrip_isolator_sort_flight() -> datafusion::common::Result<()> { + let codec = DistributedCodec; + let registry = MemoryFunctionRegistry::new(); + let schema = schema_i32("d"); let flight = Arc::new(ArrowFlightReadExec::new( Partitioning::UnknownPartitioning(1), schema.clone(), 0, )); + let sort_expr = PhysicalSortExpr { expr: col("d", &schema)?, options: Default::default(), }; let sort = Arc::new(SortExec::new(vec![sort_expr].into(), flight.clone())); + let plan: Arc = Arc::new(PartitionIsolatorExec::new(sort.clone(), 2)); - cases.push(("isolator_sort_flight", plan, vec![sort])); - // Test each case - for (name, original, inputs) in cases { - let mut buf = Vec::new(); - codec.try_encode(original.clone(), &mut buf)?; + let mut buf = Vec::new(); + codec.try_encode(plan.clone(), &mut buf)?; - let decoded = codec.try_decode(&buf, &inputs, ®istry)?; + let decoded = codec.try_decode(&buf, &[sort], ®istry)?; + assert_eq!(repr(&plan), repr(&decoded)); - assert_eq!( - repr(&original), - repr(&decoded), - "mismatch after round-trip for {name}" - ); - } Ok(()) } }