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..d2e495c31df80 100644 --- a/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs +++ b/datafusion/substrait/src/logical_plan/producer/expr/aggregate_function.rs @@ -15,10 +15,10 @@ // specific language governing permissions and limitations // under the License. -use crate::logical_plan::producer::SubstraitProducer; +use crate::logical_plan::producer::{SubstraitProducer, to_substrait_type_from_field}; use datafusion::common::DFSchemaRef; -use datafusion::logical_expr::expr; use datafusion::logical_expr::expr::AggregateFunctionParams; +use datafusion::logical_expr::{Expr, ExprSchemable, expr}; use substrait::proto::aggregate_function::AggregationInvocation; use substrait::proto::aggregate_rel::Measure; use substrait::proto::function_argument::ArgType; @@ -54,13 +54,15 @@ pub fn from_aggregate_function( }); } let function_anchor = producer.register_function(func.name().to_string()); + let (_, output_field) = Expr::AggregateFunction(agg_fn.clone()).to_field(schema)?; + let output_type = to_substrait_type_from_field(producer, &output_field)?; #[expect(deprecated)] Ok(Measure { measure: Some(AggregateFunction { function_reference: function_anchor, arguments, sorts, - output_type: None, + output_type: Some(output_type), invocation: match distinct { true => AggregationInvocation::Distinct as i32, false => AggregationInvocation::All as i32, @@ -93,3 +95,51 @@ fn to_substrait_sort_field( sort_kind: Some(SortKind::Direction(sort_kind.into())), }) } + +#[cfg(test)] +mod tests { + use crate::logical_plan::producer::{ + DefaultSubstraitProducer, SubstraitProducer, to_substrait_type, + }; + use datafusion::arrow::datatypes::{DataType, Field, Schema}; + use datafusion::common::{DFSchema, DFSchemaRef}; + use datafusion::execution::SessionStateBuilder; + use datafusion::functions_aggregate::expr_fn::{avg, count, min, sum}; + use datafusion::logical_expr::Expr; + use datafusion::prelude::col; + + #[test] + fn aggregate_function_output_type() -> datafusion::common::Result<()> { + let state = SessionStateBuilder::default().build(); + let schema = + DFSchemaRef::new(DFSchema::try_from(Schema::new(vec![Field::new( + "i", + DataType::Int64, + false, + )]))?); + let mut producer = DefaultSubstraitProducer::new(&state); + + // (aggregate, expected output type, expected nullability) + let cases = [ + (count(col("i")), DataType::Int64, false), + (sum(col("i")), DataType::Int64, true), + (avg(col("i")), DataType::Float64, true), + (min(col("i")), DataType::Int64, true), + ]; + + for (expr, expected_type, expected_nullable) in cases { + let Expr::AggregateFunction(agg_fn) = &expr else { + panic!("AggregateFunction expected, got {expr}") + }; + let measure = producer.handle_aggregate_function(agg_fn, &schema)?; + let expected = + to_substrait_type(&mut producer, &expected_type, expected_nullable)?; + let output_type = measure + .measure + .expect("Measure should contain an AggregateFunction") + .output_type; + assert_eq!(output_type, Some(expected), "output_type for {expr}"); + } + Ok(()) + } +}