From a4f18fb25f1dc4b35369c89c8d55ad001aafe2e4 Mon Sep 17 00:00:00 2001 From: rachithrr Date: Thu, 15 Dec 2022 01:11:29 +0530 Subject: [PATCH] FB-1812: implement stringsplit() (#2362) --- sql3/planner/expression.go | 2 + sql3/planner/expressionanalyzercall.go | 2 + sql3/planner/inbuiltfunctionsstring.go | 50 +++++++++++++++++++++++++ sql3/test/defs/defs_string_functions.go | 26 +++++++++++++ 4 files changed, 80 insertions(+) diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index bf89d9f90..d41c3d670 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -1484,6 +1484,8 @@ func (n *callPlanExpression) Evaluate(currentRow []interface{}) (interface{}, er return n.EvaluateReverse(currentRow) case "UPPER": return n.EvaluateUpper(currentRow) + case "STRINGSPLIT": + return n.EvaluateStringSplit(currentRow) case "SUBSTRING": return n.EvaluateSubstring(currentRow) case "LOWER": diff --git a/sql3/planner/expressionanalyzercall.go b/sql3/planner/expressionanalyzercall.go index 3aa315e53..bb13fc548 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 "UPPER": return p.analyzeFunctionUpper(call, scope) + case "STRINGSPLIT": + return p.analyseFunctionStringSplit(call, scope) case "SUBSTRING": return p.analyseFunctionSubstring(call, scope) case "LOWER": diff --git a/sql3/planner/inbuiltfunctionsstring.go b/sql3/planner/inbuiltfunctionsstring.go index 24b84510e..9ad34270b 100644 --- a/sql3/planner/inbuiltfunctionsstring.go +++ b/sql3/planner/inbuiltfunctionsstring.go @@ -53,6 +53,29 @@ func (p *ExecutionPlanner) analyzeFunctionLower(call *parser.Call, scope parser. 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() @@ -188,6 +211,33 @@ func (n *callPlanExpression) EvaluateReplaceAll(currentRow []interface{}) (inter 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) + if err != nil { + return nil, err + } + seperator, err := evaluateStringArg(n.args[1], currentRow) + if err != nil { + return nil, err + } + if len(n.args) == 2 { + return strings.Split(inputString, seperator)[0], nil + } + pos, err := strconv.Atoi(n.args[2].String()) + if err != nil { + return nil, sql3.NewErrInternalf("unexpected type converion %T", n.args[2]) + } + + res := strings.Split(inputString, seperator) + if pos <= 0 { + return res[0], nil + } else if 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 { diff --git a/sql3/test/defs/defs_string_functions.go b/sql3/test/defs/defs_string_functions.go index b616196f2..a407e4afc 100644 --- a/sql3/test/defs/defs_string_functions.go +++ b/sql3/test/defs/defs_string_functions.go @@ -56,6 +56,32 @@ 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(