Skip to content
Open
Show file tree
Hide file tree
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
90 changes: 44 additions & 46 deletions datafusion/common/src/functional_dependencies.rs
Original file line number Diff line number Diff line change
Expand Up @@ -452,9 +452,15 @@ impl Deref for FunctionalDependencies {
}

/// Calculates functional dependencies for aggregate output, when there is a GROUP BY expression.
///
/// `group_by_input_indices[i]` is the index of the input field that GROUP BY
/// expression `i` passes through unchanged (a column reference), or `None` if
/// the expression computes a new value. Only a column reference can carry an
/// input dependency; a computed expression such as `CAST(pk AS INT)` cannot,
/// even though its output name may equal the column's name.
pub fn aggregate_functional_dependencies(
aggr_input_schema: &DFSchema,
group_by_expr_names: &[String],
group_by_input_indices: &[Option<usize>],
aggr_schema: &DFSchema,
) -> FunctionalDependencies {
let mut aggregate_func_dependencies = vec![];
Expand All @@ -466,10 +472,9 @@ pub fn aggregate_functional_dependencies(
// The loop below only re-expresses input dependencies. Skip it when the
// input has none. The GROUP BY-key dependency below always runs.
if !func_dependencies.is_empty() {
let aggr_input_fields = aggr_input_schema.field_names();
// Compute once: this does not change in the loop.
let existing_target_indices =
get_target_functional_dependencies(aggr_input_schema, group_by_expr_names);
get_target_functional_dependencies(aggr_input_schema, group_by_input_indices);
for FunctionalDependence {
source_indices,
nullable,
Expand All @@ -480,24 +485,21 @@ pub fn aggregate_functional_dependencies(
{
// Indices into the GROUP BY list for this determinant:
let mut new_source_indices = vec![];
let mut new_source_field_names = vec![];
let source_field_names = source_indices
.iter()
.map(|&idx| &aggr_input_fields[idx])
.collect::<Vec<_>>();

for (idx, group_by_expr_name) in group_by_expr_names.iter().enumerate() {
// When one of the input determinant expressions matches with
// the GROUP BY expression, add the index of the GROUP BY
// expression as a new determinant key:
if source_field_names.contains(&group_by_expr_name) {
let mut new_source_input_indices = vec![];
for (idx, input_idx) in group_by_input_indices.iter().enumerate() {
// When one of the input determinant columns is a GROUP BY
// expression, add the index of the GROUP BY expression as a new
// determinant key:
if let Some(input_idx) = input_idx
&& source_indices.contains(input_idx)
{
Comment on lines +493 to +495

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A key column that is in the GROUP BY list two times loses its dependency. This is a regression from main, and only the builder / DataFrame API can reach it (SQL removes aliases from GROUP BY).

.aggregate(vec![col("id"), col("state"), col("id").alias("k")], ...) with PRIMARY KEY (id):

Output dependency
main [0] -> [0, 1, 2]
This PR [0, 1, 2] -> [0, 1, 2]

id and id AS k both map to input index 0, so the length check below fails. Results stay correct, but parent plans can no longer prune with id.

Suggested change
if let Some(input_idx) = input_idx
&& source_indices.contains(input_idx)
{
if let Some(input_idx) = input_idx
&& source_indices.contains(input_idx)
// A key column can be in the GROUP BY list more than one
// time (`x, x AS k`). Count it one time.
&& !new_source_input_indices.contains(&Some(*input_idx))
{

A unit test is optional here, because only the builder API reaches this. If you want one, this fails on this PR and passes with the suggestion:

Optional unit test for plan.rs
#[test]
fn aggregate_group_by_repeated_key_keeps_key_dependency() -> Result<()> {
    let constraints =
        Constraints::new_unverified(vec![Constraint::PrimaryKey(vec![0])]);
    let source = Arc::new(
        LogicalTableSource::new(Arc::new(employee_schema()))
            .with_constraints(constraints),
    );
    // `id` is in the GROUP BY list two times: as itself and as `k`.
    let plan = LogicalPlanBuilder::scan("employee_csv", source, None)?
        .aggregate(
            vec![col("id"), col("state"), col("id").alias("k")],
            Vec::<Expr>::new(),
        )?
        .build()?;

    // `id` alone still determines the row.
    let deps = plan.schema().functional_dependencies();
    assert_eq!(deps.len(), 1);
    assert_eq!(deps[0].source_indices, vec![0]);
    assert_eq!(deps[0].target_indices, vec![0, 1, 2]);

    Ok(())
}

With the suggestion, the full sqllogictest suite and the datafusion-common, datafusion-expr and datafusion-optimizer unit tests pass locally.

new_source_indices.push(idx);
new_source_field_names.push(group_by_expr_name.clone());
new_source_input_indices.push(Some(*input_idx));
}
}
let new_target_indices = get_target_functional_dependencies(
aggr_input_schema,
&new_source_field_names,
&new_source_input_indices,
);
let mode = if existing_target_indices == new_target_indices
&& new_target_indices.is_some()
Expand All @@ -513,7 +515,7 @@ pub fn aggregate_functional_dependencies(
// GROUP BY treats NULLs as equal: a determinant covering the
// complete grouping key gets at most one output row per NULL too.
let output_null_equality =
if new_source_indices.len() == group_by_expr_names.len() {
if new_source_indices.len() == group_by_input_indices.len() {
NullEquality::NullEqualsNull
} else {
*null_equality
Expand All @@ -533,8 +535,8 @@ pub fn aggregate_functional_dependencies(

// When we have a GROUP BY key, we can guarantee uniqueness after
// aggregation:
if !group_by_expr_names.is_empty() {
let count = group_by_expr_names.len();
if !group_by_input_indices.is_empty() {
let count = group_by_input_indices.len();
let source_indices = (0..count).collect::<Vec<_>>();
let nullable = source_indices
.iter()
Expand Down Expand Up @@ -566,32 +568,30 @@ pub fn aggregate_functional_dependencies(

/// Returns target indices, for the determinant keys that are inside
/// group by expressions.
///
/// `group_by_input_indices` holds, for each GROUP BY expression, the index of
/// the `schema` field it references, or `None` for a computed expression.
pub fn get_target_functional_dependencies(
schema: &DFSchema,
group_by_expr_names: &[String],
group_by_input_indices: &[Option<usize>],
) -> Option<Vec<usize>> {
let dependencies = schema.functional_dependencies();
if dependencies.is_empty() {
return None;
}
let mut combined_target_indices = HashSet::new();
let field_names = schema.field_names();
for FunctionalDependence {
source_indices,
target_indices,
..
} in &dependencies.deps
{
let source_key_names = source_indices
.iter()
.map(|id_key_idx| &field_names[*id_key_idx])
.collect::<Vec<_>>();
// If the GROUP BY expression contains a determinant key, we can use
// the associated fields after aggregation even if they are not part
// of the GROUP BY expression.
if source_key_names
if source_indices
.iter()
.all(|source_key_name| group_by_expr_names.contains(source_key_name))
.all(|source_idx| group_by_input_indices.contains(&Some(*source_idx)))
{
combined_target_indices.extend(target_indices.iter());
}
Expand All @@ -605,19 +605,18 @@ pub fn get_target_functional_dependencies(

/// Returns indices for the minimal subset of GROUP BY expressions that are
/// functionally equivalent to the original set of GROUP BY expressions.
///
/// `group_by_input_indices` holds, for each GROUP BY expression, the index of
/// the `schema` field it references, or `None` for a computed expression. If
/// any GROUP BY expression is computed, returns `None`.
pub fn get_required_group_by_exprs_indices(
schema: &DFSchema,
group_by_expr_names: &[String],
group_by_input_indices: &[Option<usize>],
) -> Option<Vec<usize>> {
let dependencies = schema.functional_dependencies();
let field_names = schema.field_names();
let mut groupby_expr_indices = group_by_expr_names
let mut groupby_expr_indices = group_by_input_indices
.iter()
.map(|group_by_expr_name| {
field_names
.iter()
.position(|field_name| field_name == group_by_expr_name)
})
.copied()
.collect::<Option<Vec<_>>>()?;

groupby_expr_indices.sort_unstable();
Expand All @@ -642,33 +641,32 @@ pub fn get_required_group_by_exprs_indices(
groupby_expr_indices
.iter()
.map(|idx| {
group_by_expr_names
group_by_input_indices
.iter()
.position(|name| &field_names[*idx] == name)
.position(|input_idx| *input_idx == Some(*idx))
})
.collect()
}

/// Returns indices for the minimal subset of ORDER BY expressions that are
/// functionally equivalent to the original set of ORDER BY expressions.
///
/// `sort_input_indices` holds, for each ORDER BY expression, the index of the
/// `schema` field it references, or `None` for a computed expression.
pub fn get_required_sort_exprs_indices(
schema: &DFSchema,
sort_expr_names: &[String],
sort_input_indices: &[Option<usize>],
) -> Vec<usize> {
let dependencies = schema.functional_dependencies();
let field_names = schema.field_names();

let mut known_field_indices = HashSet::new();
let mut required_sort_expr_indices = Vec::new();

for (sort_expr_idx, sort_expr_name) in sort_expr_names.iter().enumerate() {
// If the sort expression doesn't correspond to a known schema field
// (e.g. a computed expression), we can't reason about it via functional
for (sort_expr_idx, field_idx) in sort_input_indices.iter().enumerate() {
// If the sort expression doesn't reference a schema field (e.g. a
// computed expression), we can't reason about it via functional
// dependencies, so conservatively keep it.
let Some(field_idx) = field_names
.iter()
.position(|field_name| field_name == sort_expr_name)
else {
let Some(field_idx) = *field_idx else {
required_sort_expr_indices.push(sort_expr_idx);
continue;
};
Expand Down
20 changes: 9 additions & 11 deletions datafusion/expr/src/logical_plan/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ use crate::utils::{
Columnizer, can_hash, check_all_columns_from_schema, compare_sort_expr,
expand_qualified_wildcard, expand_wildcard, expr_to_columns,
find_valid_equijoin_key_pair, group_window_expr_by_sort_keys,
split_conjunction_owned,
passthrough_field_index, split_conjunction_owned,
};
use crate::{
BinaryExpr, DmlStatement, ExplainOption, Expr, ExprSchemable, Operator,
Expand Down Expand Up @@ -1988,22 +1988,20 @@ pub fn add_group_by_exprs_from_dependencies(
return Ok(group_expr);
}

// Names of the fields produced by the GROUP BY exprs for example, `GROUP BY
// c1 + 1` produces an output field named `"c1 + 1"`
let mut group_by_field_names = group_expr
// The input field that each GROUP BY expression passes through, or `None`
// for a computed expression such as `c1 + 1`
let mut group_by_input_indices = group_expr
.iter()
.map(|e| e.schema_name().to_string())
.map(|e| passthrough_field_index(e, schema))
.collect::<Vec<_>>();

if let Some(target_indices) =
get_target_functional_dependencies(schema, &group_by_field_names)
get_target_functional_dependencies(schema, &group_by_input_indices)
{
for idx in target_indices {
let expr = Expr::Column(Column::from(schema.qualified_field(idx)));
let expr_name = expr.schema_name().to_string();
if !group_by_field_names.contains(&expr_name) {
group_by_field_names.push(expr_name);
group_expr.push(expr);
if !group_by_input_indices.contains(&Some(idx)) {
group_by_input_indices.push(Some(idx));
group_expr.push(Expr::Column(Column::from(schema.qualified_field(idx))));
}
}
}
Expand Down
59 changes: 22 additions & 37 deletions datafusion/expr/src/logical_plan/plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ use crate::utils::{
check_aggregate_and_window_nesting, check_no_window_functions,
enumerate_grouping_sets, expr_to_columns, exprlist_to_fields,
find_out_reference_exprs, grouping_set_expr_count, grouping_set_to_exprlist,
merge_schema, split_conjunction,
merge_schema, passthrough_field_index, split_conjunction,
};
use crate::{
BinaryExpr, CreateMemoryTable, CreateView, Execute, Expr, ExprSchemable, GroupingSet,
Expand All @@ -69,7 +69,6 @@ use datafusion_common::{
aggregate_functional_dependencies, assert_eq_or_internal_err, assert_or_internal_err,
internal_err, plan_datafusion_err, plan_err, validate_range_split_points,
};
use indexmap::IndexSet;
use itertools::Itertools as _;

// backwards compatibility
Expand Down Expand Up @@ -4378,15 +4377,14 @@ fn calc_func_dependencies_for_aggregate(
// that GROUP BY expression results will be unique.
// - Otherwise, it may be possible to propagate functional dependencies.
if !contains_grouping_set(group_expr) {
let group_by_expr_names = group_expr
.iter()
.map(|item| item.schema_name().to_string())
.collect::<IndexSet<_>>()
// One entry per GROUP BY output field, in the order of `aggr_schema`
let group_by_input_indices = grouping_set_to_exprlist(group_expr)?
.into_iter()
.map(|item| passthrough_field_index(item, input.schema()))
.collect::<Vec<_>>();
let aggregate_func_dependencies = aggregate_functional_dependencies(
input.schema(),
&group_by_expr_names,
&group_by_input_indices,
aggr_schema,
);
Ok(aggregate_func_dependencies)
Expand All @@ -4412,19 +4410,9 @@ fn calc_func_dependencies_for_project(
return Ok(FunctionalDependencies::empty());
}

// Map each input field name to its first index so that projection
// expressions resolve with a hash lookup instead of a linear scan.
let input_fields = input.schema().field_names();
let mut input_index_by_name: HashMap<&str, usize> =
HashMap::with_capacity(input_fields.len());
for (index, name) in input_fields.iter().enumerate() {
input_index_by_name.entry(name.as_str()).or_insert(index);
}
let input_index = |name: &str| {
input_index_by_name
.get(name)
.copied()
.unwrap_or(COMPUTED_EXPR_INDEX)
let input_schema = input.schema();
let input_index = |expr: &Expr| {
passthrough_field_index(expr, input_schema).unwrap_or(COMPUTED_EXPR_INDEX)
};

// Map each projection output position to its input column index.
Expand All @@ -4445,16 +4433,14 @@ fn calc_func_dependencies_for_project(
wildcard_fields
.into_iter()
.map(|(qualifier, f)| {
let flat_name = qualifier
.map(|t| format!("{}.{}", t, f.name()))
.unwrap_or_else(|| f.name().clone());
input_index(&flat_name)
input_schema
.index_of_column_by_name(qualifier.as_ref(), f.name())
.unwrap_or(COMPUTED_EXPR_INDEX)
})
.collect::<Vec<_>>(),
)
}
Expr::Alias(alias) => Ok(vec![input_index(&format!("{}", alias.expr))]),
_ => Ok(vec![input_index(&format!("{expr}"))]),
_ => Ok(vec![input_index(expr)]),
})
.collect::<Result<Vec<_>>>()?
.into_iter()
Expand Down Expand Up @@ -5557,16 +5543,12 @@ mod tests {
}

#[test]
fn projection_duplicate_flattened_name_uses_first_input_index() -> Result<()> {
fn projection_resolves_columns_not_flattened_names() -> Result<()> {
// Build an input schema where a qualified field (`orders`.`id`) and an
// unqualified field that is literally named `"orders.id"` flatten to
// the exact same lookup key that `calc_func_dependencies_for_project`
// uses to resolve projection expressions against input fields. This is
// the only way two entries of `DFSchema::field_names()` can collide
// (`DFSchema::check_names` otherwise forbids duplicate names), and it
// pins that the hash-map based lookup resolves such a collision to the
// *first* matching index, exactly like the linear `position()` scan it
// replaces.
// unqualified field that is literally named `"orders.id"` (the name of
// e.g. `CAST(orders.id AS INT)`) have the same flattened name. The
// projection must resolve the column it references, not the first
// field with the same flattened name.
let schema = DFSchema::new_with_metadata(
vec![
(
Expand All @@ -5589,11 +5571,14 @@ mod tests {
schema: Arc::new(schema),
});

// References the *unqualified* second field, whose flattened name
// ("orders.id") collides with the first (qualified) field's.
// References the *unqualified* second field, which is not a key.
let exprs = vec![Expr::Column(Column::new_unqualified("orders.id"))];
let deps = calc_func_dependencies_for_project(&exprs, &input)?;
assert!(deps.is_empty());

// References the qualified first field, which is a key.
let exprs = vec![Expr::Column(Column::new(Some("orders"), "id"))];
let deps = calc_func_dependencies_for_project(&exprs, &input)?;
assert_eq!(deps.len(), 1);
assert_eq!(deps[0].source_indices, vec![0]);

Expand Down
Loading
Loading