Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(())
}
}