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
Original file line number Diff line number Diff line change
Expand Up @@ -36,13 +36,13 @@
import org.apache.flink.table.types.DataType;
import org.apache.flink.table.types.logical.LogicalType;
import org.apache.flink.table.types.logical.LogicalTypeFamily;
import org.apache.flink.table.types.logical.LogicalTypeRoot;
import org.apache.flink.table.types.logical.RowType;

import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Deque;
import java.util.List;
import java.util.Objects;
import java.util.Optional;
import java.util.function.BiFunction;
import java.util.regex.Matcher;
Expand Down Expand Up @@ -74,74 +74,93 @@ public PredicateConverter(PredicateBuilder builder) {

@Override
public Predicate visit(CallExpression call) {
return visit(call, false);
}

private Predicate visit(CallExpression call, boolean negated) {
FunctionDefinition func = call.getFunctionDefinition();
List<Expression> children = call.getChildren();

if (func == BuiltInFunctionDefinitions.AND) {
return PredicateBuilder.and(flattenAndConvert(children, func));
requireAtLeastArity(children, 2);
List<Predicate> predicates = flattenAndConvert(children, func, negated);
return negated ? PredicateBuilder.or(predicates) : PredicateBuilder.and(predicates);
} else if (func == BuiltInFunctionDefinitions.OR) {
return PredicateBuilder.or(flattenAndConvert(children, func));
requireAtLeastArity(children, 2);
List<Predicate> predicates = flattenAndConvert(children, func, negated);
return negated ? PredicateBuilder.and(predicates) : PredicateBuilder.or(predicates);
} else if (func == BuiltInFunctionDefinitions.NOT) {
requireArity(children, 1);
return visit(children.get(0), !negated);
} else if (func == BuiltInFunctionDefinitions.EQUALS) {
return visitBiFunction(children, builder::equal, builder::equal);
return negated
? visitBiFunction(children, builder::notEqual, builder::notEqual)
: visitBiFunction(children, builder::equal, builder::equal);
} else if (func == BuiltInFunctionDefinitions.NOT_EQUALS) {
return visitBiFunction(children, builder::notEqual, builder::notEqual);
return negated
? visitBiFunction(children, builder::equal, builder::equal)
: visitBiFunction(children, builder::notEqual, builder::notEqual);
} else if (func == BuiltInFunctionDefinitions.GREATER_THAN) {

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.

[P1] Preserve Flink NaN semantics for negated comparisons

For FLOAT and DOUBLE, this rewrite is not equivalent to the original expression. Flink numeric comparisons use Java operators, so NOT (NaN > 1.0) evaluates to true, while the Paimon LessOrEqual predicate orders values through Double.compare and rejects NaN. Because this predicate is pushed down before the remaining Flink filter runs, the row is discarded and the query returns incomplete results.

Please keep negated floating-point comparisons unsupported/residual, or construct NaN-aware equivalents for every comparison direction and add row/source tests containing NaN.

return visitBiFunction(children, builder::greaterThan, builder::lessThan);
return negated
? visitBiFunction(children, builder::lessOrEqual, builder::greaterOrEqual)
: visitBiFunction(children, builder::greaterThan, builder::lessThan);
} else if (func == BuiltInFunctionDefinitions.GREATER_THAN_OR_EQUAL) {
return visitBiFunction(children, builder::greaterOrEqual, builder::lessOrEqual);
return negated
? visitBiFunction(children, builder::lessThan, builder::greaterThan)
: visitBiFunction(children, builder::greaterOrEqual, builder::lessOrEqual);
} else if (func == BuiltInFunctionDefinitions.LESS_THAN) {
return visitBiFunction(children, builder::lessThan, builder::greaterThan);
return negated
? visitBiFunction(children, builder::greaterOrEqual, builder::lessOrEqual)
: visitBiFunction(children, builder::lessThan, builder::greaterThan);
} else if (func == BuiltInFunctionDefinitions.LESS_THAN_OR_EQUAL) {
return visitBiFunction(children, builder::lessOrEqual, builder::greaterOrEqual);
return negated
? visitBiFunction(children, builder::greaterThan, builder::lessThan)
: visitBiFunction(children, builder::lessOrEqual, builder::greaterOrEqual);
} else if (func == BuiltInFunctionDefinitions.IN) {
FieldReferenceExpression fieldRefExpr =
extractFieldReference(children.get(0)).orElseThrow(UnsupportedExpression::new);
requireAtLeastArity(children, 2);
ResolvedField field = resolveField(children.get(0));
List<Object> literals = new ArrayList<>();
for (int i = 1; i < children.size(); i++) {
literals.add(extractLiteral(fieldRefExpr.getOutputDataType(), children.get(i)));
literals.add(extractLiteral(field.expression.getOutputDataType(), children.get(i)));
}
return builder.in(builder.indexOf(fieldRefExpr.getName()), literals);
return negated
? builder.notIn(field.index, literals)

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.

[P2] Short-circuit NOT IN lists containing NULL

Under SQL WHERE semantics, v NOT IN (1, NULL, 3) can never be true. Passing the NULL literal to builder.notIn is also unsafe for file-index evaluation: the BSI reader can unbox a null mapped value, and the range-bitmap reader can pass null to its comparator, causing the query to fail when either index is enabled. Before this change, the unsupported expression stayed as a Flink residual filter.

Please return an always-false predicate when a negated IN list contains NULL, or reject the conversion so it remains residual. An index-enabled regression test would cover both paths.

: builder.in(field.index, literals);
} else if (func == BuiltInFunctionDefinitions.IS_NULL) {
return extractFieldReference(children.get(0))
.map(FieldReferenceExpression::getName)
.map(builder::indexOf)
.map(builder::isNull)
.orElseThrow(UnsupportedExpression::new);
requireArity(children, 1);
ResolvedField field = resolveField(children.get(0));
return negated ? builder.isNotNull(field.index) : builder.isNull(field.index);
} else if (func == BuiltInFunctionDefinitions.IS_NOT_NULL) {
return extractFieldReference(children.get(0))
.map(FieldReferenceExpression::getName)
.map(builder::indexOf)
.map(builder::isNotNull)
.orElseThrow(UnsupportedExpression::new);
requireArity(children, 1);
ResolvedField field = resolveField(children.get(0));
return negated ? builder.isNull(field.index) : builder.isNotNull(field.index);
} else if (func == BuiltInFunctionDefinitions.BETWEEN) {
FieldReferenceExpression fieldRefExpr =
extractFieldReference(children.get(0)).orElseThrow(UnsupportedExpression::new);
DataType fieldType = fieldRefExpr.getOutputDataType();
return builder.between(
builder.indexOf(fieldRefExpr.getName()),
extractLiteral(fieldType, children.get(1)),
extractLiteral(fieldType, children.get(2)));
requireArity(children, 3);
ResolvedField field = resolveField(children.get(0));
Object lower = extractLiteral(field.expression.getOutputDataType(), children.get(1));
Object upper = extractLiteral(field.expression.getOutputDataType(), children.get(2));
Predicate between = builder.between(field.index, lower, upper);
return negated ? negate(between) : between;
} else if (func == BuiltInFunctionDefinitions.LIKE) {
FieldReferenceExpression fieldRefExpr =
extractFieldReference(children.get(0)).orElseThrow(UnsupportedExpression::new);
if (fieldRefExpr
if (children.size() != 2 && children.size() != 3) {
throw new UnsupportedExpression();
}
ResolvedField field = resolveField(children.get(0));
if (field.expression
.getOutputDataType()
.getLogicalType()
.getTypeRoot()
.getFamilies()
.contains(LogicalTypeFamily.CHARACTER_STRING)) {
String sqlPattern =
Objects.requireNonNull(
extractLiteral(
fieldRefExpr.getOutputDataType(), children.get(1)))
extractNonNullLiteral(field.expression.getOutputDataType(), children.get(1))
.toString();
String escape =
children.size() <= 2
? null
: Objects.requireNonNull(
extractLiteral(
fieldRefExpr.getOutputDataType(),
children.get(2)))
: extractNonNullLiteral(
field.expression.getOutputDataType(),
children.get(2))
.toString();
String escapedSqlPattern = sqlPattern;
boolean allowQuick = false;
Expand Down Expand Up @@ -185,38 +204,69 @@ public Predicate visit(CallExpression call) {
if (allowQuick) {
Matcher beginMatcher = BEGIN_PATTERN.matcher(escapedSqlPattern);
if (beginMatcher.matches()) {
if (negated) {
// StartsWith has no negated predicate, so NOT LIKE must remain a
// residual filter evaluated by Flink.
throw new UnsupportedExpression();
}
return builder.startsWith(
builder.indexOf(fieldRefExpr.getName()),
BinaryString.fromString(beginMatcher.group(1)));
field.index, BinaryString.fromString(beginMatcher.group(1)));
}
}
}
} else if (func == BuiltInFunctionDefinitions.IS_TRUE) {
FieldReferenceExpression fieldRefExpr =
extractFieldReference(children.get(0)).orElseThrow(UnsupportedExpression::new);
return builder.equal(builder.indexOf(fieldRefExpr.getName()), Boolean.TRUE);
requireArity(children, 1);
return booleanTest(resolveField(children.get(0)), true, negated);
} else if (func == BuiltInFunctionDefinitions.IS_FALSE) {
FieldReferenceExpression fieldRefExpr =
extractFieldReference(children.get(0)).orElseThrow(UnsupportedExpression::new);
return builder.equal(builder.indexOf(fieldRefExpr.getName()), Boolean.FALSE);
requireArity(children, 1);
return booleanTest(resolveField(children.get(0)), false, negated);
} else if (func == BuiltInFunctionDefinitions.IS_NOT_TRUE) {
requireArity(children, 1);
return booleanTest(resolveField(children.get(0)), true, !negated);
} else if (func == BuiltInFunctionDefinitions.IS_NOT_FALSE) {
requireArity(children, 1);
return booleanTest(resolveField(children.get(0)), false, !negated);
}

// TODO is_xxx, between_xxx, similar, in, not_in, not?
throw new UnsupportedExpression();
}

private Predicate visit(Expression expression, boolean negated) {
if (expression instanceof CallExpression) {
return visit((CallExpression) expression, negated);
}
throw new UnsupportedExpression();
}

private Predicate booleanTest(ResolvedField field, boolean expected, boolean complement) {
if (field.expression.getOutputDataType().getLogicalType().getTypeRoot()
!= LogicalTypeRoot.BOOLEAN) {
throw new UnsupportedExpression();
}
Predicate equals = builder.equal(field.index, expected);
if (!complement) {
return equals;
}
return PredicateBuilder.or(
builder.isNull(field.index), builder.notEqual(field.index, expected));
}

private Predicate negate(Predicate predicate) {
return predicate.negate().orElseThrow(UnsupportedExpression::new);
}

/**
* Iteratively flattens a nested AND/OR expression tree into a flat list of child predicates,
* avoiding stack overflow caused by recursive {@code accept} calls on deeply nested trees (e.g.
* when Flink expands a large IN clause into nested OR expressions).
*
* @param children the children of the top-level AND/OR {@link CallExpression}
* @param targetFunc the function definition to flatten ({@code AND} or {@code OR})
* @param negated whether to negate every flattened child and combine them using De Morgan's law
* @return a flat list of converted child predicates in original order
*/
private List<Predicate> flattenAndConvert(
List<Expression> children, FunctionDefinition targetFunc) {
List<Expression> children, FunctionDefinition targetFunc, boolean negated) {
List<Predicate> result = new ArrayList<>();
Deque<Expression> stack = new ArrayDeque<>();
for (int i = children.size() - 1; i >= 0; i--) {
Expand All @@ -228,14 +278,15 @@ private List<Predicate> flattenAndConvert(
CallExpression ce = (CallExpression) expr;
if (ce.getFunctionDefinition() == targetFunc) {
List<Expression> ceChildren = ce.getChildren();
requireAtLeastArity(ceChildren, 2);
for (int i = ceChildren.size() - 1; i >= 0; i--) {
stack.push(ceChildren.get(i));
}
} else {
result.add(ce.accept(this));
result.add(visit(ce, negated));
}
} else {
result.add(expr.accept(this));
result.add(visit(expr, negated));
}
}
return result;
Expand All @@ -245,23 +296,52 @@ private Predicate visitBiFunction(
List<Expression> children,
BiFunction<Integer, Object, Predicate> visit1,
BiFunction<Integer, Object, Predicate> visit2) {
requireArity(children, 2);
Optional<FieldReferenceExpression> fieldRefExpr = extractFieldReference(children.get(0));
if (fieldRefExpr.isPresent()) {
int fieldIndex = resolveFieldIndex(fieldRefExpr.get());
Object literal =
extractLiteral(fieldRefExpr.get().getOutputDataType(), children.get(1));
return visit1.apply(builder.indexOf(fieldRefExpr.get().getName()), literal);
return visit1.apply(fieldIndex, literal);
} else {
fieldRefExpr = extractFieldReference(children.get(1));
if (fieldRefExpr.isPresent()) {
int fieldIndex = resolveFieldIndex(fieldRefExpr.get());
Object literal =
extractLiteral(fieldRefExpr.get().getOutputDataType(), children.get(0));
return visit2.apply(builder.indexOf(fieldRefExpr.get().getName()), literal);
return visit2.apply(fieldIndex, literal);
}
}

throw new UnsupportedExpression();
}

private ResolvedField resolveField(Expression expression) {
FieldReferenceExpression field =
extractFieldReference(expression).orElseThrow(UnsupportedExpression::new);
return new ResolvedField(field, resolveFieldIndex(field));
}

private int resolveFieldIndex(FieldReferenceExpression field) {
int index = builder.indexOf(field.getName());
if (index < 0) {
throw new UnsupportedExpression();
}
return index;
}

private void requireArity(List<Expression> children, int expected) {
if (children.size() != expected) {
throw new UnsupportedExpression();
}
}

private void requireAtLeastArity(List<Expression> children, int minimum) {
if (children.size() < minimum) {
throw new UnsupportedExpression();
}
}

private Optional<FieldReferenceExpression> extractFieldReference(Expression expression) {
if (expression instanceof FieldReferenceExpression) {
return Optional.of((FieldReferenceExpression) expression);
Expand Down Expand Up @@ -304,6 +384,14 @@ private Object extractLiteral(DataType expectedType, Expression expression) {
throw new UnsupportedExpression();
}

private Object extractNonNullLiteral(DataType expectedType, Expression expression) {
Object literal = extractLiteral(expectedType, expression);
if (literal == null) {
throw new UnsupportedExpression();
}
return literal;
}

private boolean supportsPredicate(LogicalType type) {
switch (type.getTypeRoot()) {
case CHAR:
Expand Down Expand Up @@ -331,6 +419,17 @@ private boolean supportsPredicate(LogicalType type) {
}
}

private static class ResolvedField {

private final FieldReferenceExpression expression;
private final int index;

private ResolvedField(FieldReferenceExpression expression, int index) {
this.expression = expression;
this.index = index;
}
}

@Override
public Predicate visit(ValueLiteralExpression valueLiteralExpression) {
throw new UnsupportedExpression();
Expand Down
Loading
Loading