From 2d8457e43672bfe842ce08a850ce950b9f3ee2c3 Mon Sep 17 00:00:00 2001 From: Kai Huang Date: Thu, 30 Jul 2026 10:52:12 -0700 Subject: [PATCH 1/3] Fix checked long sum analytics conversion Signed-off-by: Kai Huang --- .../DataFusionFragmentConvertor.java | 57 ++++++++++++- .../DataFusionFragmentConvertorTests.java | 84 +++++++++++++++++++ .../analytics/qa/StatsCommandIT.java | 4 + 3 files changed, 144 insertions(+), 1 deletion(-) diff --git a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java index 4e265ff0d09b6..584b4c006c2dc 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java @@ -28,6 +28,8 @@ import org.apache.calcite.rex.RexBuilder; import org.apache.calcite.rex.RexLiteral; import org.apache.calcite.rex.RexNode; +import org.apache.calcite.rex.RexOver; +import org.apache.calcite.rex.RexWindow; import org.apache.calcite.schema.ColumnStrategy; import org.apache.calcite.sql.SqlAggFunction; import org.apache.calcite.sql.SqlFunction; @@ -70,6 +72,7 @@ import io.substrait.isthmus.TypeConverter; import io.substrait.isthmus.expression.AggregateFunctionConverter; import io.substrait.isthmus.expression.FunctionMappings; +import io.substrait.isthmus.expression.RexExpressionConverter; import io.substrait.isthmus.expression.ScalarFunctionConverter; import io.substrait.isthmus.expression.WindowFunctionConverter; import io.substrait.plan.Plan; @@ -713,7 +716,8 @@ public Optional convert( AggregateCall call, Function rexConverter ) { - Optional bound = super.convert(input, inputType, call, rexConverter); + AggregateCall substraitCall = canonicalizeAggregate(call); + Optional bound = super.convert(input, inputType, substraitCall, rexConverter); if (bound.isEmpty()) { return bound; } @@ -746,6 +750,26 @@ public Optional convert( if (rewritten == null) return bound; return Optional.of(ImmutableAggregateFunctionInvocation.builder().from(fn).arguments(rewritten).build()); } + + private AggregateCall canonicalizeAggregate(AggregateCall call) { + if (call.getAggregation() == SqlStdOperatorTable.SUM || call.getAggregation().getKind() != SqlKind.SUM) { + return call; + } + return AggregateCall.create( + call.getParserPosition(), + SqlStdOperatorTable.SUM, + call.isDistinct(), + call.isApproximate(), + call.ignoreNulls(), + call.rexList, + call.getArgList(), + call.filterArg, + call.distinctKeys, + call.collation, + call.getType(), + call.getName() + ); + } }; // Same APPROX_COUNT_DISTINCT filter as aggConverter — let our `approx_distinct` entry win. WindowFunctionConverter windowConverter = new WindowFunctionConverter( @@ -760,6 +784,37 @@ protected ImmutableList getSigs() { .filter(sig -> sig.operator != SqlStdOperatorTable.APPROX_COUNT_DISTINCT) .collect(ImmutableList.toImmutableList()); } + + @Override + public Optional convert( + RexOver call, + Function rexConverter, + RexExpressionConverter rexExpressionConverter + ) { + return super.convert(canonicalizeWindow(call), rexConverter, rexExpressionConverter); + } + + private RexOver canonicalizeWindow(RexOver call) { + if (call.getAggOperator() == SqlStdOperatorTable.SUM || call.getAggOperator().getKind() != SqlKind.SUM) { + return call; + } + RexWindow window = call.getWindow(); + return (RexOver) new RexBuilder(typeFactory).makeOver( + call.getType(), + SqlStdOperatorTable.SUM, + call.getOperands(), + window.partitionKeys, + window.orderKeys, + window.getLowerBound(), + window.getUpperBound(), + window.getExclude(), + window.isRows(), + true, + false, + call.isDistinct(), + call.ignoreNulls() + ); + } }; ConverterProvider converterProvider = new ConverterProvider( typeFactory, diff --git a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java index 3c5e4b723822d..a47a3d25ce235 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java @@ -765,6 +765,90 @@ public void testOtherFunctionsNotRenamed() throws Exception { assertTrue("must find sum in extension declarations", foundSum); } + /** + * A custom aggregate with SUM semantics is canonicalized to Calcite's standard SUM for + * Substrait binding. This covers SQL's CHECKED_LONG_SUM operator. + */ + public void testCustomSumUsesNativeSumBinding() throws Exception { + RelNode scan = buildTableScan("test_index", "A"); + SqlAggFunction checkedLongSum = checkedLongSumOperator(); + AggregateCall checkedSumCall = AggregateCall.create( + checkedLongSum, + false, + List.of(0), + -1, + typeFactory.createTypeWithNullability(typeFactory.createSqlType(SqlTypeName.BIGINT), true), + "sum_col" + ); + LogicalAggregate agg = LogicalAggregate.create(scan, List.of(), ImmutableBitSet.of(), null, List.of(checkedSumCall)); + + Plan plan = decodeSubstrait(newConvertor().convertFragment(agg)); + + assertUsesNativeSum(plan); + } + + /** A custom SUM window call uses the same native SUM binding as an aggregate call. */ + public void testCustomWindowSumUsesNativeSumBinding() throws Exception { + RelNode scan = buildTableScan("test_index", "A"); + RelDataType bigintType = typeFactory.createTypeWithNullability(typeFactory.createSqlType(SqlTypeName.BIGINT), true); + RexNode checkedSumOver = rexBuilder.makeOver( + bigintType, + checkedLongSumOperator(), + List.of(rexBuilder.makeInputRef(scan, 0)), + List.of(), + com.google.common.collect.ImmutableList.of(), + org.apache.calcite.rex.RexWindowBounds.UNBOUNDED_PRECEDING, + org.apache.calcite.rex.RexWindowBounds.CURRENT_ROW, + true, + true, + false, + false, + false + ); + RelNode project = org.apache.calcite.rel.logical.LogicalProject.create( + scan, + List.of(), + List.of(rexBuilder.makeInputRef(scan, 0), checkedSumOver), + List.of("A", "running_sum"), + java.util.Set.of() + ); + + Plan plan = decodeSubstrait(newConvertor().convertFragment(project)); + + assertUsesNativeSum(plan); + } + + private SqlAggFunction checkedLongSumOperator() { + return new SqlAggFunction( + "CHECKED_LONG_SUM", + null, + SqlKind.SUM, + ReturnTypes.BIGINT_NULLABLE, + null, + OperandTypes.NUMERIC, + SqlFunctionCategory.USER_DEFINED_FUNCTION, + false, + false, + Optionality.FORBIDDEN + ) { + }; + } + + private void assertUsesNativeSum(Plan plan) { + boolean foundSum = false; + for (SimpleExtensionDeclaration decl : plan.getExtensionsList()) { + if (decl.hasExtensionFunction()) { + String name = decl.getExtensionFunction().getName(); + String baseName = name.contains(":") ? name.substring(0, name.indexOf(':')) : name; + assertNotEquals("custom operator name must not reach Substrait", "checked_long_sum", baseName); + if (baseName.equals("sum")) { + foundSum = true; + } + } + } + assertTrue("custom SUM must use the native sum extension", foundSum); + } + /** * Substrait's stdlib only defines min/max for i8..fp64; opensearch_aggregate_functions.yaml * adds str and bool overloads so PPL `stats min/max` over varchar / boolean fields binds. diff --git a/sandbox/qa/analytics-engine-rest/src/test/java/org/opensearch/analytics/qa/StatsCommandIT.java b/sandbox/qa/analytics-engine-rest/src/test/java/org/opensearch/analytics/qa/StatsCommandIT.java index 5706c4d12c08c..d7880bd553f02 100644 --- a/sandbox/qa/analytics-engine-rest/src/test/java/org/opensearch/analytics/qa/StatsCommandIT.java +++ b/sandbox/qa/analytics-engine-rest/src/test/java/org/opensearch/analytics/qa/StatsCommandIT.java @@ -67,6 +67,10 @@ public void testStatsSum() throws IOException { assertRowsEqual("source=" + DATASET.indexName + " | stats sum(num0)", row(10.0)); } + public void testStatsIntegralSum() throws IOException { + assertRowsEqual("source=" + DATASET.indexName + " | stats sum(int0)", row(68)); + } + public void testStatsAvg() throws IOException { // 10.0 / 8 = 1.25. assertRowsEqual("source=" + DATASET.indexName + " | stats avg(num0)", row(1.25)); From a01e3bdb58c477c0a078343fab32de4d5fc5f427 Mon Sep 17 00:00:00 2001 From: Kai Huang Date: Thu, 30 Jul 2026 13:13:44 -0700 Subject: [PATCH 2/3] Fix checked long sum schema collisions Signed-off-by: Kai Huang --- .../rust/src/udaf/checked_long_sum.rs | 159 ++++++++++++++++++ .../rust/src/udaf/mod.rs | 2 + .../DataFusionFragmentConvertor.java | 37 +++- .../opensearch_aggregate_functions.yaml | 13 ++ .../opensearch_window_functions.yaml | 12 ++ .../DataFusionFragmentConvertorTests.java | 54 +++--- 6 files changed, 249 insertions(+), 28 deletions(-) create mode 100644 sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/checked_long_sum.rs diff --git a/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/checked_long_sum.rs b/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/checked_long_sum.rs new file mode 100644 index 0000000000000..29a61a251c386 --- /dev/null +++ b/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/checked_long_sum.rs @@ -0,0 +1,159 @@ +/* + * SPDX-License-Identifier: Apache-2.0 + * + * The OpenSearch Contributors require contributions made to + * this file be licensed under the Apache-2.0 license or a + * compatible open source license. + */ + +//! Distinctly named analytics binding for SQL's `CHECKED_LONG_SUM`. +//! +//! Analytics routing intentionally keeps DataFusion's native SUM semantics. This wrapper delegates +//! every SUM execution path while retaining the `checked_long_sum` name, preventing collisions when +//! a distributed intermediate schema contains both SUM and CHECKED_LONG_SUM over the same field. + +use std::sync::Arc; + +use datafusion::arrow::datatypes::{DataType, FieldRef}; +use datafusion::common::Result; +use datafusion::execution::context::SessionContext; +use datafusion::functions_aggregate::sum::sum_udaf; +use datafusion::logical_expr::expr::AggregateFunction; +use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; +use datafusion::logical_expr::utils::AggregateOrderSensitivity; +use datafusion::logical_expr::{ + Accumulator, AggregateUDF, AggregateUDFImpl, Documentation, Expr, GroupsAccumulator, Operator, + ReversedUDAF, SetMonotonicity, Signature, +}; + +pub fn register_all(ctx: &SessionContext) { + ctx.register_udaf(AggregateUDF::from(CheckedLongSum::new())); +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct CheckedLongSum { + native_sum: Arc, +} + +impl CheckedLongSum { + fn new() -> Self { + Self { + native_sum: sum_udaf(), + } + } +} + +impl AggregateUDFImpl for CheckedLongSum { + fn name(&self) -> &str { + "checked_long_sum" + } + + fn signature(&self) -> &Signature { + self.native_sum.inner().signature() + } + + fn return_type(&self, arg_types: &[DataType]) -> Result { + self.native_sum.inner().return_type(arg_types) + } + + fn accumulator(&self, args: AccumulatorArgs) -> Result> { + self.native_sum.inner().accumulator(args) + } + + fn state_fields(&self, args: StateFieldsArgs) -> Result> { + self.native_sum.inner().state_fields(args) + } + + fn groups_accumulator_supported(&self, args: AccumulatorArgs) -> bool { + self.native_sum.inner().groups_accumulator_supported(args) + } + + fn create_groups_accumulator( + &self, + args: AccumulatorArgs, + ) -> Result> { + self.native_sum.inner().create_groups_accumulator(args) + } + + fn create_sliding_accumulator(&self, args: AccumulatorArgs) -> Result> { + self.native_sum.inner().create_sliding_accumulator(args) + } + + fn reverse_expr(&self) -> ReversedUDAF { + self.native_sum.inner().reverse_expr() + } + + fn order_sensitivity(&self) -> AggregateOrderSensitivity { + self.native_sum.inner().order_sensitivity() + } + + fn documentation(&self) -> Option<&Documentation> { + self.native_sum.inner().documentation() + } + + fn set_monotonicity(&self, data_type: &DataType) -> SetMonotonicity { + self.native_sum.inner().set_monotonicity(data_type) + } + + fn simplify_expr_op_literal( + &self, + aggregate: &AggregateFunction, + arg: &Expr, + op: Operator, + literal: &Expr, + arg_is_left: bool, + ) -> Result> { + self.native_sum + .inner() + .simplify_expr_op_literal(aggregate, arg, op, literal, arg_is_left) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::arrow::array::{Array, Int64Array, RecordBatch}; + use datafusion::arrow::datatypes::{Field, Schema}; + + #[tokio::test] + async fn native_and_checked_sum_keep_distinct_names() { + let ctx = SessionContext::new(); + register_all(&ctx); + let batch = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)])), + vec![Arc::new(Int64Array::from(vec![1, 2, 3]))], + ) + .unwrap(); + ctx.register_batch("t", batch).unwrap(); + + let batches = ctx + .sql("SELECT SUM(x), CHECKED_LONG_SUM(x) FROM t") + .await + .unwrap() + .collect() + .await + .unwrap(); + let result = &batches[0]; + + assert_eq!(result.schema().field(0).name(), "sum(t.x)"); + assert_eq!(result.schema().field(1).name(), "checked_long_sum(t.x)"); + assert_eq!( + result + .column(0) + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + 6 + ); + assert_eq!( + result + .column(1) + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + 6 + ); + } +} diff --git a/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/mod.rs b/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/mod.rs index 54b408a6b8b53..fd04784bafa56 100644 --- a/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/mod.rs +++ b/sandbox/plugins/analytics-backend-datafusion/rust/src/udaf/mod.rs @@ -13,12 +13,14 @@ use datafusion::execution::context::SessionContext; pub mod approx_distinct_safe; +pub mod checked_long_sum; pub mod internal_pattern; pub mod list_merge; pub mod os_count_distinct; pub mod take; pub fn register_all(ctx: &SessionContext) { + checked_long_sum::register_all(ctx); take::register_all(ctx); list_merge::register_all(ctx); internal_pattern::register_all(ctx); diff --git a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java index 584b4c006c2dc..84e1109f6a6bf 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java @@ -434,6 +434,25 @@ public RexNode normaliseLiteralArg(int argIndex, RexLiteral lit, RexBuilder rexB ) { }; + /** + * Analytics-engine binding for SQL's reflective {@code CHECKED_LONG_SUM}. The runtime + * implementation delegates to DataFusion's native SUM but keeps this distinct function name + * so plans containing both SUM and CHECKED_LONG_SUM do not produce duplicate Arrow field names. + */ + static final SqlAggFunction LOCAL_CHECKED_LONG_SUM_OP = new SqlAggFunction( + "checked_long_sum", + null, + SqlKind.SUM, + ReturnTypes.BIGINT_NULLABLE, + null, + OperandTypes.NUMERIC, + SqlFunctionCategory.USER_DEFINED_FUNCTION, + false, + false, + Optionality.FORBIDDEN + ) { + }; + private static final List ADDITIONAL_AGGREGATE_SIGS = List.of( FunctionMappings.s(SqlStdOperatorTable.APPROX_COUNT_DISTINCT, "approx_distinct"), FunctionMappings.s(LOCAL_TAKE_OP, "take"), @@ -444,14 +463,16 @@ public RexNode normaliseLiteralArg(int argIndex, RexLiteral lit, RexBuilder rexB FunctionMappings.s(LOCAL_LIST_MERGE_DISTINCT_OP, "list_merge_distinct"), FunctionMappings.s(LOCAL_PERCENTILE_APPROX_OP, "approx_percentile_cont"), FunctionMappings.s(LOCAL_INTERNAL_PATTERN_OP, "internal_pattern"), - FunctionMappings.s(LOCAL_OS_COUNT_DISTINCT_OP, "os_count_distinct") + FunctionMappings.s(LOCAL_OS_COUNT_DISTINCT_OP, "os_count_distinct"), + FunctionMappings.s(LOCAL_CHECKED_LONG_SUM_OP, "checked_long_sum") ); private static final List ADDITIONAL_WINDOW_SIGS = List.of( FunctionMappings.s(LOCAL_INTERNAL_PATTERN_WINDOW_OP, "internal_pattern"), // Mirror ADDITIONAL_AGGREGATE_SIGS: rename APPROX_COUNT_DISTINCT to DataFusion's `approx_distinct`. FunctionMappings.s(SqlStdOperatorTable.APPROX_COUNT_DISTINCT, "approx_distinct"), - FunctionMappings.s(LOCAL_OS_COUNT_DISTINCT_OP, "os_count_distinct") + FunctionMappings.s(LOCAL_OS_COUNT_DISTINCT_OP, "os_count_distinct"), + FunctionMappings.s(LOCAL_CHECKED_LONG_SUM_OP, "checked_long_sum") ); /** @@ -716,7 +737,7 @@ public Optional convert( AggregateCall call, Function rexConverter ) { - AggregateCall substraitCall = canonicalizeAggregate(call); + AggregateCall substraitCall = bindCheckedLongSum(call); Optional bound = super.convert(input, inputType, substraitCall, rexConverter); if (bound.isEmpty()) { return bound; @@ -751,13 +772,13 @@ public Optional convert( return Optional.of(ImmutableAggregateFunctionInvocation.builder().from(fn).arguments(rewritten).build()); } - private AggregateCall canonicalizeAggregate(AggregateCall call) { + private AggregateCall bindCheckedLongSum(AggregateCall call) { if (call.getAggregation() == SqlStdOperatorTable.SUM || call.getAggregation().getKind() != SqlKind.SUM) { return call; } return AggregateCall.create( call.getParserPosition(), - SqlStdOperatorTable.SUM, + LOCAL_CHECKED_LONG_SUM_OP, call.isDistinct(), call.isApproximate(), call.ignoreNulls(), @@ -791,17 +812,17 @@ public Optional convert( Function rexConverter, RexExpressionConverter rexExpressionConverter ) { - return super.convert(canonicalizeWindow(call), rexConverter, rexExpressionConverter); + return super.convert(bindCheckedLongSum(call), rexConverter, rexExpressionConverter); } - private RexOver canonicalizeWindow(RexOver call) { + private RexOver bindCheckedLongSum(RexOver call) { if (call.getAggOperator() == SqlStdOperatorTable.SUM || call.getAggOperator().getKind() != SqlKind.SUM) { return call; } RexWindow window = call.getWindow(); return (RexOver) new RexBuilder(typeFactory).makeOver( call.getType(), - SqlStdOperatorTable.SUM, + LOCAL_CHECKED_LONG_SUM_OP, call.getOperands(), window.partitionKeys, window.orderKeys, diff --git a/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_aggregate_functions.yaml b/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_aggregate_functions.yaml index c0ec1000a4337..f3ffcb2e30901 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_aggregate_functions.yaml +++ b/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_aggregate_functions.yaml @@ -2,6 +2,19 @@ --- urn: extension:org.opensearch:aggregate_functions aggregate_functions: + - name: checked_long_sum + description: >- + Analytics-engine binding for SQL's CHECKED_LONG_SUM. The DataFusion runtime + delegates this function to its native SUM implementation while retaining a + distinct output name for intermediate schemas. + impls: + - args: + - name: x + value: any + nullability: DECLARED_OUTPUT + decomposable: MANY + intermediate: i64? + return: i64? - name: approx_distinct description: >- Approximate distinct count using HyperLogLog. Maps to DataFusion's diff --git a/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_window_functions.yaml b/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_window_functions.yaml index 132de1448ecbd..829e30199abc0 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_window_functions.yaml +++ b/sandbox/plugins/analytics-backend-datafusion/src/main/resources/opensearch_window_functions.yaml @@ -2,6 +2,18 @@ --- urn: extension:org.opensearch:window_functions window_functions: + - name: checked_long_sum + description: >- + Window form of the analytics-engine CHECKED_LONG_SUM binding. Execution + delegates to DataFusion's native SUM window accumulator. + impls: + - args: + - name: x + value: any + decomposable: MANY + intermediate: i64? + return: i64? + window_type: STREAMING - name: os_count_distinct description: >- Window form of `os_count_distinct(x)`. Emitted as the substrait window diff --git a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java index a47a3d25ce235..638a2cb863756 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java @@ -766,10 +766,10 @@ public void testOtherFunctionsNotRenamed() throws Exception { } /** - * A custom aggregate with SUM semantics is canonicalized to Calcite's standard SUM for - * Substrait binding. This covers SQL's CHECKED_LONG_SUM operator. + * A custom aggregate with SUM semantics binds to the distinctly named analytics adapter. The + * adapter delegates execution to native DataFusion SUM. */ - public void testCustomSumUsesNativeSumBinding() throws Exception { + public void testCustomSumUsesCheckedLongSumBinding() throws Exception { RelNode scan = buildTableScan("test_index", "A"); SqlAggFunction checkedLongSum = checkedLongSumOperator(); AggregateCall checkedSumCall = AggregateCall.create( @@ -784,11 +784,11 @@ public void testCustomSumUsesNativeSumBinding() throws Exception { Plan plan = decodeSubstrait(newConvertor().convertFragment(agg)); - assertUsesNativeSum(plan); + assertUsesCheckedLongSum(plan); } - /** A custom SUM window call uses the same native SUM binding as an aggregate call. */ - public void testCustomWindowSumUsesNativeSumBinding() throws Exception { + /** A custom SUM window call uses the same checked-long-sum adapter as an aggregate call. */ + public void testCustomWindowSumUsesCheckedLongSumBinding() throws Exception { RelNode scan = buildTableScan("test_index", "A"); RelDataType bigintType = typeFactory.createTypeWithNullability(typeFactory.createSqlType(SqlTypeName.BIGINT), true); RexNode checkedSumOver = rexBuilder.makeOver( @@ -815,7 +815,23 @@ public void testCustomWindowSumUsesNativeSumBinding() throws Exception { Plan plan = decodeSubstrait(newConvertor().convertFragment(project)); - assertUsesNativeSum(plan); + assertUsesCheckedLongSum(plan); + } + + /** Native and checked SUM keep distinct extension names when they target the same field. */ + public void testNativeAndCheckedSumUseDistinctBindings() throws Exception { + RelNode scan = buildTableScan("test_index", "A"); + RelDataType integerType = typeFactory.createTypeWithNullability(typeFactory.createSqlType(SqlTypeName.INTEGER), true); + RelDataType bigintType = typeFactory.createTypeWithNullability(typeFactory.createSqlType(SqlTypeName.BIGINT), true); + AggregateCall nativeSum = AggregateCall.create(SqlStdOperatorTable.SUM, false, List.of(0), -1, integerType, "native_sum"); + AggregateCall checkedSum = AggregateCall.create(checkedLongSumOperator(), false, List.of(0), -1, bigintType, "checked_sum"); + LogicalAggregate agg = LogicalAggregate.create(scan, List.of(), ImmutableBitSet.of(), null, List.of(nativeSum, checkedSum)); + + Plan plan = decodeSubstrait(newConvertor().convertFragment(agg)); + java.util.Set functionNames = extensionFunctionNames(plan); + + assertTrue("native SUM binding is required", functionNames.contains("sum")); + assertTrue("checked SUM binding is required", functionNames.contains("checked_long_sum")); } private SqlAggFunction checkedLongSumOperator() { @@ -834,19 +850,17 @@ private SqlAggFunction checkedLongSumOperator() { }; } - private void assertUsesNativeSum(Plan plan) { - boolean foundSum = false; - for (SimpleExtensionDeclaration decl : plan.getExtensionsList()) { - if (decl.hasExtensionFunction()) { - String name = decl.getExtensionFunction().getName(); - String baseName = name.contains(":") ? name.substring(0, name.indexOf(':')) : name; - assertNotEquals("custom operator name must not reach Substrait", "checked_long_sum", baseName); - if (baseName.equals("sum")) { - foundSum = true; - } - } - } - assertTrue("custom SUM must use the native sum extension", foundSum); + private void assertUsesCheckedLongSum(Plan plan) { + assertTrue("custom SUM must use the checked_long_sum adapter", extensionFunctionNames(plan).contains("checked_long_sum")); + } + + private java.util.Set extensionFunctionNames(Plan plan) { + return plan.getExtensionsList() + .stream() + .filter(SimpleExtensionDeclaration::hasExtensionFunction) + .map(decl -> decl.getExtensionFunction().getName()) + .map(name -> name.contains(":") ? name.substring(0, name.indexOf(':')) : name) + .collect(java.util.stream.Collectors.toSet()); } /** From d0cc9030f3788623065df63ba4348272b4276e68 Mon Sep 17 00:00:00 2001 From: Kai Huang Date: Fri, 31 Jul 2026 09:42:40 -0700 Subject: [PATCH 3/3] Limit checked sum operator rebinding Signed-off-by: Kai Huang --- .../be/datafusion/DataFusionFragmentConvertor.java | 10 ++++++++-- .../DataFusionFragmentConvertorTests.java | 13 ++++++++++++- 2 files changed, 20 insertions(+), 3 deletions(-) diff --git a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java index 84e1109f6a6bf..b13a64185e23c 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/main/java/org/opensearch/be/datafusion/DataFusionFragmentConvertor.java @@ -453,6 +453,12 @@ public RexNode normaliseLiteralArg(int argIndex, RexLiteral lit, RexBuilder rexB ) { }; + static boolean isUnboundCheckedLongSum(SqlAggFunction operator) { + return operator != LOCAL_CHECKED_LONG_SUM_OP + && operator.getKind() == SqlKind.SUM + && "CHECKED_LONG_SUM".equalsIgnoreCase(operator.getName()); + } + private static final List ADDITIONAL_AGGREGATE_SIGS = List.of( FunctionMappings.s(SqlStdOperatorTable.APPROX_COUNT_DISTINCT, "approx_distinct"), FunctionMappings.s(LOCAL_TAKE_OP, "take"), @@ -773,7 +779,7 @@ public Optional convert( } private AggregateCall bindCheckedLongSum(AggregateCall call) { - if (call.getAggregation() == SqlStdOperatorTable.SUM || call.getAggregation().getKind() != SqlKind.SUM) { + if (!isUnboundCheckedLongSum(call.getAggregation())) { return call; } return AggregateCall.create( @@ -816,7 +822,7 @@ public Optional convert( } private RexOver bindCheckedLongSum(RexOver call) { - if (call.getAggOperator() == SqlStdOperatorTable.SUM || call.getAggOperator().getKind() != SqlKind.SUM) { + if (!isUnboundCheckedLongSum(call.getAggOperator())) { return call; } RexWindow window = call.getWindow(); diff --git a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java index 638a2cb863756..c7f069658951d 100644 --- a/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java +++ b/sandbox/plugins/analytics-backend-datafusion/src/test/java/org/opensearch/be/datafusion/DataFusionFragmentConvertorTests.java @@ -834,9 +834,20 @@ public void testNativeAndCheckedSumUseDistinctBindings() throws Exception { assertTrue("checked SUM binding is required", functionNames.contains("checked_long_sum")); } + public void testOnlyCheckedLongSumIsRebound() { + assertTrue(DataFusionFragmentConvertor.isUnboundCheckedLongSum(checkedLongSumOperator())); + assertFalse(DataFusionFragmentConvertor.isUnboundCheckedLongSum(SqlStdOperatorTable.SUM)); + assertFalse(DataFusionFragmentConvertor.isUnboundCheckedLongSum(DataFusionFragmentConvertor.LOCAL_CHECKED_LONG_SUM_OP)); + assertFalse(DataFusionFragmentConvertor.isUnboundCheckedLongSum(sumKindOperator("OTHER_SUM"))); + } + private SqlAggFunction checkedLongSumOperator() { + return sumKindOperator("CHECKED_LONG_SUM"); + } + + private SqlAggFunction sumKindOperator(String name) { return new SqlAggFunction( - "CHECKED_LONG_SUM", + name, null, SqlKind.SUM, ReturnTypes.BIGINT_NULLABLE,