dd-annarose commented on code in PR #25337:
URL: https://github.com/apache/datafusion/pull/25337#discussion_r4231568792
##########
datafusion/functions-aggregate/src/utils.rs:
##########
@@ -37,41 +41,215 @@ pub(crate) fn get_scalar_value(expr: &Arc<dyn
PhysicalExpr>) -> Result<ScalarVal
}
}
-/// Validates that a percentile expression is a literal float value between
0.0 and 1.0.
-///
-/// Used by both `percentile_cont` and `approx_percentile_cont` to validate
their
-/// percentile parameters.
-pub(crate) fn validate_percentile_expr(
- expr: &Arc<dyn PhysicalExpr>,
- fn_name: &str,
-) -> Result<f64> {
- let scalar_value = get_scalar_value(expr).map_err(|_e| {
- DataFusionError::Plan(format!(
- "Percentile value for '{fn_name}' must be a literal"
- ))
- })?;
-
+/// Validates that a percentile scalar is a Float32/Float64 value between 0.0
and 1.0.
+fn scalar_to_percentile(scalar_value: ScalarValue, fn_name: &str) ->
Result<f64> {
let percentile = match scalar_value {
ScalarValue::Float32(Some(value)) => value as f64,
ScalarValue::Float64(Some(value)) => value,
ScalarValue::Float32(None) | ScalarValue::Float64(None) => {
return plan_err!(
- "Percentile value for '{fn_name}' must be Float32 or Float64
literal (got null)"
+ "Percentile value for '{fn_name}' must be Float32 or Float64
(got null)"
);
}
sv => {
return plan_err!(
- "Percentile value for '{fn_name}' must be Float32 or Float64
literal (got data type {})",
+ "Percentile value for '{fn_name}' must be Float32 or Float64
(got data type {})",
sv.data_type()
);
}
};
- // Ensure the percentile is between 0 and 1.
+ check_percentile_range(percentile)
+}
+
+/// Ensures the percentile is between 0 and 1.
+fn check_percentile_range(percentile: f64) -> Result<f64> {
if !(0.0..=1.0).contains(&percentile) {
return plan_err!(
"Percentile value must be between 0.0 and 1.0 inclusive,
{percentile} is invalid"
);
}
Ok(percentile)
}
+
+/// State of the PercentileParam resolution.
+/// Either already resolved or still requiring a non-empty record batch.
+#[derive(Debug, Clone)]
+pub(crate) enum PercentileParamState {
+ Resolved(f64),
+ Pending,
+}
+
+/// Percentile argument for `aggregate_fn_name` and its state.
+#[derive(Debug)]
+pub struct PercentileParam {
+ pub aggregate_fn_name: String,
+ pub(crate) state: PercentileParamState,
+ pub(crate) is_desc: bool,
+}
+
+impl PercentileParam {
+ /// Try to resolve the percentile eagerly. If the expression can't be
+ /// evaluated without row data (i.e. it references a column), defer
+ /// resolution to the first batch instead of erroring here.
+ pub(crate) fn try_new(
+ expr: &Arc<dyn PhysicalExpr>,
+ fn_name: &str,
+ is_desc: bool,
+ ) -> Result<Self> {
+ match get_scalar_value(expr) {
+ Ok(scalar_value) => Ok(PercentileParam {
+ aggregate_fn_name: fn_name.to_string(),
+ state: PercentileParamState::Resolved(scalar_to_percentile(
+ scalar_value,
+ fn_name,
+ )?),
+ is_desc,
+ }),
+ Err(_) => Ok(PercentileParam {
+ aggregate_fn_name: fn_name.to_string(),
+ state: PercentileParamState::Pending,
+ is_desc,
+ }),
+ }
+ }
+
+ /// Resolve using the current batch if `Pending`
+ /// and validate that the argument is constant across all batches.
+ pub(crate) fn resolve(&mut self, array: &ArrayRef) -> Result<()> {
+ if array.null_count() >= array.len() {
+ return Ok(());
+ }
+
+ let agg_fn_name = self.aggregate_fn_name.clone();
+ let (batch_min, batch_max) = match array.data_type() {
+ DataType::Float64 => {
+ let float_array = downcast_value!(array, Float64Array);
+ (min(float_array), max(float_array))
+ }
+ DataType::Float32 => {
+ let float_array = downcast_value!(array, Float32Array);
+ (
+ min(float_array).map(|v| v as f64),
+ max(float_array).map(|v| v as f64),
+ )
+ }
+ data_type => {
+ return plan_err!(
+ "Percentile value for {agg_fn_name} must be Float32 or
Float64 (got {data_type})"
+ );
+ }
+ };
+ let batch_min = batch_min.ok_or_else(|| {
+ internal_datafusion_err!("expected a non-null percentile value")
+ })?;
+ let batch_max = batch_max.ok_or_else(|| {
+ internal_datafusion_err!("expected a non-null percentile value")
+ })?;
+ if batch_min != batch_max {
+ return plan_err!(
+ "Percentile value for '{agg_fn_name}' must be constant across
the aggregation, found differing values"
+ );
+ }
+
+ match self.state {
+ PercentileParamState::Resolved(resolved) => {
+ if batch_min != resolved {
+ return plan_err!(
+ "Percentile value for '{agg_fn_name}' must be constant
across the aggregation, found differing values"
+ );
+ }
+ }
+ PercentileParamState::Pending => {
+ let resolved = scalar_to_percentile(
+ ScalarValue::Float64(Some(batch_min)),
+ &agg_fn_name,
+ )?;
+ self.state = PercentileParamState::Resolved(resolved);
+ }
+ }
+
+ Ok(())
+ }
+
+ /// Returns the resolved percentile.
+ /// Errors if it has not been yet resolved.
+ pub(crate) fn get(&self) -> Result<f64> {
+ match self.state {
+ PercentileParamState::Resolved(value) => Ok(value),
+ PercentileParamState::Pending => {
+ let aggregate_fn_name = self.aggregate_fn_name.clone();
+ plan_err!(
+ "Percentile value for '{aggregate_fn_name}' could not be
determined: no non-null percentile value was seen"
+ )
+ }
+ }
+ }
+
+ /// The percentile to use, applying the `1.0 - p` flip for descending
+ /// `WITHIN GROUP (ORDER BY ... DESC)`.
+ pub(crate) fn effective_percentile(&self) -> Result<f64> {
+ Ok(self.apply_order(self.get()?))
+ }
+
+ /// Applies the `1.0 - p` flip for descending `WITHIN GROUP (ORDER BY ...
DESC)`.
+ pub(crate) fn apply_order(&self, percentile: f64) -> f64 {
+ if self.is_desc {
+ 1.0 - percentile
+ } else {
+ percentile
+ }
+ }
+
+ /// Returns `true` if the percentile must be read from the input rows,
+ /// i.e. it was given as a column reference rather than a literal.
+ pub(crate) fn is_pending(&self) -> bool {
+ matches!(self.state, PercentileParamState::Pending)
+ }
+
+ /// Casts a percentile argument array to `Float64`.
+ pub(crate) fn to_float64_array(&self, array: &ArrayRef) ->
Result<Float64Array> {
+ match array.data_type() {
+ DataType::Float64 | DataType::Float32 => {
+ let array = cast(array, &DataType::Float64)?;
+ Ok(downcast_value!(array, Float64Array).clone())
Review Comment:
I could use `to_owned` but it's essentially cloning
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]