diff --git a/.github/workflows/rust.yml b/.github/workflows/rust.yml
index e51215f99174f..726be0c5e57f0 100644
--- a/.github/workflows/rust.yml
+++ b/.github/workflows/rust.yml
@@ -610,6 +610,48 @@ jobs:
# and this command can be run without filters.
run: cargo xtask ci step test substrait
+ # Runs a plan this crate produces through substrait-java and Spark. A
+ # DataFusion round trip cannot check a field the consumer never reads, so
+ # this is what catches a wrong `AggregationPhase`. It pulls a Spark sized
+ # dependency set, so it only runs when the Substrait crate changes.
+ substrait-java-interop:
+ name: "Run the Substrait java interop test"
+ needs: linux-build-lib
+ runs-on: ubuntu-latest
+ steps:
+ - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
+ - name: Check whether the Substrait crate changed
+ id: filter
+ uses: dorny/paths-filter@fbd0ab8f3e69293af611ebaee6363fc25e6d187d # v3.0.2
+ with:
+ filters: |
+ substrait:
+ - 'datafusion/substrait/**'
+ - '.github/workflows/rust.yml'
+ # `setup-builder` assumes the amd64/rust container this job does not use,
+ # so the two things the Substrait build needs are installed directly.
+ - name: Install protobuf compiler
+ if: steps.filter.outputs.substrait == 'true'
+ run: |
+ sudo apt-get update
+ sudo apt-get install -y protobuf-compiler
+ - name: Setup Rust toolchain
+ if: steps.filter.outputs.substrait == 'true'
+ run: rustup component add rustfmt
+ - name: Setup Java
+ if: steps.filter.outputs.substrait == 'true'
+ uses: actions/setup-java@de7274f081f381c8f8158605e0321c36c376e2e6 # v6
+ with:
+ distribution: temurin
+ java-version: '17'
+ cache: maven
+ - name: Write the plan
+ if: steps.filter.outputs.substrait == 'true'
+ run: cargo test -p datafusion-substrait --test substrait_integration -- --ignored write_java_interop_plan
+ - name: Convert it with substrait-java and run it in Spark
+ if: steps.filter.outputs.substrait == 'true'
+ run: mvn -B -f datafusion/substrait/java-interop/pom.xml test
+
# Temporarily commenting out the Windows flow, the reason is enormously slow running build
# Waiting for new Windows 2025 github runner
# Details: https://github.com/apache/datafusion/issues/13726
diff --git a/datafusion/substrait/java-interop/pom.xml b/datafusion/substrait/java-interop/pom.xml
new file mode 100644
index 0000000000000..c6e21e3685a2b
--- /dev/null
+++ b/datafusion/substrait/java-interop/pom.xml
@@ -0,0 +1,119 @@
+
+
+
+ 4.0.0
+
+ org.apache.datafusion
+ datafusion-substrait-java-interop
+ 1-SNAPSHOT
+ jar
+ DataFusion Substrait java interop test
+
+ Reads a Substrait plan produced by DataFusion, converts it with
+ substrait-java and runs it in Spark, so that fields DataFusion writes but
+ never reads back are checked against another implementation.
+
+
+
+ 17
+ UTF-8
+ 0.103.0
+ 3.5.4
+
+ 2.15.2
+
+
+
+
+ io.substrait
+ spark35_2.12
+ ${substrait.version}
+
+
+ org.apache.spark
+ spark-sql_2.12
+ ${spark.version}
+
+
+ com.fasterxml.jackson.core
+ jackson-databind
+ ${jackson.version}
+
+
+ com.fasterxml.jackson.core
+ jackson-core
+ ${jackson.version}
+
+
+ com.fasterxml.jackson.core
+ jackson-annotations
+ ${jackson.version}
+
+
+ com.fasterxml.jackson.dataformat
+ jackson-dataformat-yaml
+ ${jackson.version}
+
+
+ com.fasterxml.jackson.datatype
+ jackson-datatype-jdk8
+ ${jackson.version}
+
+
+ org.junit.jupiter
+ junit-jupiter
+ 5.11.3
+ test
+
+
+
+
+
+
+ org.apache.maven.plugins
+ maven-surefire-plugin
+ 3.5.2
+
+
+
+ --add-opens=java.base/java.lang=ALL-UNNAMED
+ --add-opens=java.base/java.lang.invoke=ALL-UNNAMED
+ --add-opens=java.base/java.lang.reflect=ALL-UNNAMED
+ --add-opens=java.base/java.io=ALL-UNNAMED
+ --add-opens=java.base/java.net=ALL-UNNAMED
+ --add-opens=java.base/java.nio=ALL-UNNAMED
+ --add-opens=java.base/java.util=ALL-UNNAMED
+ --add-opens=java.base/java.util.concurrent=ALL-UNNAMED
+ --add-opens=java.base/java.util.concurrent.atomic=ALL-UNNAMED
+ --add-opens=java.base/sun.nio.ch=ALL-UNNAMED
+ --add-opens=java.base/sun.nio.cs=ALL-UNNAMED
+ --add-opens=java.base/sun.security.action=ALL-UNNAMED
+ --add-opens=java.base/sun.util.calendar=ALL-UNNAMED
+
+
+ 127.0.0.1
+
+
+
+
+
+
diff --git a/datafusion/substrait/java-interop/src/test/java/org/apache/datafusion/substrait/SubstraitJavaInteropTest.java b/datafusion/substrait/java-interop/src/test/java/org/apache/datafusion/substrait/SubstraitJavaInteropTest.java
new file mode 100644
index 0000000000000..4eb4e65900eae
--- /dev/null
+++ b/datafusion/substrait/java-interop/src/test/java/org/apache/datafusion/substrait/SubstraitJavaInteropTest.java
@@ -0,0 +1,206 @@
+/*
+ * 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.
+ */
+
+package org.apache.datafusion.substrait;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+
+import com.google.protobuf.Descriptors.FieldDescriptor;
+import com.google.protobuf.Message;
+import io.substrait.plan.ProtoPlanConverter;
+import io.substrait.proto.AggregateFunction;
+import io.substrait.proto.Plan;
+import io.substrait.proto.SimpleExtensionDeclaration;
+import io.substrait.proto.SimpleExtensionURN;
+import io.substrait.proto.Type;
+import io.substrait.spark.logical.ToLogicalPlan;
+import java.nio.file.Files;
+import java.nio.file.Path;
+import java.nio.file.Paths;
+import java.util.List;
+import java.util.Map;
+import java.util.regex.Matcher;
+import java.util.regex.Pattern;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Row;
+import org.apache.spark.sql.SparkSession;
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan;
+import org.junit.jupiter.api.Test;
+
+/**
+ * Runs a Substrait plan produced by DataFusion through substrait-java and Spark.
+ *
+ *
A DataFusion to DataFusion round trip cannot check fields the consumer never reads, so it
+ * cannot see a wrong {@code AggregationPhase}: the consumer ignores the field and rebuilds a
+ * complete aggregation either way. substrait-spark does read it, and maps the default
+ * {@code UNSPECIFIED} to Spark's {@code Final}, which makes Spark reject the plan because the
+ * inputs are raw rows rather than partial aggregation buffers.
+ *
+ *
The plan comes from {@code tests/cases/java_interop.rs}. Two workarounds are applied here,
+ * each for a known open issue, so that this test fails for the reason it is about and not for an
+ * unrelated one.
+ */
+public class SubstraitJavaInteropTest {
+
+ /** The rows of `t` on this side; the plan reads a table named `t` with one i64 column. */
+ private static final String TABLE = "CREATE OR REPLACE TEMP VIEW t AS SELECT * FROM VALUES (1L), (2L), (3L) AS t(i)";
+
+ /**
+ * apache/datafusion#11545: the producer writes {@code extension_urn_reference = u32::MAX} and a
+ * bare function name, so no URN resolves. Point each function at the extension that defines it
+ * and give it the compound name substrait-java looks up.
+ */
+ private static final Map URNS = Map.of(
+ "count", new String[] {"extension:io.substrait:functions_aggregate_generic", "count:any"},
+ "sum", new String[] {"extension:io.substrait:functions_arithmetic", "sum:i64"},
+ "avg", new String[] {"extension:io.substrait:functions_arithmetic", "avg:fp64"});
+
+ /**
+ * apache/datafusion#25049: the producer leaves {@code AggregateFunction.output_type} unset and
+ * substrait-java rejects the call. Remove this once that fix lands.
+ */
+ private static final Map OUTPUT_TYPES = Map.of(
+ "count", type(Type.I64.newBuilder().setNullability(Type.Nullability.NULLABILITY_REQUIRED)),
+ "sum", type(Type.I64.newBuilder().setNullability(Type.Nullability.NULLABILITY_NULLABLE)),
+ "avg", type(Type.FP64.newBuilder().setNullability(Type.Nullability.NULLABILITY_NULLABLE)));
+
+ private static Type type(Type.I64.Builder i64) {
+ return Type.newBuilder().setI64(i64).build();
+ }
+
+ private static Type type(Type.FP64.Builder fp64) {
+ return Type.newBuilder().setFp64(fp64).build();
+ }
+
+ private static Path planPath() {
+ String path = System.getProperty("substrait.interop.plan");
+ return path != null ? Paths.get(path) : Paths.get("target", "aggregate_plan.bin");
+ }
+
+ /** Applies both workarounds and returns the plan substrait-java can read. */
+ private static Plan patched(Plan plan) {
+ Plan.Builder builder = plan.toBuilder();
+ Map urnAnchors = new java.util.LinkedHashMap<>();
+ Map typesByAnchor = new java.util.HashMap<>();
+
+ for (int i = 0; i < builder.getExtensionsCount(); i++) {
+ SimpleExtensionDeclaration declaration = builder.getExtensions(i);
+ if (!declaration.hasExtensionFunction()) {
+ continue;
+ }
+ String name = declaration.getExtensionFunction().getName();
+ String[] urn = URNS.get(name);
+ assertTrue(urn != null, "plan declares an unexpected function: " + name);
+ typesByAnchor.put(declaration.getExtensionFunction().getFunctionAnchor(), OUTPUT_TYPES.get(name));
+ int anchor = urnAnchors.computeIfAbsent(urn[0], key -> urnAnchors.size() + 1);
+ builder.setExtensions(
+ i,
+ declaration.toBuilder()
+ .setExtensionFunction(
+ declaration.getExtensionFunction().toBuilder()
+ .setExtensionUrnReference(anchor)
+ .setName(urn[1])));
+ }
+ urnAnchors.forEach(
+ (urn, anchor) ->
+ builder.addExtensionUrns(
+ SimpleExtensionURN.newBuilder().setExtensionUrnAnchor(anchor).setUrn(urn)));
+ return (Plan) fillOutputTypes(builder.build(), typesByAnchor);
+ }
+
+ /** Sets `output_type` on every aggregate call, for as long as #25049 is open. */
+ private static Message fillOutputTypes(Message message, Map byAnchor) {
+ Message.Builder builder = message.toBuilder();
+ for (Map.Entry field : message.getAllFields().entrySet()) {
+ FieldDescriptor descriptor = field.getKey();
+ if (descriptor.getJavaType() != FieldDescriptor.JavaType.MESSAGE) {
+ continue;
+ }
+ if (descriptor.isRepeated()) {
+ builder.clearField(descriptor);
+ for (Object element : (List>) field.getValue()) {
+ builder.addRepeatedField(descriptor, fillOutputTypes((Message) element, byAnchor));
+ }
+ } else {
+ builder.setField(descriptor, fillOutputTypes((Message) field.getValue(), byAnchor));
+ }
+ }
+ if (builder instanceof AggregateFunction.Builder) {
+ AggregateFunction.Builder call = (AggregateFunction.Builder) builder;
+ Type type = byAnchor.get(call.getFunctionReference());
+ if (type != null && !call.hasOutputType()) {
+ call.setOutputType(type);
+ }
+ }
+ return builder.build();
+ }
+
+ @Test
+ public void sparkRunsTheAggregatePlanDataFusionProduced() throws Exception {
+ Path plan = planPath();
+ assertTrue(
+ Files.exists(plan),
+ "run `cargo test -p datafusion-substrait --test substrait_integration -- --ignored "
+ + "write_java_interop_plan` first; expected " + plan.toAbsolutePath());
+
+ SparkSession spark =
+ SparkSession.builder()
+ .master("local[1]")
+ .config("spark.ui.enabled", "false")
+ .config("spark.sql.shuffle.partitions", "1")
+ .getOrCreate();
+ try {
+ spark.sparkContext().setLogLevel("ERROR");
+ spark.sql(TABLE);
+
+ Plan proto = patched(Plan.parseFrom(Files.readAllBytes(plan)));
+ LogicalPlan logical = new ToLogicalPlan(spark).convert(new ProtoPlanConverter().from(proto));
+
+ // The phase decides this: INITIAL_TO_RESULT is Spark's Complete, while the
+ // UNSPECIFIED this producer used to write becomes Final, which Spark rejects
+ // because the inputs are raw rows rather than partial aggregation buffers.
+ List modes = aggregateModes(logical);
+ assertEquals(
+ List.of("Complete", "Complete", "Complete"),
+ modes,
+ "the plan's aggregation phase should reach Spark as Complete");
+
+ List rows = Dataset.ofRows(spark, logical).collectAsList();
+ assertEquals(1, rows.size(), "expected one row");
+ Row row = rows.get(0);
+ assertEquals(3L, row.getLong(0), "count(i)");
+ assertEquals(6L, row.getLong(1), "sum(i)");
+ assertEquals(2.0d, row.getDouble(2), 0.0d, "avg(i)");
+ } finally {
+ spark.stop();
+ }
+ }
+
+ /** The Spark aggregate mode of each aggregate expression, in plan order. */
+ private static List aggregateModes(LogicalPlan plan) {
+ Matcher matcher =
+ Pattern.compile("aggregate\\.(Complete|Final|Partial|PartialMerge)\\$").matcher(plan.toJSON());
+ List modes = new java.util.ArrayList<>();
+ while (matcher.find()) {
+ modes.add(matcher.group(1));
+ }
+ return modes;
+ }
+}
diff --git a/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs b/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs
index 3713f8934f19f..d96f33f49f108 100644
--- a/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs
+++ b/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs
@@ -65,7 +65,7 @@ pub fn from_aggregate_function(
true => AggregationInvocation::Distinct as i32,
false => AggregationInvocation::All as i32,
},
- phase: AggregationPhase::Unspecified as i32,
+ phase: AggregationPhase::InitialToResult as i32,
args: vec![],
options: vec![],
}),
diff --git a/datafusion/substrait/src/logical_plan/producer/expr/window_function.rs b/datafusion/substrait/src/logical_plan/producer/expr/window_function.rs
index d35771bf099d3..f449b8100c34d 100644
--- a/datafusion/substrait/src/logical_plan/producer/expr/window_function.rs
+++ b/datafusion/substrait/src/logical_plan/producer/expr/window_function.rs
@@ -20,6 +20,7 @@ use crate::logical_plan::producer::utils::substrait_sort_field;
use datafusion::common::{DFSchemaRef, ScalarValue, not_impl_err};
use datafusion::logical_expr::expr::{WindowFunction, WindowFunctionParams};
use datafusion::logical_expr::{WindowFrame, WindowFrameBound, WindowFrameUnits};
+use substrait::proto::AggregationPhase;
use substrait::proto::aggregate_function::AggregationInvocation;
use substrait::proto::expression::RexType;
use substrait::proto::expression::WindowFunction as SubstraitWindowFunction;
@@ -108,7 +109,7 @@ fn make_substrait_window_function(
sorts,
options: vec![],
output_type: None,
- phase: 0, // default to AGGREGATION_PHASE_UNSPECIFIED
+ phase: AggregationPhase::InitialToResult as i32,
invocation: if distinct {
AggregationInvocation::Distinct as i32
} else {
diff --git a/datafusion/substrait/tests/cases/java_interop.rs b/datafusion/substrait/tests/cases/java_interop.rs
new file mode 100644
index 0000000000000..1dd0ae196b79a
--- /dev/null
+++ b/datafusion/substrait/tests/cases/java_interop.rs
@@ -0,0 +1,79 @@
+// 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.
+
+//! Writes the Substrait plan that `java-interop` reads.
+//!
+//! The Java side converts the plan with substrait-java, runs it in Spark and
+//! compares the rows, which is the only way to catch a field that DataFusion
+//! writes but never reads back, such as `AggregationPhase`. Producing the plan
+//! needs no JVM, so it lives here; the test is ignored by default because the
+//! file is only useful to that Java project.
+//!
+//! ```shell
+//! cargo test -p datafusion-substrait --test substrait_integration -- --ignored write_java_interop_plan
+//! mvn -f datafusion/substrait/java-interop/pom.xml test
+//! ```
+
+#[cfg(test)]
+mod tests {
+ use datafusion::arrow::datatypes::{DataType, Field, Schema};
+ use datafusion::common::Result;
+ use datafusion::datasource::empty::EmptyTable;
+ use datafusion::prelude::SessionContext;
+ use datafusion_substrait::logical_plan::producer::to_substrait_plan;
+ use prost::Message;
+ use std::path::PathBuf;
+ use std::sync::Arc;
+
+ /// The query the Java side runs. `t` holds 1, 2 and 3 there, so the
+ /// expected rows are 3, 6 and 2.0.
+ const SQL: &str = "SELECT count(i), sum(i), avg(i) FROM t";
+
+ /// Where the plan is written, overridable with `SUBSTRAIT_INTEROP_PLAN`.
+ fn plan_path() -> PathBuf {
+ match std::env::var_os("SUBSTRAIT_INTEROP_PLAN") {
+ Some(path) => PathBuf::from(path),
+ None => PathBuf::from(env!("CARGO_MANIFEST_DIR"))
+ .join("java-interop/target/aggregate_plan.bin"),
+ }
+ }
+
+ #[tokio::test]
+ #[ignore = "writes a file for the java-interop project"]
+ async fn write_java_interop_plan() -> Result<()> {
+ let ctx = SessionContext::new();
+ ctx.register_table(
+ "t",
+ Arc::new(EmptyTable::new(Arc::new(Schema::new(vec![Field::new(
+ "i",
+ DataType::Int64,
+ true,
+ )])))),
+ )?;
+
+ let plan = ctx.sql(SQL).await?.into_optimized_plan()?;
+ let proto = to_substrait_plan(&plan, &ctx.state())?;
+
+ let path = plan_path();
+ if let Some(parent) = path.parent() {
+ std::fs::create_dir_all(parent)?;
+ }
+ std::fs::write(&path, proto.encode_to_vec())?;
+ println!("wrote {}", path.display());
+ Ok(())
+ }
+}
diff --git a/datafusion/substrait/tests/cases/mod.rs b/datafusion/substrait/tests/cases/mod.rs
index 0870c56cd3ba2..2c3ae35c38de3 100644
--- a/datafusion/substrait/tests/cases/mod.rs
+++ b/datafusion/substrait/tests/cases/mod.rs
@@ -20,6 +20,7 @@ mod builtin_expr_semantics_tests;
mod consumer_integration;
mod emit_kind_tests;
mod function_test;
+mod java_interop;
mod logical_plans;
mod roundtrip_logical_plan;
#[cfg(feature = "physical")]
diff --git a/datafusion/substrait/tests/cases/serialize.rs b/datafusion/substrait/tests/cases/serialize.rs
index 4a8413718edb9..809ed612a66d7 100644
--- a/datafusion/substrait/tests/cases/serialize.rs
+++ b/datafusion/substrait/tests/cases/serialize.rs
@@ -35,7 +35,7 @@ mod tests {
use substrait::proto::plan_rel::RelType;
use substrait::proto::rel_common::{Emit, EmitKind};
use substrait::proto::r#type::{I64, Kind as TypeKind, List, Nullability, Struct};
- use substrait::proto::{Expression, RelCommon, Type, rel};
+ use substrait::proto::{AggregationPhase, Expression, RelCommon, Type, rel};
use crate::cases::roundtrip_logical_plan::higher_order_function_ctx;
@@ -321,6 +321,71 @@ mod tests {
Ok(())
}
+ // Collects the `phase` of every aggregate and window function call in a plan.
+ fn collect_phases(rel_type: &rel::RelType, out: &mut Vec) {
+ let input = match rel_type {
+ rel::RelType::Aggregate(aggregate) => {
+ for measure in &aggregate.measures {
+ if let Some(function) = &measure.measure {
+ out.push(function.phase);
+ }
+ }
+ aggregate.input.as_ref()
+ }
+ rel::RelType::Project(project) => {
+ for expr in &project.expressions {
+ if let Some(RexType::WindowFunction(window)) = &expr.rex_type {
+ out.push(window.phase);
+ }
+ }
+ project.input.as_ref()
+ }
+ rel::RelType::Filter(filter) => filter.input.as_ref(),
+ rel::RelType::Sort(sort) => sort.input.as_ref(),
+ rel::RelType::Fetch(fetch) => fetch.input.as_ref(),
+ _ => None,
+ };
+ if let Some(rel_type) = input.and_then(|input| input.rel_type.as_ref()) {
+ collect_phases(rel_type, out);
+ }
+ }
+
+ /// Substrait requires `phase` on aggregate and window function calls, and
+ /// requires `INITIAL_TO_RESULT` for a complete invocation. A DataFusion
+ /// logical plan only ever describes complete aggregations, so that is the
+ /// phase every produced call should carry.
+ #[tokio::test]
+ async fn aggregate_and_window_functions_declare_initial_to_result() -> Result<()> {
+ let ctx = create_context().await?;
+
+ for sql in [
+ "SELECT sum(a) FROM data",
+ "SELECT a, count(*) FROM data GROUP BY a",
+ "SELECT RANK() OVER (PARTITION BY a) FROM data",
+ ] {
+ let plan = ctx.sql(sql).await?.into_optimized_plan()?;
+ let proto = to_substrait_plan(&plan, &ctx.state())?;
+
+ let root = match proto.relations.first().unwrap().rel_type.as_ref() {
+ Some(RelType::Root(root)) => root.input.as_ref().unwrap(),
+ _ => panic!("expected Root"),
+ };
+ let mut phases = vec![];
+ collect_phases(root.rel_type.as_ref().unwrap(), &mut phases);
+
+ assert!(!phases.is_empty(), "no function call found for `{sql}`");
+ for phase in phases {
+ assert_eq!(
+ phase,
+ AggregationPhase::InitialToResult as i32,
+ "phase for `{sql}`"
+ );
+ }
+ }
+
+ Ok(())
+ }
+
fn assert_emit(rel_common: Option<&RelCommon>, output_mapping: Vec) {
assert_eq!(
rel_common.unwrap().emit_kind.clone(),