From 188b61b3cd934bd0f3e3781c2dfd49b0e778d0f3 Mon Sep 17 00:00:00 2001 From: rachithrr Date: Fri, 27 Jan 2023 11:04:39 -0600 Subject: [PATCH] FB-1817: Implement FORMAT() (#2220) --- sql3/errors.go | 20 +++- sql3/planner/expression.go | 2 + sql3/planner/expressionanalyzercall.go | 2 + sql3/planner/inbuiltfunctionsstring.go | 121 +++++++++++++++++------- sql3/test/defs/defs_string_functions.go | 81 ++++++++++++++++ 5 files changed, 190 insertions(+), 36 deletions(-) diff --git a/sql3/errors.go b/sql3/errors.go index 6e069bd18..895dc0139 100644 --- a/sql3/errors.go +++ b/sql3/errors.go @@ -50,6 +50,7 @@ const ( ErrIntegerLiteral errors.Code = "ErrIntegerLiteral" ErrStringLiteral errors.Code = "ErrStringLiteral" ErrBoolLiteral errors.Code = "ErrBoolLiteral" + ErrLiteralNullNotAllowed errors.Code = "ErrLiteralNullNotAllowed" ErrLiteralEmptySetNotAllowed errors.Code = "ErrLiteralEmptySetNotAllowed" ErrLiteralEmptyTupleNotAllowed errors.Code = "ErrLiteralEmptyTupleNotAllowed" ErrSetLiteralMustContainIntOrString errors.Code = "ErrSetLiteralMustContainIntOrString" @@ -121,8 +122,9 @@ const ( ErrAggregateNotAllowedInGroupBy errors.Code = "ErrIdPercentileNotAllowedInGroupBy" // function evaluation - ErrValueOutOfRange errors.Code = "ErrValueOutOfRange" - ErrStringLengthMismatch errors.Code = "ErrStringLengthMismatch" + ErrValueOutOfRange errors.Code = "ErrValueOutOfRange" + ErrStringLengthMismatch errors.Code = "ErrStringLengthMismatch" + ErrUnexpectedTypeConversion errors.Code = "ErrUnexpectedTypeConversion" ) func NewErrDuplicateColumn(line int, col int, column string) error { @@ -268,6 +270,13 @@ func NewErrSetLiteralMustContainIntOrString(line, col int) error { ) } +func NewErrLiteralNullNotAllowed(line, col int) error { + return errors.New( + ErrLiteralNullNotAllowed, + fmt.Sprintf("[%d:%d] null literal not allowed", line, col), + ) +} + func NewErrInvalidColumnInFilterExpression(line, col int, column string, op string) error { return errors.New( ErrInvalidColumnInFilterExpression, @@ -751,3 +760,10 @@ func NewErrStringLengthMismatch(line, col, len int, val interface{}) error { fmt.Sprintf("[%d:%d] value '%v' should be of the length %d", line, col, val, len), ) } + +func NewErrUnexpectedTypeConversion(line, col int, val interface{}) error { + return errors.New( + ErrUnexpectedTypeConversion, + NewErrInternalf("unexpected type conversion %T", val).Error(), + ) +} diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index 01926400c..892b9f176 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -1520,6 +1520,8 @@ func (n *callPlanExpression) Evaluate(currentRow []interface{}) (interface{}, er return n.EvaluateLen(currentRow) case "REPLICATE": return n.EvaluateReplicate(currentRow) + case "FORMAT": + return n.EvaluateFormat(currentRow) default: return nil, sql3.NewErrInternalf("unhandled function name '%s'", n.name) } diff --git a/sql3/planner/expressionanalyzercall.go b/sql3/planner/expressionanalyzercall.go index 6dce440f2..f238a2310 100644 --- a/sql3/planner/expressionanalyzercall.go +++ b/sql3/planner/expressionanalyzercall.go @@ -273,6 +273,8 @@ func (p *ExecutionPlanner) analyzeCallExpression(call *parser.Call, scope parser return p.analyseFunctionLen(call, scope) case "REPLICATE": return p.analyseFunctionPrefixSuffixReplicate(call, scope) + case "FORMAT": + return p.analyseFunctionFormat(call, scope) default: return nil, sql3.NewErrCallUnknownFunction(call.Name.NamePos.Line, call.Name.NamePos.Column, call.Name.Name) } diff --git a/sql3/planner/inbuiltfunctionsstring.go b/sql3/planner/inbuiltfunctionsstring.go index ccc31b90c..2b2e91b11 100644 --- a/sql3/planner/inbuiltfunctionsstring.go +++ b/sql3/planner/inbuiltfunctionsstring.go @@ -1,6 +1,7 @@ package planner import ( + "fmt" "strings" "github.com/featurebasedb/featurebase/v3/sql3" @@ -207,7 +208,27 @@ func (p *ExecutionPlanner) analyseFunctionLen(call *parser.Call, scope parser.St return call, nil } -// reverses the string +// format(format_string, args...) +func (p *ExecutionPlanner) analyseFunctionFormat(call *parser.Call, scope parser.Statement) (parser.Expr, error) { + // should have at least one argument to check call.args[0] string + 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) + } + + for i := 1; i < len(call.Args); i++ { + if typeIsVoid(call.Args[i].DataType()) { + return nil, sql3.NewErrLiteralNullNotAllowed(call.Args[i].Pos().Line, call.Args[i].Pos().Column) + } + } + call.ResultDataType = parser.NewDataTypeString() + return call, nil +} + +// EvaluateReverse reverses the string func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -218,7 +239,7 @@ func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interfac } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // reverse the string @@ -239,13 +260,13 @@ func (n *callPlanExpression) EvaluateLower(currentRow []interface{}) (interface{ } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } return strings.ToLower(stringArg), nil } -// Convert string to Upper case +// EvaluateUpper converts string to Upper case func (n *callPlanExpression) EvaluateUpper(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -256,7 +277,7 @@ func (n *callPlanExpression) EvaluateUpper(currentRow []interface{}) (interface{ } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // convert to Upper @@ -274,7 +295,7 @@ func (n *callPlanExpression) EvaluateChar(currentRow []interface{}) (interface{} } intArg, ok := argEval.(int64) if !ok { - return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return 0, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // ascii range is [0-255] if intArg < 0 || intArg > 255 { @@ -285,7 +306,7 @@ func (n *callPlanExpression) EvaluateChar(currentRow []interface{}) (interface{} return string(rune(intArg)), nil } -// this takes a string and returns the ascii value. +// EvaluateAscii 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 @@ -298,7 +319,7 @@ func (n *callPlanExpression) EvaluateAscii(currentRow []interface{}) (interface{ } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if len(stringArg) == 0 { @@ -313,7 +334,7 @@ func (n *callPlanExpression) EvaluateAscii(currentRow []interface{}) (interface{ return int64(res[0]), nil } -// Takes string, startIndex and length and returns the substring. +// EvaluateSubstring 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 { @@ -324,7 +345,7 @@ func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interf } stringArgOne, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // this takes a sliding window approach to evaluate substring. @@ -338,7 +359,7 @@ func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interf startIndex, ok := argEval.(int64) if !ok { - return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return 0, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if startIndex < 0 || startIndex >= int64(len(stringArgOne)) { @@ -356,7 +377,7 @@ func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interf } ln, ok := argEval.(int64) if !ok { - return 0, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return 0, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } endIndex = startIndex + ln } @@ -368,7 +389,7 @@ func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interf return stringArgOne[startIndex:endIndex], nil } -// takes string, findstring, replacestring. +// EvaluateReplaceAll takes string, findstring, replacestring. // replaces all occurances of findstring with replacestring func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) @@ -380,7 +401,7 @@ func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (inter } stringArgOne, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[1].Evaluate(currentRow) if err != nil { @@ -391,7 +412,7 @@ func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (inter } stringArgTwo, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[2].Evaluate(currentRow) if err != nil { @@ -402,12 +423,12 @@ func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (inter } stringArgThree, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } return strings.ReplaceAll(stringArgOne, stringArgTwo, stringArgThree), nil } -// takes a string, seperator and the position `n`, splits the string and returns n'th substring +// EvaluateStringSplit takes a string, seperator and the position `n`, splits the string and returns n'th substring func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -418,7 +439,7 @@ func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (inte } inputString, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[1].Evaluate(currentRow) @@ -430,7 +451,7 @@ func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (inte } seperator, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if len(n.args) == 2 { @@ -445,7 +466,7 @@ func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (inte } pos, ok := argEval.(int64) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } res := strings.Split(inputString, seperator) @@ -457,7 +478,7 @@ func (n *callPlanExpression) EvaluateStringSplit(currentRow []interface{}) (inte return "", nil } -// Execute Trim function to remove whitespaces from string +// EvaluateTrim function to remove whitespaces from string func (n *callPlanExpression) EvaluateTrim(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -468,14 +489,14 @@ func (n *callPlanExpression) EvaluateTrim(currentRow []interface{}) (interface{} } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // Trim the whitespace from string return strings.TrimSpace(stringArg), nil } -// Execute RTrim function to remove trailing whitespaces from string +// EvaluateRTrim function removes trailing whitespaces from string func (n *callPlanExpression) EvaluateRTrim(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -486,14 +507,14 @@ func (n *callPlanExpression) EvaluateRTrim(currentRow []interface{}) (interface{ } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // Trim the trailing whitespace from string return strings.TrimRight(stringArg, " "), nil } -// Execute LTrim function to remove leading whitespaces from string +// EvaluateLTrim function removes leading whitespaces from string func (n *callPlanExpression) EvaluateLTrim(currentRow []interface{}) (interface{}, error) { argEval, err := n.args[0].Evaluate(currentRow) if err != nil { @@ -504,7 +525,7 @@ func (n *callPlanExpression) EvaluateLTrim(currentRow []interface{}) (interface{ } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // Trim the leading whitespace from string @@ -521,7 +542,7 @@ func (n *callPlanExpression) EvaluatePrefix(currentRow []interface{}) (interface } stringArgOne, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[1].Evaluate(currentRow) @@ -533,7 +554,7 @@ func (n *callPlanExpression) EvaluatePrefix(currentRow []interface{}) (interface } intArgTwo, ok := argEval.(int64) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if intArgTwo < 0 || intArgTwo > int64(len(stringArgOne)) { @@ -553,7 +574,7 @@ func (n *callPlanExpression) EvaluateSuffix(currentRow []interface{}) (interface } stringArgOne, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[1].Evaluate(currentRow) @@ -565,7 +586,7 @@ func (n *callPlanExpression) EvaluateSuffix(currentRow []interface{}) (interface } intArgTwo, ok := argEval.(int64) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if intArgTwo < 0 || intArgTwo > int64(len(stringArgOne)) { @@ -586,7 +607,7 @@ func (n *callPlanExpression) EvaluateSpace(currentRow []interface{}) (interface{ } intArg, ok := argEval.(int64) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } // Return a string containing a number of spaces equal to the integer value @@ -607,7 +628,7 @@ func (n *callPlanExpression) EvaluateLen(currentRow []interface{}) (interface{}, } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } return int64(len([]rune(stringArg))), nil @@ -622,7 +643,7 @@ func (n *callPlanExpression) EvaluateReplicate(currentRow []interface{}) (interf } stringArg, ok := argEval.(string) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } argEval, err = n.args[1].Evaluate(currentRow) @@ -634,7 +655,7 @@ func (n *callPlanExpression) EvaluateReplicate(currentRow []interface{}) (interf } intArg, ok := argEval.(int64) if !ok { - return nil, sql3.NewErrInternalf("unexpected type converion %T", argEval) + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) } if intArg < 0 { @@ -644,3 +665,35 @@ func (n *callPlanExpression) EvaluateReplicate(currentRow []interface{}) (interf return strings.Repeat(stringArg, int(intArg)), nil } + +// EvaluateFormat function formats according to a format specifier and returns resulting string. +func (n *callPlanExpression) EvaluateFormat(currentRow []interface{}) (interface{}, error) { + // first arg must be a string + argEval, err := n.args[0].Evaluate(currentRow) + if err != nil { + return nil, err + } + if argEval == nil { + return nil, nil + } + formatString, ok := argEval.(string) + if !ok { + return nil, sql3.NewErrUnexpectedTypeConversion(0, 0, argEval) + } + + var args []interface{} + + // loop, since args can be of any length. + for _, arg := range n.args[1:] { + argEval, err = arg.Evaluate(currentRow) + if err != nil { + return nil, err + } + if argEval == nil { + // this should never happen. + return nil, sql3.NewErrLiteralNullNotAllowed(0, 0) + } + args = append(args, argEval) + } + return fmt.Sprintf(formatString, args...), nil +} diff --git a/sql3/test/defs/defs_string_functions.go b/sql3/test/defs/defs_string_functions.go index 916844b9a..d52e87693 100644 --- a/sql3/test/defs/defs_string_functions.go +++ b/sql3/test/defs/defs_string_functions.go @@ -820,5 +820,86 @@ var stringScalarFunctionsTests = TableTest{ ), ExpErr: "[0:0] value '-1' out of range", }, + { + name: "FormatString", + SQLs: sqls( + "select format('this or %s', 'that')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("this or that")), + ), + Compare: CompareExactOrdered, + }, + { + name: "FormatBoolean", + SQLs: sqls( + "select format('is this %t?', true)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("is this true?")), + ), + Compare: CompareExactOrdered, + }, + { + name: "FormatInteger", + SQLs: sqls( + "select format('%d > %d', 11 , 9)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("11 > 9")), + ), + Compare: CompareExactOrdered, + }, + { + name: "FormatNullString", + SQLs: sqls( + "select format(null,'this')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(nil), + ), + Compare: CompareExactOrdered, + }, + { + name: "FormatNullArgument", + SQLs: sqls( + "select format('format = %d', null)", + ), + ExpErr: "[1:30] null literal not allowed", + Compare: CompareExactOrdered, + }, + { + name: "FormatLengthZero", + SQLs: sqls( + "select format()", + ), + ExpErr: "[1:15] 'format': count of formal parameters (1) does not match count of actual parameters (0)", + Compare: CompareExactOrdered, + }, + { + name: "FormatLengthOne", + SQLs: sqls( + "select format('noArg')", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("noArg")), + ), + Compare: CompareExactOrdered, + }, }, }