diff --git a/Cargo.lock b/Cargo.lock index 6914453b3da2c..18a7cf7863d34 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2531,6 +2531,7 @@ dependencies = [ "pretty_assertions", "prost", "rand 0.9.2", + "recursive", "serde", "serde_json", "tokio", diff --git a/datafusion/proto/Cargo.toml b/datafusion/proto/Cargo.toml index a2a0a4546188d..5ff4a8bd5dab8 100644 --- a/datafusion/proto/Cargo.toml +++ b/datafusion/proto/Cargo.toml @@ -37,6 +37,7 @@ name = "datafusion_proto" [features] default = ["parquet"] +recursive_protection = ["dep:recursive"] json = ["pbjson", "serde", "serde_json", "datafusion-proto-common/json"] parquet = ["datafusion-datasource-parquet", "datafusion-common/parquet", "datafusion/parquet"] avro = ["datafusion-datasource-avro", "datafusion-common/avro"] @@ -69,6 +70,7 @@ pbjson = { workspace = true, optional = true } prost = { workspace = true } rand = { workspace = true } serde = { version = "1.0", optional = true } +recursive = { workspace = true, optional = true } serde_json = { workspace = true, optional = true } [dev-dependencies] diff --git a/datafusion/proto/src/logical_plan/mod.rs b/datafusion/proto/src/logical_plan/mod.rs index cd19cb7bcf61e..e8819def06381 100644 --- a/datafusion/proto/src/logical_plan/mod.rs +++ b/datafusion/proto/src/logical_plan/mod.rs @@ -104,7 +104,30 @@ pub trait AsLogicalPlan: Debug + Send + Sync + Clone { Self: Sized; } -pub trait LogicalExtensionCodec: Debug + Send + Sync { +// In debug builds, keep each [de]serializer arm's local temporaries out of the +// recursive dispatcher frame. Without this call boundary, they inflate the +// frame of every recursive invocation. +#[cfg_attr(debug_assertions, inline(never))] +fn serde_logical_plan_arm(f: F) -> Result +where + F: FnOnce() -> Result, +{ + f() +} + +macro_rules! dispatch_logical_plan { + ($plan:expr, { $($pattern:pat => $body:expr $(,)?)+ }) => { + match $plan { + $( + $pattern => serde_logical_plan_arm(|| { + $body + }), + )+ + } + }; +} + +pub trait LogicalExtensionCodec: Debug + Send + Sync + std::any::Any { fn try_decode( &self, buf: &[u8], @@ -383,6 +406,7 @@ impl AsLogicalPlan for LogicalPlanNode { .map_err(|e| internal_datafusion_err!("failed to encode logical plan: {e:?}")) } + #[cfg_attr(feature = "recursive_protection", recursive::recursive)] fn try_into_logical_plan( &self, ctx: &TaskContext, @@ -393,7 +417,7 @@ impl AsLogicalPlan for LogicalPlanNode { "logical_plan::from_proto() Unsupported logical plan '{self:?}'" )) })?; - match plan { + dispatch_logical_plan!(plan, { LogicalPlanType::Values(values) => { let n_cols = values.n_cols as usize; let values: Vec> = if values.values_list.is_empty() { @@ -1084,9 +1108,10 @@ impl AsLogicalPlan for LogicalPlanNode { Arc::new(into_logical_plan!(dml_node.input, ctx, extension_codec)?), ))) } - } + }) } + #[cfg_attr(feature = "recursive_protection", recursive::recursive)] fn try_from_logical_plan( plan: &LogicalPlan, extension_codec: &dyn LogicalExtensionCodec, @@ -1094,7 +1119,7 @@ impl AsLogicalPlan for LogicalPlanNode { where Self: Sized, { - match plan { + dispatch_logical_plan!(plan, { LogicalPlan::Values(Values { values, .. }) => { let n_cols = if values.is_empty() { 0 @@ -1909,6 +1934,6 @@ impl AsLogicalPlan for LogicalPlanNode { ))), }) } - } + }) } } diff --git a/datafusion/proto/tests/cases/mod.rs b/datafusion/proto/tests/cases/mod.rs index aec6c1de30309..18885f90ef409 100644 --- a/datafusion/proto/tests/cases/mod.rs +++ b/datafusion/proto/tests/cases/mod.rs @@ -34,6 +34,7 @@ use std::sync::Arc; mod roundtrip_logical_plan; mod roundtrip_physical_plan; mod serialize; +mod stack_safety; #[derive(Debug, PartialEq, Eq, Hash)] struct MyRegexUdf { diff --git a/datafusion/proto/tests/cases/stack_safety.rs b/datafusion/proto/tests/cases/stack_safety.rs new file mode 100644 index 0000000000000..95e43549ee82d --- /dev/null +++ b/datafusion/proto/tests/cases/stack_safety.rs @@ -0,0 +1,119 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::process::Command; +use std::sync::Arc; + +use datafusion_common::DFSchema; +use datafusion_execution::TaskContext; +use datafusion_expr::logical_plan::{EmptyRelation, LogicalPlan, LogicalPlanBuilder}; +use datafusion_proto::bytes::{logical_plan_from_bytes, logical_plan_to_bytes}; + +const CHILD_ENV: &str = "DATAFUSION_PROTO_ISSUE_23823_CHILD"; +const ALIAS_DEPTH_ENV: &str = "DATAFUSION_PROTO_ISSUE_23823_ALIAS_DEPTH"; +const TWO_MIB_TEST_NAME: &str = + "cases::stack_safety::logical_plan_serde_fits_a_two_mib_stack"; +#[cfg(feature = "recursive_protection")] +const GROWABLE_STACK_TEST_NAME: &str = + "cases::stack_safety::deeply_nested_logical_plan_serde_uses_a_growable_stack"; + +const RECURSION_LIMIT_EXPECTED_DEPTH: usize = 20; + +fn deeply_aliased_plan(alias_depth: usize) -> LogicalPlan { + let mut plan = LogicalPlan::EmptyRelation(EmptyRelation { + produce_one_row: false, + schema: Arc::new(DFSchema::empty()), + }); + + for level in 0..alias_depth { + plan = LogicalPlanBuilder::from(plan) + .alias(format!("level_{level}")) + .unwrap() + .build() + .unwrap(); + } + + plan +} + +fn serde_on_two_mib_stack(alias_depth: usize) { + let plan = deeply_aliased_plan(alias_depth); + std::thread::Builder::new() + .name("two-megabyte-stack".into()) + .stack_size(2 * 1024 * 1024) + .spawn(move || { + let bytes = logical_plan_to_bytes(&plan).unwrap(); + match logical_plan_from_bytes(&bytes, &TaskContext::default()) { + Ok(_) => {} + Err(err) + if err.to_string().contains("recursion limit") + && alias_depth > RECURSION_LIMIT_EXPECTED_DEPTH => + { + // Ignore the error as it is expected for this depth. + } + res => { + res.unwrap(); + } + } + }) + .unwrap() + .join() + .unwrap(); +} + +fn run_in_child(test_name: &str, alias_depth: usize) { + if std::env::var_os(CHILD_ENV).is_some() { + let alias_depth = std::env::var(ALIAS_DEPTH_ENV).unwrap().parse().unwrap(); + serde_on_two_mib_stack(alias_depth); + return; + } + + // A native stack overflow aborts the process. Re-run this exact test in a + // child process so a regression produces a normal test failure. + let output = Command::new(std::env::current_exe().unwrap()) + .args(["--exact", test_name, "--nocapture"]) + .env(CHILD_ENV, "1") + .env(ALIAS_DEPTH_ENV, alias_depth.to_string()) + .output() + .unwrap(); + + assert!( + output.status.success(), + "child process failed with status {}\nstdout:\n{}\nstderr:\n{}", + output.status, + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr), + ); +} + +#[test] +fn logical_plan_serde_fits_a_two_mib_stack() { + // Ten aliases reproduce #23823. Use 100 to provide a safety margin while + // verifying the dispatcher reduction without runtime stack growth. + run_in_child(TWO_MIB_TEST_NAME, 100); + // Check also for a depth that does not trigger a recursion limit error + // on deserialization. + run_in_child(TWO_MIB_TEST_NAME, RECURSION_LIMIT_EXPECTED_DEPTH); +} + +#[cfg(feature = "recursive_protection")] +#[test] +fn deeply_nested_logical_plan_serde_uses_a_growable_stack() { + // This depth exceeds the 2 MiB thread stack without recursive protection, + // exercising the `recursive` stack-growth checkpoint. + run_in_child(GROWABLE_STACK_TEST_NAME, 2_000); +}