From 1cdcaf472fdaadc426e28992c1c50c0f1eae27ea Mon Sep 17 00:00:00 2001 From: rachithrr Date: Wed, 11 Jan 2023 12:32:27 -0500 Subject: [PATCH] FB-1814: Implement ASCII() (#2378) * FB-1814: Implement ASCII() (cherry picked from commit f59022747105b937a54778fb8e74c174e524ee82) --- sql3/errors.go | 10 +- sql3/planner/expression.go | 2 + sql3/planner/expressionanalyzercall.go | 2 + sql3/planner/expressiontypes.go | 9 + sql3/planner/inbuiltfunctionsstring.go | 513 +++++++++++++++--------- sql3/test/defs/defs_string_functions.go | 287 +++++++++++-- 6 files changed, 613 insertions(+), 210 deletions(-) diff --git a/sql3/errors.go b/sql3/errors.go index 91535a277..10c5ceff3 100644 --- a/sql3/errors.go +++ b/sql3/errors.go @@ -115,7 +115,8 @@ const ( ErrAggregateNotAllowedInGroupBy errors.Code = "ErrIdPercentileNotAllowedInGroupBy" // function evaluation - ErrValueOutOfRange errors.Code = "ErrValueOutOfRange" + ErrValueOutOfRange errors.Code = "ErrValueOutOfRange" + ErrStringLengthMismatch errors.Code = "ErrStringLengthMismatch" ) func NewErrDuplicateColumn(line int, col int, column string) error { @@ -705,3 +706,10 @@ func NewErrValueOutOfRange(line, col int, val interface{}) error { fmt.Sprintf("[%d:%d] value '%v' out of range", line, col, val), ) } + +func NewErrStringLengthMismatch(line, col, len int, val interface{}) error { + return errors.New( + ErrStringLengthMismatch, + fmt.Sprintf("[%d:%d] value '%v' should be of the length %d", line, col, val, len), + ) +} diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index 70dd3e736..e0ad68c1f 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -1496,6 +1496,8 @@ func (n *callPlanExpression) Evaluate(currentRow []interface{}) (interface{}, er return n.EvaluateStringSplit(currentRow) case "CHAR": return n.EvaluateChar(currentRow) + case "ASCII": + return n.EvaluateAscii(currentRow) case "SUBSTRING": return n.EvaluateSubstring(currentRow) case "LOWER": diff --git a/sql3/planner/expressionanalyzercall.go b/sql3/planner/expressionanalyzercall.go index 65e2b81e1..a2b6d04b1 100644 --- a/sql3/planner/expressionanalyzercall.go +++ b/sql3/planner/expressionanalyzercall.go @@ -245,6 +245,8 @@ func (p *ExecutionPlanner) analyzeCallExpression(call *parser.Call, scope parser return p.analyseFunctionReverse(call, scope) case "CHAR": return p.analyseFunctionChar(call, scope) + case "ASCII": + return p.analyseFunctionAscii(call, scope) case "UPPER": return p.analyzeFunctionUpper(call, scope) case "STRINGSPLIT": diff --git a/sql3/planner/expressiontypes.go b/sql3/planner/expressiontypes.go index 03698174b..8be5d2fb1 100644 --- a/sql3/planner/expressiontypes.go +++ b/sql3/planner/expressiontypes.go @@ -475,6 +475,15 @@ func typeIsString(testType parser.ExprDataType) bool { } } +func typeIsVoid(testType parser.ExprDataType) bool { + switch testType.(type) { + case *parser.DataTypeVoid: + return true + default: + return false + } +} + // returns true if the type is timestamp func typeIsTimestamp(testType parser.ExprDataType) bool { switch testType.(type) { diff --git a/sql3/planner/inbuiltfunctionsstring.go b/sql3/planner/inbuiltfunctionsstring.go index e5bee003f..9cb83c88b 100644 --- a/sql3/planner/inbuiltfunctionsstring.go +++ b/sql3/planner/inbuiltfunctionsstring.go @@ -14,7 +14,35 @@ func (p *ExecutionPlanner) analyseFunctionReverse(call *parser.Call, scope parse return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) } - if !typeIsString(call.Args[0].DataType()) { + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) + } + + call.ResultDataType = parser.NewDataTypeString() + + return call, nil +} + +func (p *ExecutionPlanner) analyzeFunctionLower(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + if len(call.Args) != 1 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) + } + + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) + } + call.ResultDataType = parser.NewDataTypeString() + + return call, nil +} + +func (p *ExecutionPlanner) analyzeFunctionUpper(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + //one argument for Upper Function + if len(call.Args) != 1 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) + } + + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } @@ -29,7 +57,7 @@ func (p *ExecutionPlanner) analyseFunctionChar(call *parser.Call, scope parser.S return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) } - if !typeIsInteger(call.Args[0].DataType()) { + if !typeIsInteger(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrIntExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } @@ -38,18 +66,33 @@ func (p *ExecutionPlanner) analyseFunctionChar(call *parser.Call, scope parser.S return call, nil } +func (p *ExecutionPlanner) analyseFunctionAscii(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + //one argument + if len(call.Args) != 1 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) + } + + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) + } + + call.ResultDataType = parser.NewDataTypeInt() + + return call, nil +} + func (p *ExecutionPlanner) analyseFunctionSubstring(call *parser.Call, scope parser.Statement) (parser.Expr, error) { if len(call.Args) <= 1 || len(call.Args) > 3 { return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) } - if !typeIsString(call.Args[0].DataType()) { + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } // the third parameter is optional for i := 1; i < len(call.Args); i++ { - if !typeIsInteger(call.Args[i].DataType()) { + if !typeIsInteger(call.Args[i].DataType()) && !typeIsVoid(call.Args[i].DataType()) { return nil, sql3.NewErrIntExpressionExpected(call.Args[i].Pos().Line, call.Args[i].Pos().Column) } } @@ -59,57 +102,20 @@ func (p *ExecutionPlanner) analyseFunctionSubstring(call *parser.Call, scope par return call, nil } -func (p *ExecutionPlanner) analyzeFunctionLower(call *parser.Call, scope parser.Statement) (parser.Expr, error) { - if len(call.Args) != 1 { - return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) - } - - if !typeIsString(call.Args[0].DataType()) { - return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) - } - call.ResultDataType = parser.NewDataTypeString() - - return call, nil -} - -func (p *ExecutionPlanner) analyseFunctionStringSplit(call *parser.Call, scope parser.Statement) (parser.Expr, error) { - if len(call.Args) <= 1 || len(call.Args) > 3 { - return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) - } - - if !typeIsString(call.Args[0].DataType()) { - return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) - } - - // string seperator - if !typeIsString(call.Args[1].DataType()) { - return nil, sql3.NewErrStringExpressionExpected(call.Args[1].Pos().Line, call.Args[1].Pos().Column) - } - - // third argument is the position. optional, defaults to 0 - if len(call.Args) == 3 && !typeIsInteger(call.Args[2].DataType()) { - return nil, sql3.NewErrIntExpressionExpected(call.Args[2].Pos().Line, call.Args[2].Pos().Column) - } - - call.ResultDataType = parser.NewDataTypeString() - - return call, nil -} - func (p *ExecutionPlanner) analyseFunctionReplaceAll(call *parser.Call, scope parser.Statement) (parser.Expr, error) { if len(call.Args) != 3 { return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 3, len(call.Args)) } // input string - if !typeIsString(call.Args[0].DataType()) { + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } // string to find and replace - if !typeIsString(call.Args[1].DataType()) { + if !typeIsString(call.Args[1].DataType()) && !typeIsVoid(call.Args[1].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[1].Pos().Line, call.Args[1].Pos().Column) } // string to replace with - if !typeIsString(call.Args[2].DataType()) { + if !typeIsString(call.Args[2].DataType()) && !typeIsVoid(call.Args[2].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[2].Pos().Line, call.Args[2].Pos().Column) } @@ -118,28 +124,38 @@ func (p *ExecutionPlanner) analyseFunctionReplaceAll(call *parser.Call, scope pa return call, nil } -// reverses the string -func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) - if err != nil { - return nil, err +func (p *ExecutionPlanner) analyseFunctionStringSplit(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + if len(call.Args) <= 1 || len(call.Args) > 3 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) } - // reverse the string - runes := []rune(stringArgOne) - for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { - runes[i], runes[j] = runes[j], runes[i] + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } - return string(runes), nil + + // string seperator + if !typeIsString(call.Args[1].DataType()) && !typeIsVoid(call.Args[1].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[1].Pos().Line, call.Args[1].Pos().Column) + } + + // third argument is the position. optional, defaults to 0 + if len(call.Args) == 3 && !typeIsInteger(call.Args[2].DataType()) && !typeIsVoid(call.Args[2].DataType()) { + return nil, sql3.NewErrIntExpressionExpected(call.Args[2].Pos().Line, call.Args[2].Pos().Column) + } + + call.ResultDataType = parser.NewDataTypeString() + + return call, nil } -func (p *ExecutionPlanner) analyzeFunctionUpper(call *parser.Call, scope parser.Statement) (parser.Expr, error) { - //one argument for Upper Function +// Analyze function for Trim/RTrim/LTrim +func (p *ExecutionPlanner) analyseFunctionTrim(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + //one argument for Trim Functions if len(call.Args) != 1 { return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) } - if !typeIsString(call.Args[0].DataType()) { + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } @@ -148,249 +164,364 @@ func (p *ExecutionPlanner) analyzeFunctionUpper(call *parser.Call, scope parser. return call, nil } +func (p *ExecutionPlanner) analyseFunctionPrefixSuffix(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + if len(call.Args) != 2 { + return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) + } + + if !typeIsString(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { + return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) + } + + if !typeIsInteger(call.Args[1].DataType()) && !typeIsVoid(call.Args[1].DataType()) { + return nil, sql3.NewErrIntExpressionExpected(call.Args[1].Pos().Line, call.Args[1].Pos().Column) + } + + call.ResultDataType = parser.NewDataTypeString() + + return call, nil +} + func (p *ExecutionPlanner) analyseFunctionSpace(call *parser.Call, scope parser.Statement) (parser.Expr, error) { //one argument if len(call.Args) != 1 { return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) } - if !typeIsInteger(call.Args[0].DataType()) { + if !typeIsInteger(call.Args[0].DataType()) && !typeIsVoid(call.Args[0].DataType()) { return nil, sql3.NewErrIntExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) } call.ResultDataType = parser.NewDataTypeString() - return call, nil } +// reverses the string +func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interface{}, error) { + argEval, err := n.args[0].Evaluate(currentRow) + if err != nil { + return nil, err + } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + + // reverse the string + runes := []rune(stringArg) + for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { + runes[i], runes[j] = runes[j], runes[i] + } + return string(runes), nil +} + +func (n *callPlanExpression) EvaluateLower(currentRow []interface{}) (interface{}, error) { + argEval, err := n.args[0].Evaluate(currentRow) + if err != nil { + return nil, err + } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + + return strings.ToLower(stringArg), nil +} + // Convert string to Upper case func (n *callPlanExpression) EvaluateUpper(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } // convert to Upper - return strings.ToUpper(stringArgOne), nil + return strings.ToUpper(stringArg), nil } func (n *callPlanExpression) EvaluateChar(currentRow []interface{}) (interface{}, error) { // Get the integer argument from the function call - intArg, err := evaluateIntArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { - return "", err + return 0, err + } + if argEval == nil { + return nil, nil + } + intArg, ok := argEval.(int64) + if !ok { + return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) } // Return the character that corresponds to the integer value return string(rune(intArg)), nil } -// Takes string, startIndex and length and returns the substring. -func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) +// this takes a string and returns the ascii value. +// sthe string should be of the length 1. +func (n *callPlanExpression) EvaluateAscii(currentRow []interface{}) (interface{}, error) { + // Get the string argument from the function call + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + + if len(stringArg) == 0 { + return "", nil + } + + if len(stringArg) != 1 { + return nil, sql3.NewErrStringLengthMismatch(0, 0, 1, stringArg) + } + + res := []rune(stringArg) + return int64(res[0]), nil +} + +// Takes string, startIndex and length and returns the substring. +func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interface{}, error) { + argEval, err := n.args[0].Evaluate(currentRow) + if err != nil { + return nil, err + } + if argEval == nil { + return nil, nil + } + stringArgOne, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } // this takes a sliding window approach to evaluate substring. - startIndex, err := evaluateIntArg(n.args[1], currentRow) + argEval, err = n.args[1].Evaluate(currentRow) if err != nil { - return nil, err + return 0, err + } + if argEval == nil { + return nil, nil } - if startIndex < 0 { + startIndex, ok := argEval.(int64) + if !ok { + return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + + if startIndex < 0 || startIndex >= int64(len(stringArgOne)) { return nil, sql3.NewErrValueOutOfRange(0, 0, startIndex) } - if startIndex >= len(stringArgOne) { - return nil, sql3.NewErrValueOutOfRange(0, 0, startIndex) - } - - endIndex := len(stringArgOne) + endIndex := int64(len(stringArgOne)) if len(n.args) > 2 { - ln, err := evaluateIntArg(n.args[2], currentRow) + argEval, err = n.args[2].Evaluate(currentRow) if err != nil { - return nil, err + return 0, err + } + if argEval == nil { + return nil, nil + } + ln, ok := argEval.(int64) + if !ok { + return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) } endIndex = startIndex + ln } - if endIndex < startIndex { - return nil, sql3.NewErrValueOutOfRange(0, 0, endIndex) - } - - if endIndex > len(stringArgOne) { + if endIndex < startIndex || endIndex > int64(len(stringArgOne)) { return nil, sql3.NewErrValueOutOfRange(0, 0, endIndex) } return stringArgOne[startIndex:endIndex], nil } -func (n *callPlanExpression) EvaluateLower(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) - if err != nil { - return nil, err - } - - return strings.ToLower(stringArgOne), nil -} - // takes string, findstring, replacestring. // replaces all occurances of findstring with replacestring func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } - stringArgTwo, err := evaluateStringArg(n.args[1], currentRow) + if argEval == nil { + return nil, nil + } + stringArgOne, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + argEval, err = n.args[1].Evaluate(currentRow) if err != nil { return nil, err } - stringArgThree, err := evaluateStringArg(n.args[2], currentRow) + if argEval == nil { + return nil, nil + } + stringArgTwo, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + argEval, err = n.args[2].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArgThree, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } return strings.ReplaceAll(stringArgOne, stringArgTwo, stringArgThree), nil } // takes a string, seperator and the position `n`, splits the string and returns n'th substring func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (interface{}, error) { - inputString, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } - seperator, err := evaluateStringArg(n.args[1], currentRow) + if argEval == nil { + return nil, nil + } + inputString, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + + argEval, err = n.args[1].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + seperator, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } + if len(n.args) == 2 { return strings.Split(inputString, seperator)[0], nil } - pos, err := evaluateIntArg(n.args[2], currentRow) + argEval, err = n.args[2].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + pos, ok := argEval.(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } res := strings.Split(inputString, seperator) if pos <= 0 { return res[0], nil - } else if len(res) > pos { + } else if int64(len(res)) > pos { return res[pos], nil } return "", nil } -func evaluateStringArg(n types.PlanExpression, currentRow []interface{}) (string, error) { - argOneEval, err := n.Evaluate(currentRow) - if err != nil { - return "", err - } - stringArgOne, ok := argOneEval.(string) - if !ok { - return "", sql3.NewErrInternalf("unexpected type converion %T", argOneEval) - } - - return stringArgOne, nil -} - -func evaluateIntArg(n types.PlanExpression, currentRow []interface{}) (int, error) { - argOneEval, err := n.Evaluate(currentRow) - if err != nil { - return 0, err - } - - intArgOne, ok := argOneEval.(int64) - if !ok { - return 0, sql3.NewErrInternalf("unexpected type converion %T", argOneEval) - } - - return int(intArgOne), nil -} - -// Analyze function for Trim/RTrim/LTrim -func (p *ExecutionPlanner) analyseFunctionTrim(call *parser.Call, scope parser.Statement) (parser.Expr, error) { - //one argument for Trim Functions - if len(call.Args) != 1 { - return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 1, len(call.Args)) - } - - if !typeIsString(call.Args[0].DataType()) { - return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) - } - - call.ResultDataType = parser.NewDataTypeString() - - return call, nil -} - // Execute Trim function to remove whitespaces from string func (n *callPlanExpression) EvaluateTrim(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } // Trim the whitespace from string - return strings.TrimSpace(stringArgOne), nil + return strings.TrimSpace(stringArg), nil } // Execute RTrim function to remove trailing whitespaces from string func (n *callPlanExpression) EvaluateRTrim(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } // Trim the trailing whitespace from string - return strings.TrimRight(stringArgOne, " "), nil + return strings.TrimRight(stringArg, " "), nil } // Execute LTrim function to remove leading whitespaces from string func (n *callPlanExpression) EvaluateLTrim(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArg, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } // Trim the leading whitespace from string - return strings.TrimLeft(stringArgOne, " "), nil -} - -func (p *ExecutionPlanner) analyseFunctionPrefixSuffix(call *parser.Call, scope parser.Statement) (parser.Expr, error) { - if len(call.Args) != 2 { - return nil, sql3.NewErrCallParameterCountMismatch(call.Rparen.Line, call.Rparen.Column, call.Name.Name, 2, len(call.Args)) - } - - if !typeIsString(call.Args[0].DataType()) { - return nil, sql3.NewErrStringExpressionExpected(call.Args[0].Pos().Line, call.Args[0].Pos().Column) - } - - if !typeIsInteger(call.Args[1].DataType()) { - return nil, sql3.NewErrIntExpressionExpected(call.Args[1].Pos().Line, call.Args[1].Pos().Column) - } - - call.ResultDataType = parser.NewDataTypeString() - - return call, nil + return strings.TrimLeft(stringArg, " "), nil } func (n *callPlanExpression) EvaluatePrefix(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArgOne, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } - intArgTwo, err := evaluateIntArg(n.args[1], currentRow) + argEval, err = n.args[1].Evaluate(currentRow) if err != nil { return nil, err } - - // if the length is less than zero, out of range - if intArgTwo < 0 { - return nil, sql3.NewErrValueOutOfRange(0, 0, intArgTwo) + if argEval == nil { + return nil, nil + } + intArgTwo, ok := argEval.(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) } - if intArgTwo > len(stringArgOne) { + if intArgTwo < 0 || intArgTwo > int64(len(stringArgOne)) { return nil, sql3.NewErrValueOutOfRange(0, 0, intArgTwo) } @@ -398,38 +529,54 @@ func (n *callPlanExpression) EvaluatePrefix(currentRow []interface{}) (interface } func (n *callPlanExpression) EvaluateSuffix(currentRow []interface{}) (interface{}, error) { - stringArgOne, err := evaluateStringArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + stringArgOne, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } - intArgTwo, err := evaluateIntArg(n.args[1], currentRow) + argEval, err = n.args[1].Evaluate(currentRow) if err != nil { return nil, err } + if argEval == nil { + return nil, nil + } + intArgTwo, ok := argEval.(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + } - // if the length is less than zero, out of range - if intArgTwo < 0 { + if intArgTwo < 0 || intArgTwo > int64(len(stringArgOne)) { return nil, sql3.NewErrValueOutOfRange(0, 0, intArgTwo) } - if intArgTwo > len(stringArgOne) { - return nil, sql3.NewErrValueOutOfRange(0, 0, intArgTwo) - } - - return stringArgOne[len(stringArgOne)-intArgTwo:], nil + return stringArgOne[int64(len(stringArgOne))-intArgTwo:], nil } func (n *callPlanExpression) EvaluateSpace(currentRow []interface{}) (interface{}, error) { // Get the integer argument from the function call - intArg, err := evaluateIntArg(n.args[0], currentRow) + argEval, err := n.args[0].Evaluate(currentRow) if err != nil { - return "", err + return nil, err + } + if argEval == nil { + return nil, nil + } + intArg, ok := argEval.(int64) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) } // Return a string containing a number of spaces equal to the integer value spaces := "" - for i := 0; i < intArg; i++ { + for i := int64(0); i < intArg; i++ { spaces += " " } return spaces, nil diff --git a/sql3/test/defs/defs_string_functions.go b/sql3/test/defs/defs_string_functions.go index 6c339ac2a..1790fd8e4 100644 --- a/sql3/test/defs/defs_string_functions.go +++ b/sql3/test/defs/defs_string_functions.go @@ -17,6 +17,32 @@ var stringScalarFunctionsTests = TableTest{ ), ), SQLTests: []SQLTest{ + { + name: "ReverseNull", + SQLs: sqls( + "select reverse(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, + { + name: "ReverseEmpty", + SQLs: sqls( + "select reverse('')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("")), + ), + Compare: CompareExactUnordered, + }, { name: "ReverseString", SQLs: sqls( @@ -43,6 +69,32 @@ var stringScalarFunctionsTests = TableTest{ ), Compare: CompareExactUnordered, }, + { + name: "SubstringNull", + SQLs: sqls( + "select substring(null, 1, 3)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, + { + name: "SubstringNullInt", + SQLs: sqls( + "select substring('some_string', null, 3)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, { name: "SubstringPositiveIndex", SQLs: sqls( @@ -56,32 +108,6 @@ var stringScalarFunctionsTests = TableTest{ ), Compare: CompareExactUnordered, }, - { - name: "StringSplitNoPos", - SQLs: sqls( - "select stringsplit('string,split', ',')", - ), - ExpHdrs: hdrs( - hdr("", fldTypeString), - ), - ExpRows: rows( - row(string("string")), - ), - Compare: CompareExactUnordered, - }, - { - name: "StringSplitPos", - SQLs: sqls( - "select stringsplit('string,split,now', stringsplit(',mid,', 'mid', 1), 2)", - ), - ExpHdrs: hdrs( - hdr("", fldTypeString), - ), - ExpRows: rows( - row(string("now")), - ), - Compare: CompareExactUnordered, - }, { name: "SubstringNegativeIndex", SQLs: sqls( @@ -122,6 +148,58 @@ var stringScalarFunctionsTests = TableTest{ ), Compare: CompareExactUnordered, }, + { + name: "StringSplitNull", + SQLs: sqls( + "select stringsplit(null, ',')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, + { + name: "StringSplitNoPos", + SQLs: sqls( + "select stringsplit('string,split', ',')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("string")), + ), + Compare: CompareExactUnordered, + }, + { + name: "StringSplitPos", + SQLs: sqls( + "select stringsplit('string,split,now', stringsplit(',mid,', 'mid', 1), 2)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("now")), + ), + Compare: CompareExactUnordered, + }, + { + name: "CharNull", + SQLs: sqls( + "select char(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, { name: "CharInt", SQLs: sqls( @@ -142,6 +220,59 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "integer expression expected", }, + { + name: "ASCIINull", + SQLs: sqls( + "select ascii(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeInt), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, + { + name: "ASCIILengthMisMatch", + SQLs: sqls( + "select ascii('longer')", + ), + ExpErr: "[0:0] value 'longer' should be of the length 1", + }, + { + name: "ASCIIString", + SQLs: sqls( + "select ascii('R')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeInt), + ), + ExpRows: rows( + row(int64(82)), + ), + Compare: CompareExactUnordered, + }, + { + name: "ASCIIInt", + SQLs: sqls( + "select ascii(32)", + ), + ExpErr: "string expression expected", + }, + { + name: "UpperNull", + SQLs: sqls( + "select upper(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "ConvertingStringtoUpper", SQLs: sqls( @@ -169,6 +300,19 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "string expression expected", }, + { + name: "LowerNull", + SQLs: sqls( + "select lower(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "StringLower", SQLs: sqls( @@ -196,6 +340,32 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "string expression expected", }, + { + name: "ReplaceAllNullString", + SQLs: sqls( + "select replaceall(null,'data','feature')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, + { + name: "ReplaceAllNullArg", + SQLs: sqls( + "select replaceall('hello database',null,'feature')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactUnordered, + }, { name: "ReplaceAllString", SQLs: sqls( @@ -249,6 +419,19 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "string expression expected", }, + { + name: "TrimNull", + SQLs: sqls( + "select trim(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "RemovingWhitespacefromStringusingTrim", SQLs: sqls( @@ -277,6 +460,19 @@ var stringScalarFunctionsTests = TableTest{ ExpErr: "string expression expected", }, //Prefix() + { + name: "PrefixNull", + SQLs: sqls( + "SELECT PREFIX(NULL, 34)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "IncorrectArgumentsforPrefix", SQLs: sqls( @@ -345,6 +541,19 @@ var stringScalarFunctionsTests = TableTest{ Compare: CompareExactOrdered, }, //Suffix() + { + name: "SuffixNull", + SQLs: sqls( + "SELECT SUFFIX(NULL, 23)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "IncorrectArgumentsforSuffix", SQLs: sqls( @@ -412,6 +621,19 @@ var stringScalarFunctionsTests = TableTest{ ), Compare: CompareExactOrdered, }, + { + name: "RTrimNull", + SQLs: sqls( + "select rtrim(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "RemovingTrailingspacefromStringusingRTrim", SQLs: sqls( @@ -439,6 +661,19 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "string expression expected", }, + { + name: "LTrimNull", + SQLs: sqls( + "select ltrim(null)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, { name: "RemovingLeadingspacefromStringusingLTrim", SQLs: sqls(