diff --git a/pkg/sql/colexec/evalExpression.go b/pkg/sql/colexec/evalExpression.go index bbaa366e253e6..307f0078e93ff 100644 --- a/pkg/sql/colexec/evalExpression.go +++ b/pkg/sql/colexec/evalExpression.go @@ -231,12 +231,29 @@ type FixedVectorExpressionExecutor struct { type FunctionExpressionExecutor struct { m *mpool.MPool + // resultType is the declared function return type. Some built-ins refine + // result metadata (for example temporal scale or decimal width/scale) at + // runtime, so reusable result vectors must start each evaluation from this + // stable type before the function applies the current runtime metadata. + resultType types.Type functionInformationForEval folded functionFolding selectList1 []bool selectList2 []bool selectList function.FunctionSelectList + // A function implementation cannot be required to interpret selectList + // correctly: many built-ins predate it, and evaluating them on masked rows + // can still raise errors or perform side effects. For a partial selection we + // therefore compact row-aligned parameters, evaluate only selected rows, and + // scatter the result back to the original row positions. These buffers are + // allocated lazily and reused across batches. + selectedRows []int64 + selectedParameterResults []*vector.Vector + selectedParameterVectors []*vector.Vector + selectedResult vector.FunctionResultWrapper + selectedNullResult *vector.Vector + resultVector vector.FunctionResultWrapper // parameters related parameterResults []*vector.Vector @@ -267,14 +284,23 @@ func (expr *ColumnExpressionExecutor) GetColIndex() int { type ParamExpressionExecutor struct { mp *mpool.MPool null *vector.Vector - vec *vector.Vector - pos int - typ types.Type + // maskedNull is separate from null/vec because it is not a resolved + // parameter value and must never participate in the folded-value cache. + maskedNull *vector.Vector + vec *vector.Vector + pos int + typ types.Type folded bool } -func (expr *ParamExpressionExecutor) Eval(proc *process.Process, _ []*batch.Batch, _ []bool) (*vector.Vector, error) { +func (expr *ParamExpressionExecutor) Eval(proc *process.Process, batches []*batch.Batch, selectList []bool) (*vector.Vector, error) { + if noRowsSelected(selectList, expressionRowCount(batches)) { + if expr.maskedNull == nil { + expr.maskedNull = vector.NewConstNull(expr.typ, 1, proc.GetMPool()) + } + return expr.maskedNull, nil + } if expr.folded { if expr.null != nil { return expr.null, nil @@ -331,6 +357,10 @@ func (expr *ParamExpressionExecutor) Free() { expr.null.Free(expr.mp) expr.null = nil } + if expr.maskedNull != nil { + expr.maskedNull.Free(expr.mp) + expr.maskedNull = nil + } reuse.Free[ParamExpressionExecutor](expr, nil) } @@ -341,7 +371,10 @@ func (expr *ParamExpressionExecutor) IsColumnExpr() bool { type VarExpressionExecutor struct { mp *mpool.MPool null *vector.Vector - vec *vector.Vector + // maskedNull lets a skipped variable avoid the resolver without changing + // the value cache used by a later selected evaluation. + maskedNull *vector.Vector + vec *vector.Vector name string system bool @@ -349,7 +382,13 @@ type VarExpressionExecutor struct { typ types.Type } -func (expr *VarExpressionExecutor) Eval(proc *process.Process, batches []*batch.Batch, _ []bool) (*vector.Vector, error) { +func (expr *VarExpressionExecutor) Eval(proc *process.Process, batches []*batch.Batch, selectList []bool) (*vector.Vector, error) { + if noRowsSelected(selectList, expressionRowCount(batches)) { + if expr.maskedNull == nil { + expr.maskedNull = vector.NewConstNull(expr.typ, 1, proc.GetMPool()) + } + return expr.maskedNull, nil + } resolveVariableFunc := proc.GetResolveVariableFunc() if resolveVariableFunc == nil { return nil, moerr.NewInternalErrorf(proc.Ctx, "resolve variable function is not set for variable %s", expr.name) @@ -406,6 +445,10 @@ func (expr *VarExpressionExecutor) Free() { expr.null.Free(expr.mp) expr.null = nil } + if expr.maskedNull != nil { + expr.maskedNull.Free(expr.mp) + expr.maskedNull = nil + } reuse.Free[VarExpressionExecutor](expr, nil) } @@ -496,6 +539,7 @@ func (expr *FunctionExpressionExecutor) Init( m := proc.Mp() expr.m = m + expr.resultType = retType expr.parameterResults = make([]*vector.Vector, parameterNum) expr.parameterExecutor = make([]ExpressionExecutor, parameterNum) @@ -503,56 +547,77 @@ func (expr *FunctionExpressionExecutor) Init( return err } +func (expr *FunctionExpressionExecutor) resetResultType(result vector.FunctionResultWrapper) { + if result == nil { + return + } + if vec := result.GetResultVector(); vec != nil { + vec.SetType(expr.resultType) + vec.SetIsBin(false) + } +} + +func expressionRowCount(batches []*batch.Batch) int { + if len(batches) > 0 { + return batches[0].RowCount() + } + return 1 +} + func (expr *FunctionExpressionExecutor) EvalIff(proc *process.Process, batches []*batch.Batch, selectList []bool) (err error) { expr.parameterResults[0], err = expr.parameterExecutor[0].Eval(proc, batches, selectList) if err != nil { return err } - rowCount := batches[0].RowCount() + rowCount := expressionRowCount(batches) if len(expr.selectList1) < rowCount { expr.selectList1 = make([]bool, rowCount) expr.selectList2 = make([]bool, rowCount) } + trueBranch := expr.selectList1[:rowCount] + falseBranch := expr.selectList2[:rowCount] bs := vector.GenerateFunctionFixedTypeParameter[bool](expr.parameterResults[0]) for i := 0; i < rowCount; i++ { b, null := bs.GetValue(uint64(i)) if selectList != nil { - expr.selectList1[i] = selectList[i] - expr.selectList2[i] = selectList[i] + trueBranch[i] = selectList[i] + falseBranch[i] = selectList[i] } else { - expr.selectList1[i] = true - expr.selectList2[i] = true + trueBranch[i] = true + falseBranch[i] = true } if !null && b { - expr.selectList2[i] = false + falseBranch[i] = false } else { - expr.selectList1[i] = false + trueBranch[i] = false } } - expr.parameterResults[1], err = expr.parameterExecutor[1].Eval(proc, batches, expr.selectList1) + expr.parameterResults[1], err = expr.parameterExecutor[1].Eval(proc, batches, trueBranch) if err != nil { return err } - expr.parameterResults[2], err = expr.parameterExecutor[2].Eval(proc, batches, expr.selectList2) + expr.parameterResults[2], err = expr.parameterExecutor[2].Eval(proc, batches, falseBranch) return err } func (expr *FunctionExpressionExecutor) EvalCase(proc *process.Process, batches []*batch.Batch, selectList []bool) (err error) { - rowCount := batches[0].RowCount() + rowCount := expressionRowCount(batches) if len(expr.selectList1) < rowCount { expr.selectList1 = make([]bool, rowCount) expr.selectList2 = make([]bool, rowCount) } + remaining := expr.selectList1[:rowCount] + selectedBranch := expr.selectList2[:rowCount] if selectList != nil { - copy(expr.selectList1, selectList) + copy(remaining, selectList) } else { - for i := range expr.selectList1 { - expr.selectList1[i] = true + for i := range remaining { + remaining[i] = true } } for i := 0; i < len(expr.parameterExecutor); i += 2 { - expr.parameterResults[i], err = expr.parameterExecutor[i].Eval(proc, batches, expr.selectList1) + expr.parameterResults[i], err = expr.parameterExecutor[i].Eval(proc, batches, remaining) if err != nil { return err } @@ -561,14 +626,14 @@ func (expr *FunctionExpressionExecutor) EvalCase(proc *process.Process, batches for j := 0; j < rowCount; j++ { b, null := bs.GetValue(uint64(j)) - if !null && b { - expr.selectList1[j] = false - expr.selectList2[j] = true + if remaining[j] && !null && b { + remaining[j] = false + selectedBranch[j] = true } else { - expr.selectList2[j] = false + selectedBranch[j] = false } } - expr.parameterResults[i+1], err = expr.parameterExecutor[i+1].Eval(proc, batches, expr.selectList2) + expr.parameterResults[i+1], err = expr.parameterExecutor[i+1].Eval(proc, batches, selectedBranch) if err != nil { return err } @@ -577,7 +642,161 @@ func (expr *FunctionExpressionExecutor) EvalCase(proc *process.Process, batches return err } +func (expr *FunctionExpressionExecutor) EvalCoalesce(proc *process.Process, batches []*batch.Batch, selectList []bool) (err error) { + rowCount := expressionRowCount(batches) + if len(expr.selectList1) < rowCount { + expr.selectList1 = make([]bool, rowCount) + } + remaining := expr.selectList1[:rowCount] + if selectList != nil { + for i := range remaining { + remaining[i] = i < len(selectList) && selectList[i] + } + } else { + for i := range remaining { + remaining[i] = true + } + } + + for i := range expr.parameterExecutor { + expr.parameterResults[i], err = expr.parameterExecutor[i].Eval(proc, batches, remaining) + if err != nil { + return err + } + for row := range remaining { + if remaining[row] && !expr.parameterResults[i].IsNull(uint64(row)) { + remaining[row] = false + } + } + } + return nil +} + +func noRowsSelected(selectList []bool, rowCount int) bool { + if selectList == nil { + return false + } + if len(selectList) < rowCount { + return false + } + for i := 0; i < rowCount; i++ { + if selectList[i] { + return false + } + } + return true +} + +func (expr *FunctionExpressionExecutor) makeNullResult(rowCount int) (*vector.Vector, error) { + expr.resetResultType(expr.resultVector) + if err := expr.resultVector.PreExtendAndReset(rowCount); err != nil { + return nil, err + } + result := expr.resultVector.GetResultVector() + result.SetAllNulls(rowCount) + result.SetLength(rowCount) + return result, nil +} + +func (expr *FunctionExpressionExecutor) evalSelectedRows( + proc *process.Process, + rowCount int, + selectList []bool, +) (*vector.Vector, error) { + expr.selectedRows = expr.selectedRows[:0] + for row := 0; row < rowCount; row++ { + if selectList[row] { + expr.selectedRows = append(expr.selectedRows, int64(row)) + } + } + + selectedCount := len(expr.selectedRows) + if len(expr.selectedParameterResults) == 0 && len(expr.parameterResults) > 0 { + expr.selectedParameterResults = make([]*vector.Vector, len(expr.parameterResults)) + expr.selectedParameterVectors = make([]*vector.Vector, len(expr.parameterResults)) + } + for i, parameter := range expr.parameterResults { + // Constants, folded vectors, and list/vector literals are not row-aligned. + // They must be passed through unchanged; only column and non-folded + // function results map one-to-one to the input batch rows. + rowAligned := false + switch executor := expr.parameterExecutor[i].(type) { + case *ColumnExpressionExecutor: + rowAligned = true + case *FunctionExpressionExecutor: + rowAligned = !executor.folded.canFold + } + if rowAligned && !parameter.IsConst() { + selected := expr.selectedParameterVectors[i] + if selected == nil { + selected = vector.NewOffHeapVecWithType(*parameter.GetType()) + expr.selectedParameterVectors[i] = selected + } else { + selected.Reset(*parameter.GetType()) + } + selected.SetIsBin(parameter.GetIsBin()) + if err := selected.Union(parameter, expr.selectedRows, proc.Mp()); err != nil { + return nil, err + } + expr.selectedParameterResults[i] = selected + continue + } + expr.selectedParameterResults[i] = parameter + } + + expr.resetResultType(expr.resultVector) + if err := expr.resultVector.PreExtendAndReset(rowCount); err != nil { + return nil, err + } + if expr.selectedResult == nil { + expr.selectedResult = vector.NewFunctionResultWrapper(expr.resultType, expr.m) + } + expr.resetResultType(expr.selectedResult) + if err := expr.selectedResult.PreExtendAndReset(selectedCount); err != nil { + return nil, err + } + if err := expr.evalFn( + expr.selectedParameterResults, expr.selectedResult, proc, selectedCount, nil); err != nil { + return nil, err + } + + selectedResult := expr.selectedResult.GetResultVector() + runtimeType := *selectedResult.GetType() + runtimeIsBin := selectedResult.GetIsBin() + + result := expr.resultVector.GetResultVector() + result.SetType(runtimeType) + result.SetIsBin(runtimeIsBin) + result.ResetWithSameType() + if expr.selectedNullResult == nil { + expr.selectedNullResult = vector.NewConstNull(runtimeType, 1, expr.m) + } else { + expr.selectedNullResult.SetType(runtimeType) + expr.selectedNullResult.SetLength(1) + } + expr.selectedNullResult.SetIsBin(runtimeIsBin) + selectedRow := int64(0) + for row := 0; row < rowCount; row++ { + if selectList[row] { + if err := result.UnionOne(selectedResult, selectedRow, proc.Mp()); err != nil { + return nil, err + } + selectedRow++ + } else if err := result.UnionOne(expr.selectedNullResult, 0, proc.Mp()); err != nil { + return nil, err + } + } + return result, nil +} + func (expr *FunctionExpressionExecutor) Eval(proc *process.Process, batches []*batch.Batch, selectList []bool) (*vector.Vector, error) { + if len(batches) == 0 { + batches = []*batch.Batch{batch.EmptyForConstFoldBatch} + } + rowCount := expressionRowCount(batches) + if !expr.folded.canFold && noRowsSelected(selectList, rowCount) { + return expr.makeNullResult(rowCount) + } if expr.folded.needFoldingCheck { if err := expr.doFold(proc, proc.GetBaseProcessRunningStatus()); err != nil { return nil, err @@ -601,6 +820,11 @@ func (expr *FunctionExpressionExecutor) Eval(proc *process.Process, batches []*b if err != nil { return nil, err } + } else if expr.fid == function.COALESCE { + err = expr.EvalCoalesce(proc, batches, selectList) + if err != nil { + return nil, err + } } else { for i := range expr.parameterExecutor { expr.parameterResults[i], err = expr.parameterExecutor[i].Eval(proc, batches, selectList) @@ -610,12 +834,24 @@ func (expr *FunctionExpressionExecutor) Eval(proc *process.Process, batches []*b } } - if err = expr.resultVector.PreExtendAndReset(batches[0].RowCount()); err != nil { - return nil, err + if selectList != nil { + selectedCount := 0 + for row := 0; row < rowCount; row++ { + if selectList[row] { + selectedCount++ + } + } + if selectedCount < rowCount { + return expr.evalSelectedRows(proc, rowCount, selectList) + } } - if selectList != nil && len(expr.selectList.SelectList) < batches[0].RowCount() { - expr.selectList.SelectList = make([]bool, batches[0].RowCount()) + expr.resetResultType(expr.resultVector) + if err = expr.resultVector.PreExtendAndReset(rowCount); err != nil { + return nil, err + } + if selectList != nil && len(expr.selectList.SelectList) < rowCount { + expr.selectList.SelectList = make([]bool, rowCount) } if selectList == nil { expr.selectList.AnyNull = false @@ -637,7 +873,7 @@ func (expr *FunctionExpressionExecutor) Eval(proc *process.Process, batches []*b } if err = expr.evalFn( - expr.parameterResults, expr.resultVector, proc, batches[0].RowCount(), &expr.selectList); err != nil { + expr.parameterResults, expr.resultVector, proc, rowCount, &expr.selectList); err != nil { return nil, err } @@ -664,6 +900,19 @@ func (expr *FunctionExpressionExecutor) Free() { expr.resultVector.Free() expr.resultVector = nil } + if expr.selectedResult != nil { + expr.selectedResult.Free() + expr.selectedResult = nil + } + if expr.selectedNullResult != nil { + expr.selectedNullResult.Free(expr.m) + expr.selectedNullResult = nil + } + for _, parameter := range expr.selectedParameterVectors { + if parameter != nil { + parameter.Free(expr.m) + } + } for _, p := range expr.parameterExecutor { if p != nil { diff --git a/pkg/sql/colexec/evalExpressionReset.go b/pkg/sql/colexec/evalExpressionReset.go index 7b89ffa75affe..02a9145d20c01 100644 --- a/pkg/sql/colexec/evalExpressionReset.go +++ b/pkg/sql/colexec/evalExpressionReset.go @@ -18,6 +18,7 @@ import ( "context" "github.com/matrixorigin/matrixone/pkg/common/mpool" + "github.com/matrixorigin/matrixone/pkg/container/types" "github.com/matrixorigin/matrixone/pkg/container/vector" "github.com/matrixorigin/matrixone/pkg/sql/plan/function" "github.com/matrixorigin/matrixone/pkg/vm/process" @@ -108,6 +109,150 @@ func (expr *FunctionExpressionExecutor) getFoldedVector(requiredLength int) *vec return rv } +func (expr *FunctionExpressionExecutor) tryFoldParameter( + proc *process.Process, + atRuntime bool, + index int, +) (bool, error) { + parameter := expr.parameterExecutor[index] + if constant, ok := parameter.(*FixedVectorExpressionExecutor); ok { + expr.parameterResults[index] = constant.resultVector + if !constant.noNeedToSetLength { + expr.parameterResults[index].SetLength(1) + } + return true, nil + } + if functionParameter, ok := parameter.(*FunctionExpressionExecutor); ok { + if err := functionParameter.doFold(proc, atRuntime); err != nil { + return false, err + } + if functionParameter.folded.canFold { + expr.parameterResults[index] = functionParameter.getFoldedVector(1) + return true, nil + } + return false, nil + } + if atRuntime { + if parameter, ok := parameter.(*ParamExpressionExecutor); ok { + result, err := parameter.Eval(proc, nil, nil) + if err != nil { + return false, err + } + expr.parameterResults[index] = result + return true, nil + } + } + return false, nil +} + +func (expr *FunctionExpressionExecutor) tryFoldFlowControl( + proc *process.Process, + atRuntime bool, +) (bool, error) { + switch expr.fid { + case function.IFF: + folded, err := expr.tryFoldParameter(proc, atRuntime, 0) + if err != nil || !folded { + return folded, err + } + condition := vector.GenerateFunctionFixedTypeParameter[bool](expr.parameterResults[0]) + value, isNull := condition.GetValue(0) + selected := 2 + if !isNull && value { + selected = 1 + } + return expr.tryFoldParameter(proc, atRuntime, selected) + + case function.CASE: + parameterCount := len(expr.parameterExecutor) + for conditionIndex := 0; conditionIndex+1 < parameterCount; conditionIndex += 2 { + folded, err := expr.tryFoldParameter(proc, atRuntime, conditionIndex) + if err != nil || !folded { + return folded, err + } + condition := vector.GenerateFunctionFixedTypeParameter[bool](expr.parameterResults[conditionIndex]) + value, isNull := condition.GetValue(0) + if !isNull && value { + return expr.tryFoldParameter(proc, atRuntime, conditionIndex+1) + } + } + if parameterCount%2 == 1 { + return expr.tryFoldParameter(proc, atRuntime, parameterCount-1) + } + return true, nil + + case function.COALESCE: + for i := range expr.parameterExecutor { + folded, err := expr.tryFoldParameter(proc, atRuntime, i) + if err != nil || !folded { + return folded, err + } + if !expr.parameterResults[i].IsNull(0) { + return true, nil + } + } + return true, nil + } + return false, nil +} + +func (expr *FunctionExpressionExecutor) fillSkippedFlowControlParameters() func() { + // The registered kernels still receive their complete argument list. Supply + // typed NULLs for branches that lazy folding deliberately did not evaluate; + // the selected conditions make those placeholders unobservable. + var boolNull *vector.Vector + var resultNull *vector.Vector + temporaryIndexes := make([]int, 0, len(expr.parameterResults)) + parameterCount := len(expr.parameterResults) + for i := range expr.parameterResults { + if expr.parameterResults[i] != nil { + continue + } + isCondition := expr.fid == function.IFF && i == 0 + if expr.fid == function.CASE && i%2 == 0 && (parameterCount%2 == 0 || i < parameterCount-1) { + isCondition = true + } + if isCondition { + if boolNull == nil { + boolNull = vector.NewConstNull(types.T_bool.ToType(), 1, expr.m) + } + expr.parameterResults[i] = boolNull + } else { + if resultNull == nil { + resultNull = vector.NewConstNull(expr.resultType, 1, expr.m) + } + expr.parameterResults[i] = resultNull + } + temporaryIndexes = append(temporaryIndexes, i) + } + return func() { + for _, i := range temporaryIndexes { + expr.parameterResults[i] = nil + } + if boolNull != nil { + boolNull.Free(expr.m) + } + if resultNull != nil { + resultNull.Free(expr.m) + } + } +} + +func (expr *FunctionExpressionExecutor) finishFolding(proc *process.Process, execLen int) error { + expr.resetResultType(expr.resultVector) + if err := expr.resultVector.PreExtendAndReset(execLen); err != nil { + return err + } + if err := expr.evalFn(expr.parameterResults, expr.resultVector, proc, execLen, nil); err != nil { + return err + } + if execLen == 1 { + expr.resultVector.GetResultVector().ToConst() + } + expr.folded.canFold = true + return nil +} + func (expr *FunctionExpressionExecutor) doFold(proc *process.Process, atRuntime bool) (err error) { if !expr.folded.needFoldingCheck { return nil @@ -115,40 +260,29 @@ func (expr *FunctionExpressionExecutor) doFold(proc *process.Process, atRuntime expr.folded.needFoldingCheck = false expr.folded.canFold = false + if expr.fid == function.IFF || expr.fid == function.CASE || expr.fid == function.COALESCE { + if expr.volatile || (!atRuntime && expr.timeDependent) { + return nil + } + folded, err := expr.tryFoldFlowControl(proc, atRuntime) + if err != nil || !folded { + return err + } + cleanup := expr.fillSkippedFlowControlParameters() + defer cleanup() + return expr.finishFolding(proc, 1) + } // fold parameters. allParametersFolded := true - for i, param := range expr.parameterExecutor { - // constant expression. - if constant, ok := param.(*FixedVectorExpressionExecutor); ok { - expr.parameterResults[i] = constant.resultVector - if !constant.noNeedToSetLength { - expr.parameterResults[i].SetLength(1) - } - continue + for i := range expr.parameterExecutor { + folded, foldErr := expr.tryFoldParameter(proc, atRuntime, i) + if foldErr != nil { + return foldErr } - // function expression. - if fExpr, ok := param.(*FunctionExpressionExecutor); ok { - paramFoldError := fExpr.doFold(proc, atRuntime) - if paramFoldError != nil { - return err - } - if fExpr.folded.canFold { - expr.parameterResults[i] = fExpr.getFoldedVector(1) - continue - } + if !folded { + allParametersFolded = false } - if atRuntime { - if pExpr, ok := param.(*ParamExpressionExecutor); ok { - expr.parameterResults[i], err = pExpr.Eval(proc, nil, nil) - if err != nil { - return - } - continue - } - } - - allParametersFolded = false } if !allParametersFolded || expr.volatile || (!atRuntime && expr.timeDependent) { return nil @@ -164,19 +298,7 @@ func (expr *FunctionExpressionExecutor) doFold(proc *process.Process, atRuntime } } - // fold the function. - if err = expr.resultVector.PreExtendAndReset(execLen); err != nil { - return err - } - if err = expr.evalFn(expr.parameterResults, expr.resultVector, proc, execLen, nil); err != nil { - return err - } - if execLen == 1 { - expr.resultVector.GetResultVector().ToConst() - } - - expr.folded.canFold = true - return nil + return expr.finishFolding(proc, execLen) } func (expr *ParamExpressionExecutor) ResetForNextQuery() { diff --git a/pkg/sql/colexec/evalExpression_test.go b/pkg/sql/colexec/evalExpression_test.go index 913037c24186c..e161e53440521 100644 --- a/pkg/sql/colexec/evalExpression_test.go +++ b/pkg/sql/colexec/evalExpression_test.go @@ -454,6 +454,860 @@ func TestFunctionExpressionExecutor(t *testing.T) { } } +func TestFlowControlShortCircuitInvalidCast(t *testing.T) { + proc := testutil.NewProcess(t) + defer proc.Free() + + stringConst := func(value string) *plan.Expr { + return &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar), NotNullable: true}, + Expr: &plan.Expr_Lit{Lit: &plan.Literal{ + Value: &plan.Literal_Sval{Sval: value}, + }}, + } + } + uint8Const := func(value uint8) *plan.Expr { + return &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_uint8), NotNullable: true}, + Expr: &plan.Expr_Lit{Lit: &plan.Literal{ + Value: &plan.Literal_U8Val{U8Val: uint32(value)}, + }}, + } + } + bindFunction := func(name string, args ...*plan.Expr) *plan.Expr { + argTypes := make([]types.Type, len(args)) + for i := range args { + argTypes[i] = types.New(types.T(args[i].Typ.Id), args[i].Typ.Width, args[i].Typ.Scale) + } + fn, err := function.GetFunctionByName(proc.Ctx, name, argTypes) + require.NoError(t, err) + retType := fn.GetReturnType() + return &plan.Expr{ + Typ: plan.Type{Id: int32(retType.Oid), Width: retType.Width, Scale: retType.Scale}, + Expr: &plan.Expr_F{F: &plan.Function{ + Func: &plan.ObjectRef{Obj: fn.GetEncodedOverloadID(), ObjName: name}, + Args: args, + }}, + } + } + invalidCast := func() *plan.Expr { + target := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_int64), NotNullable: true}, + Expr: &plan.Expr_T{T: &plan.TargetType{}}, + } + return bindFunction("cast", stringConst("bad"), target) + } + + tests := []struct { + name string + expr *plan.Expr + want int64 + }{ + { + name: "if skips true branch", + expr: bindFunction("if", makePlan2BoolConstExprWithType(false), invalidCast(), makePlan2Int64ConstExprWithType(7)), + want: 7, + }, + { + name: "case skips then branch", + expr: bindFunction("case", makePlan2BoolConstExprWithType(false), invalidCast(), makePlan2Int64ConstExprWithType(7)), + want: 7, + }, + { + name: "coalesce skips later argument", + expr: bindFunction("coalesce", makePlan2Int64ConstExprWithType(5), invalidCast()), + want: 5, + }, + { + name: "ifnull rewrite skips second argument", + expr: bindFunction("case", + bindFunction("isnull", makePlan2Int64ConstExprWithType(5)), + invalidCast(), + makePlan2Int64ConstExprWithType(5)), + want: 5, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + executor, err := NewExpressionExecutor(proc, test.expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, nil, nil) + require.NoError(t, err) + require.Equal(t, test.want, vector.MustFixedColWithTypeCheck[int64](result)[0]) + require.True(t, executor.(*FunctionExpressionExecutor).folded.canFold) + }) + } + + t.Run("if evaluates selected branch", func(t *testing.T) { + expr := bindFunction("if", makePlan2BoolConstExprWithType(true), invalidCast(), makePlan2Int64ConstExprWithType(7)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + _, err = executor.Eval(proc, nil, nil) + require.ErrorContains(t, err, "invalid argument cast to int") + }) + + t.Run("coalesce evaluates remaining argument", func(t *testing.T) { + nullInt64 := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_int64)}, + Expr: &plan.Expr_Lit{Lit: &plan.Literal{Isnull: true}}, + } + expr := bindFunction("coalesce", nullInt64, invalidCast()) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + _, err = executor.Eval(proc, nil, nil) + require.ErrorContains(t, err, "invalid argument cast to int") + }) + + t.Run("skipped varlen function preserves batch length", func(t *testing.T) { + input := batch.New(nil) + input.SetRowCount(2) + expr := bindFunction("concat", stringConst("a"), stringConst("b")) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, []bool{false, false}) + require.NoError(t, err) + require.Equal(t, 2, result.Length()) + require.True(t, result.GetNulls().Contains(0)) + require.True(t, result.GetNulls().Contains(1)) + }) + + column := func(pos int32, typ types.Type) *plan.Expr { + return &plan.Expr{ + Typ: plan.Type{Id: int32(typ.Oid), Width: typ.Width, Scale: typ.Scale}, + Expr: &plan.Expr_Col{Col: &plan.ColRef{RelPos: 0, ColPos: pos}}, + } + } + castTo := func(source *plan.Expr, typ types.Type) *plan.Expr { + target := &plan.Expr{ + Typ: plan.Type{Id: int32(typ.Oid), Width: typ.Width, Scale: typ.Scale, NotNullable: true}, + Expr: &plan.Expr_T{T: &plan.TargetType{}}, + } + return bindFunction("cast", source, target) + } + castToInt64 := func(source *plan.Expr) *plan.Expr { + return castTo(source, types.T_int64.ToType()) + } + typedNull := func(typ types.Type) *plan.Expr { + return &plan.Expr{ + Typ: plan.Type{Id: int32(typ.Oid), Width: typ.Width, Scale: typ.Scale}, + Expr: &plan.Expr_Lit{Lit: &plan.Literal{Isnull: true}}, + } + } + + t.Run("if skips unresolved variable leaf across reuse", func(t *testing.T) { + leafProc := testutil.NewProcess(t) + defer leafProc.Free() + resolveCalls := 0 + leafProc.SetResolveVariableFunc(func(string, bool, bool) (interface{}, error) { + resolveCalls++ + return nil, moerr.NewInternalErrorNoCtx("missing variable") + }) + + variable := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_V{V: &plan.VarRef{ + Name: "missing_user_variable", + }}, + } + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + variable, + stringConst("ok")) + executor, err := NewExpressionExecutor(leafProc, expr) + require.NoError(t, err) + defer executor.Free() + + eval := func(condition bool) (*vector.Vector, error) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(1, types.T_bool.ToType(), leafProc.Mp(), false, []bool{condition}), + }, nil) + defer input.Clean(leafProc.Mp()) + return executor.Eval(leafProc, []*batch.Batch{input}, nil) + } + + result, err := eval(false) + require.NoError(t, err) + require.Equal(t, "ok", result.GetStringAt(0)) + require.Zero(t, resolveCalls) + + _, err = eval(true) + require.ErrorContains(t, err, "missing variable") + require.Equal(t, 1, resolveCalls) + + result, err = eval(false) + require.NoError(t, err) + require.Equal(t, "ok", result.GetStringAt(0)) + require.Equal(t, 1, resolveCalls) + }) + + t.Run("case and coalesce skip unresolved variable leaves", func(t *testing.T) { + leafProc := testutil.NewProcess(t) + defer leafProc.Free() + resolveCalls := 0 + leafProc.SetResolveVariableFunc(func(string, bool, bool) (interface{}, error) { + resolveCalls++ + return nil, moerr.NewInternalErrorNoCtx("missing variable") + }) + + variable := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_V{V: &plan.VarRef{ + Name: "missing_user_variable", + }}, + } + tests := []struct { + name string + expr *plan.Expr + input *vector.Vector + }{ + { + name: "case", + expr: bindFunction("case", + column(0, types.T_bool.ToType()), + variable, + stringConst("ok")), + input: testutil.NewVector(1, types.T_bool.ToType(), leafProc.Mp(), false, []bool{false}), + }, + { + name: "coalesce", + expr: bindFunction("coalesce", + column(0, types.T_varchar.ToType()), + variable), + input: testutil.NewVector(1, types.T_varchar.ToType(), leafProc.Mp(), false, []string{"ok"}), + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{test.input}, nil) + defer input.Clean(leafProc.Mp()) + executor, err := NewExpressionExecutor(leafProc, test.expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(leafProc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, "ok", result.GetStringAt(0)) + }) + } + require.Zero(t, resolveCalls) + }) + + t.Run("if skips missing parameter leaf", func(t *testing.T) { + leafProc := testutil.NewProcess(t) + defer leafProc.Free() + params := vector.NewVec(types.T_text.ToType()) + defer params.Free(leafProc.Mp()) + leafProc.SetPrepareParams(params) + + parameter := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_P{P: &plan.ParamRef{Pos: 0}}, + } + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + parameter, + stringConst("ok")) + executor, err := NewExpressionExecutor(leafProc, expr) + require.NoError(t, err) + defer executor.Free() + + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(1, types.T_bool.ToType(), leafProc.Mp(), false, []bool{false}), + }, nil) + defer input.Clean(leafProc.Mp()) + result, err := executor.Eval(leafProc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, "ok", result.GetStringAt(0)) + }) + + t.Run("parameter leaf remains valid after a skipped generation", func(t *testing.T) { + leafProc := testutil.NewProcess(t) + defer leafProc.Free() + params := vector.NewVec(types.T_text.ToType()) + require.NoError(t, vector.AppendBytes(params, []byte("parameter"), false, leafProc.Mp())) + defer params.Free(leafProc.Mp()) + leafProc.SetPrepareParams(params) + + parameter := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_P{P: &plan.ParamRef{Pos: 0}}, + } + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + parameter, + stringConst("fallback")) + executor, err := NewExpressionExecutor(leafProc, expr) + require.NoError(t, err) + defer executor.Free() + + eval := func(condition bool) string { + t.Helper() + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(1, types.T_bool.ToType(), leafProc.Mp(), false, []bool{condition}), + }, nil) + defer input.Clean(leafProc.Mp()) + result, err := executor.Eval(leafProc, []*batch.Batch{input}, nil) + require.NoError(t, err) + return result.GetStringAt(0) + } + + require.Equal(t, "fallback", eval(false)) + require.Equal(t, "parameter", eval(true)) + require.Equal(t, "parameter", eval(true)) + }) + + t.Run("runtime parameter folding follows prepared statement reset", func(t *testing.T) { + leafProc := testutil.NewProcess(t) + defer leafProc.Free() + leafProc.SetBaseProcessRunningStatus(true) + + parameter := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_P{P: &plan.ParamRef{Pos: 0}}, + } + expr := bindFunction("if", + makePlan2BoolConstExprWithType(true), + parameter, + stringConst("fallback")) + executor, err := NewExpressionExecutor(leafProc, expr) + require.NoError(t, err) + defer executor.Free() + + eval := func(value string) string { + t.Helper() + params := vector.NewVec(types.T_text.ToType()) + require.NoError(t, vector.AppendBytes(params, []byte(value), false, leafProc.Mp())) + leafProc.SetPrepareParams(params) + defer func() { + leafProc.SetPrepareParams(nil) + params.Free(leafProc.Mp()) + }() + + result, evalErr := executor.Eval(leafProc, nil, nil) + require.NoError(t, evalErr) + require.True(t, executor.(*FunctionExpressionExecutor).folded.canFold) + return result.GetStringAt(0) + } + + require.Equal(t, "first", eval("first")) + executor.ResetForNextQuery() + require.Equal(t, "second", eval("second")) + }) + + t.Run("case without else and coalesce all null still fold", func(t *testing.T) { + nullInt64 := typedNull(types.T_int64.ToType()) + expressions := []*plan.Expr{ + bindFunction("case", makePlan2BoolConstExprWithType(false), makePlan2Int64ConstExprWithType(7)), + bindFunction("coalesce", nullInt64, typedNull(types.T_int64.ToType())), + } + for _, expr := range expressions { + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + + result, err := executor.Eval(proc, nil, nil) + require.NoError(t, err) + require.True(t, result.IsConstNull()) + require.True(t, executor.(*FunctionExpressionExecutor).folded.canFold) + executor.Free() + } + }) + + t.Run("constant flow control stays allocation-free after folding", func(t *testing.T) { + expr := bindFunction("if", + makePlan2BoolConstExprWithType(true), + makePlan2Int64ConstExprWithType(7), + makePlan2Int64ConstExprWithType(9)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + input := batch.New(nil) + input.SetRowCount(8192) + batches := []*batch.Batch{input} + result, err := executor.Eval(proc, batches, nil) + require.NoError(t, err) + require.True(t, result.IsConst()) + require.Equal(t, 8192, result.Length()) + require.Equal(t, int64(7), vector.MustFixedColWithTypeCheck[int64](result)[0]) + require.True(t, executor.(*FunctionExpressionExecutor).folded.canFold) + + var evalErr error + allocations := testing.AllocsPerRun(100, func() { + _, evalErr = executor.Eval(proc, batches, nil) + }) + require.NoError(t, evalErr) + require.LessOrEqual(t, allocations, 1.0) + }) + + t.Run("partial evaluation preserves runtime result type", func(t *testing.T) { + sourceType := types.New(types.T_float32, 10, 2) + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, sourceType, proc.Mp(), false, []float32{1.25, 2.5}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := castTo(column(0, sourceType), types.T_float64.ToType()) + evalType := func(selectList []bool) types.Type { + t.Helper() + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, selectList) + require.NoError(t, err) + return *result.GetType() + } + + fullType := evalType(nil) + partialType := evalType([]bool{false, true}) + require.Equal(t, fullType, partialType) + require.Equal(t, int32(10), partialType.Width) + require.Equal(t, int32(2), partialType.Scale) + }) + + t.Run("partial runtime result type updates across reuse", func(t *testing.T) { + expr := bindFunction("sysdate", column(0, types.T_int64.ToType())) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + evalScale := func(scale int64) { + t.Helper() + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_int64.ToType(), proc.Mp(), false, []int64{6, scale}), + }, nil) + defer input.Clean(proc.Mp()) + + result, err := executor.Eval(proc, []*batch.Batch{input}, []bool{false, true}) + require.NoError(t, err) + require.Equal(t, int32(scale), result.GetType().Scale) + } + + evalScale(3) + evalScale(1) + evalScale(5) + }) + + t.Run("nested consumer observes partial runtime result type", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true}), + }, nil) + defer input.Clean(proc.Mp()) + + sysdate := bindFunction("sysdate", makePlan2Int64ConstExprWithType(3)) + asChar := castTo(sysdate, types.New(types.T_char, 64, 0)) + directExecutor, err := NewExpressionExecutor(proc, asChar) + require.NoError(t, err) + defer directExecutor.Free() + direct, err := directExecutor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + + ifExpr := bindFunction("if", + column(0, types.T_bool.ToType()), + asChar, + stringConst("fallback")) + ifExecutor, err := NewExpressionExecutor(proc, ifExpr) + require.NoError(t, err) + defer ifExecutor.Free() + partial, err := ifExecutor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + + require.Equal(t, 23, len(direct.GetStringAt(0))) + require.Equal(t, "fallback", partial.GetStringAt(0)) + require.Equal(t, len(direct.GetStringAt(0)), len(partial.GetStringAt(1))) + }) + + t.Run("if skips invalid rows within a batch", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"bad", "9"}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + castToInt64(column(1, types.T_varchar.ToType())), + makePlan2Int64ConstExprWithType(7)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, []int64{7, 9}, vector.MustFixedColWithTypeCheck[int64](result)) + }) + + t.Run("if skips invalid regexp rows within a batch", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"x", "a"}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"[", "a"}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"c", "c"}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + bindFunction("regexp_like", + column(1, types.T_varchar.ToType()), + column(2, types.T_varchar.ToType()), + column(3, types.T_varchar.ToType())), + makePlan2BoolConstExprWithType(false)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, []bool{false, true}, vector.MustFixedColWithTypeCheck[bool](result)) + }) + + t.Run("if still evaluates invalid selected regexp row", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{true, false}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"x", "a"}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"[", "a"}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"c", "c"}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + bindFunction("regexp_like", + column(1, types.T_varchar.ToType()), + column(2, types.T_varchar.ToType()), + column(3, types.T_varchar.ToType())), + makePlan2BoolConstExprWithType(false)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + _, err = executor.Eval(proc, []*batch.Batch{input}, nil) + require.Error(t, err) + }) + + t.Run("if does not execute sleep on unselected rows", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true}), + testutil.NewVector(2, types.T_float64.ToType(), proc.Mp(), false, []float64{-1, 0}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + bindFunction("sleep", column(1, types.T_float64.ToType())), + uint8Const(0)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, []uint8{0, 0}, vector.MustFixedColWithTypeCheck[uint8](result)) + }) + + t.Run("if preserves non-row-aligned in-list parameters", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(3, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true, false}), + testutil.NewVector(3, types.T_varchar.ToType(), proc.Mp(), false, []string{"x", "a", "b"}), + }, nil) + defer input.Clean(proc.Mp()) + + list := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_varchar)}, + Expr: &plan.Expr_List{List: &plan.ExprList{List: []*plan.Expr{ + stringConst("a"), + stringConst("b"), + }}}, + } + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + bindFunction("in", column(1, types.T_varchar.ToType()), list), + makePlan2BoolConstExprWithType(false)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, []bool{false, true, false}, vector.MustFixedColWithTypeCheck[bool](result)) + }) + + t.Run("case preserves first match across multiple when clauses and reuse", func(t *testing.T) { + expr := bindFunction("case", + column(0, types.T_bool.ToType()), + makePlan2Int64ConstExprWithType(1), + column(1, types.T_bool.ToType()), + castToInt64(column(2, types.T_varchar.ToType())), + makePlan2Int64ConstExprWithType(7)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + tests := []struct { + name string + firstCondition []bool + firstNulls []bool + secondCondition []bool + secondNulls []bool + values []string + parentSelect []bool + want []int64 + }{ + { + name: "multiple rows choose first later else and null conditions", + firstCondition: []bool{true, false, false, false, true}, + firstNulls: []bool{false, false, true, true, false}, + secondCondition: []bool{true, true, false, true, true}, + secondNulls: []bool{false, false, true, false, false}, + values: []string{"bad", "9", "bad", "11", "bad"}, + want: []int64{1, 9, 7, 11, 1}, + }, + { + name: "changing and shrinking batch selects later branch", + firstCondition: []bool{false, true}, + secondCondition: []bool{true, true}, + values: []string{"13", "bad"}, + want: []int64{13, 1}, + }, + { + name: "subsequent reuse keeps first match state", + firstCondition: []bool{true, false, false}, + secondCondition: []bool{true, false, true}, + values: []string{"bad", "bad", "17"}, + want: []int64{1, 7, 17}, + }, + { + name: "parent partial selection cannot be reselected", + firstCondition: []bool{true, false, false, true}, + secondCondition: []bool{true, true, true, true}, + values: []string{"bad", "19", "bad", "bad"}, + parentSelect: []bool{true, true, false, false}, + want: []int64{1, 19, 0, 0}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + require.Len(t, test.secondCondition, len(test.firstCondition)) + require.Len(t, test.values, len(test.firstCondition)) + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVectorWithNulls(len(test.firstCondition), types.T_bool.ToType(), proc.Mp(), false, test.firstNulls, test.firstCondition), + testutil.NewVectorWithNulls(len(test.secondCondition), types.T_bool.ToType(), proc.Mp(), false, test.secondNulls, test.secondCondition), + testutil.NewVector(len(test.values), types.T_varchar.ToType(), proc.Mp(), false, test.values), + }, nil) + defer input.Clean(proc.Mp()) + + result, err := executor.Eval(proc, []*batch.Batch{input}, test.parentSelect) + require.NoError(t, err) + values := vector.MustFixedColWithTypeCheck[int64](result) + require.Len(t, values, len(test.want)) + for row := range test.want { + if test.parentSelect != nil && !test.parentSelect[row] { + continue + } + require.False(t, result.IsNull(uint64(row)), "row %d", row) + require.Equal(t, test.want[row], values[row], "row %d", row) + } + }) + } + }) + + for _, test := range []struct { + name string + targetType types.Type + validValue string + fallback *plan.Expr + }{ + { + name: "bool", + targetType: types.T_bool.ToType(), + validValue: "true", + fallback: makePlan2BoolConstExprWithType(false), + }, + { + name: "uuid", + targetType: types.T_uuid.ToType(), + validValue: "00000000-0000-0000-0000-000000000001", + fallback: typedNull(types.T_uuid.ToType()), + }, + { + name: "json", + targetType: types.T_json.ToType(), + validValue: `{"ok":true}`, + fallback: typedNull(types.T_json.ToType()), + }, + } { + t.Run("if skips unselected "+test.name+" cast rows", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{false, true}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"bad", test.validValue}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + castTo(column(1, types.T_varchar.ToType()), test.targetType), + test.fallback) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.False(t, result.IsNull(1)) + if test.targetType.Oid == types.T_bool { + require.Equal(t, []bool{false, true}, vector.MustFixedColWithTypeCheck[bool](result)) + } else { + require.True(t, result.IsNull(0)) + } + }) + } + + t.Run("if still evaluates invalid selected bool cast row", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(2, types.T_bool.ToType(), proc.Mp(), false, []bool{true, false}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"bad", "true"}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("if", + column(0, types.T_bool.ToType()), + castTo(column(1, types.T_varchar.ToType()), types.T_bool.ToType()), + makePlan2BoolConstExprWithType(false)) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + _, err = executor.Eval(proc, []*batch.Batch{input}, nil) + require.ErrorContains(t, err, "not a valid bool expression") + }) + + t.Run("coalesce skips invalid rows within a batch", func(t *testing.T) { + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVectorWithNulls(2, types.T_int64.ToType(), proc.Mp(), false, []bool{false, true}, []int64{5, 0}), + testutil.NewVector(2, types.T_varchar.ToType(), proc.Mp(), false, []string{"bad", "9"}), + }, nil) + defer input.Clean(proc.Mp()) + + expr := bindFunction("coalesce", + column(0, types.T_int64.ToType()), + castToInt64(column(1, types.T_varchar.ToType()))) + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(t, err) + defer executor.Free() + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, []int64{5, 9}, vector.MustFixedColWithTypeCheck[int64](result)) + }) + + for _, test := range []struct { + name string + expr *plan.Expr + }{ + { + name: "if", + expr: bindFunction("if", + column(0, types.T_bool.ToType()), + castToInt64(column(1, types.T_varchar.ToType())), + makePlan2Int64ConstExprWithType(7)), + }, + { + name: "case", + expr: bindFunction("case", + column(0, types.T_bool.ToType()), + castToInt64(column(1, types.T_varchar.ToType())), + makePlan2Int64ConstExprWithType(7)), + }, + } { + t.Run(test.name+" reuses executor across shrinking batches", func(t *testing.T) { + executor, err := NewExpressionExecutor(proc, test.expr) + require.NoError(t, err) + defer executor.Free() + + eval := func(conditions []bool, values []string, expected []int64) { + t.Helper() + require.Len(t, values, len(conditions)) + input := testutil.NewBatchWithVectors([]*vector.Vector{ + testutil.NewVector(len(conditions), types.T_bool.ToType(), proc.Mp(), false, conditions), + testutil.NewVector(len(values), types.T_varchar.ToType(), proc.Mp(), false, values), + }, nil) + defer input.Clean(proc.Mp()) + + result, err := executor.Eval(proc, []*batch.Batch{input}, nil) + require.NoError(t, err) + require.Equal(t, expected, vector.MustFixedColWithTypeCheck[int64](result)) + } + + eval( + []bool{false, false, false, false, false}, + []string{"bad", "bad", "bad", "bad", "bad"}, + []int64{7, 7, 7, 7, 7}, + ) + eval( + []bool{true, true}, + []string{"8", "9"}, + []int64{8, 9}, + ) + eval( + []bool{true, false, true}, + []string{"10", "bad", "12"}, + []int64{10, 7, 12}, + ) + }) + } +} + +func BenchmarkConstantFlowControlExpression(b *testing.B) { + proc := testutil.NewProcess(b) + defer proc.Free() + + fn, err := function.GetFunctionByName(proc.Ctx, "if", []types.Type{ + types.T_bool.ToType(), + types.T_int64.ToType(), + types.T_int64.ToType(), + }) + require.NoError(b, err) + expr := &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_int64)}, + Expr: &plan.Expr_F{F: &plan.Function{ + Func: &plan.ObjectRef{Obj: fn.GetEncodedOverloadID(), ObjName: "if"}, + Args: []*plan.Expr{ + makePlan2BoolConstExprWithType(true), + makePlan2Int64ConstExprWithType(7), + makePlan2Int64ConstExprWithType(9), + }, + }}, + } + executor, err := NewExpressionExecutor(proc, expr) + require.NoError(b, err) + defer executor.Free() + input := batch.New(nil) + input.SetRowCount(8192) + batches := []*batch.Batch{input} + _, err = executor.Eval(proc, batches, nil) + require.NoError(b, err) + + b.ReportAllocs() + b.ResetTimer() + for range b.N { + if _, err = executor.Eval(proc, batches, nil); err != nil { + b.Fatal(err) + } + } +} + func TestExpressionReset(t *testing.T) { proc := testutil.NewProcess(t) diff --git a/test/distributed/cases/function/func_if.result b/test/distributed/cases/function/func_if.result index 352812dbbd564..8471c9d28684c 100644 --- a/test/distributed/cases/function/func_if.result +++ b/test/distributed/cases/function/func_if.result @@ -244,4 +244,11 @@ string 𝄀 123 𝄀 mixed DROP TABLE t1; +SET SESSION sql_mode = ''; +SELECT IF(0, CAST('bad' AS SIGNED), 7) AS if_skip_true_branch, +CASE WHEN 0 THEN CAST('bad' AS SIGNED) ELSE 7 END AS case_skip_then, +COALESCE(5, CAST('bad' AS SIGNED)) AS coalesce_skip_later, +IFNULL(5, CAST('bad' AS SIGNED)) AS ifnull_skip_later; +➤ if_skip_true_branch[-5,64,0] ¦ case_skip_then[-5,64,0] ¦ coalesce_skip_later[-5,64,0] ¦ ifnull_skip_later[-5,64,0] 𝄀 +7 ¦ 7 ¦ 5 ¦ 5 SET TIME_ZONE = "SYSTEM"; diff --git a/test/distributed/cases/function/func_if.test b/test/distributed/cases/function/func_if.test index c2e9e36d5e3ac..bb9c11968c252 100644 --- a/test/distributed/cases/function/func_if.test +++ b/test/distributed/cases/function/func_if.test @@ -204,5 +204,13 @@ INSERT INTO t1 VALUES(1),(2),(3); SELECT CASE WHEN id = 1 THEN 'string' WHEN id = 2 THEN 123 ELSE 'mixed' END AS mixed_case FROM t1 ORDER BY id; DROP TABLE t1; +-- @bvt:issue#25314 +SET SESSION sql_mode = ''; +SELECT IF(0, CAST('bad' AS SIGNED), 7) AS if_skip_true_branch, + CASE WHEN 0 THEN CAST('bad' AS SIGNED) ELSE 7 END AS case_skip_then, + COALESCE(5, CAST('bad' AS SIGNED)) AS coalesce_skip_later, + IFNULL(5, CAST('bad' AS SIGNED)) AS ifnull_skip_later; +-- @bvt:issue + # reset SET TIME_ZONE = "SYSTEM";