diff --git a/sql3/planner/expression.go b/sql3/planner/expression.go index e81a59860..828ebd054 100644 --- a/sql3/planner/expression.go +++ b/sql3/planner/expression.go @@ -1478,6 +1478,8 @@ func (n *callPlanExpression) Evaluate(currentRow []interface{}) (interface{}, er return n.EvaluateDatepart(currentRow) case "REVERSE": return n.EvaluateReverse(currentRow) + case "SUBSTRING": + return n.EvaluateSubstring(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 5ad8bdb91..5bb692f5c 100644 --- a/sql3/planner/expressionanalyzercall.go +++ b/sql3/planner/expressionanalyzercall.go @@ -243,6 +243,8 @@ func (p *ExecutionPlanner) analyzeCallExpression(call *parser.Call, scope parser return p.analyzeFunctionSubtable(call, scope) case "REVERSE": return p.analyseFunctionReverse(call, scope) + case "SUBSTRING": + return p.analyseFunctionSubstring(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 71143a96c..33958d636 100644 --- a/sql3/planner/inbuiltfunctionsstring.go +++ b/sql3/planner/inbuiltfunctionsstring.go @@ -1,6 +1,8 @@ package planner import ( + "strconv" + "github.com/molecula/featurebase/v3/sql3" "github.com/molecula/featurebase/v3/sql3/parser" ) @@ -20,6 +22,27 @@ func (p *ExecutionPlanner) analyseFunctionReverse(call *parser.Call, scope parse 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()) { + 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()) { + return nil, sql3.NewErrIntExpressionExpected(call.Args[i].Pos().Line, call.Args[i].Pos().Column) + } + } + + call.ResultDataType = parser.NewDataTypeString() + + return call, nil +} + // reverses the string func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interface{}, error) { argOneEval, err := n.args[0].Evaluate(currentRow) @@ -39,3 +62,47 @@ func (n *callPlanExpression) EvaluateReverse(currentRow []interface{}) (interfac } return string(runes), nil } + +// Takes string, startIndex and length and returns the substring. +func (n *callPlanExpression) EvaluateSubstring(currentRow []interface{}) (interface{}, error) { + argOneEval, err := n.args[0].Evaluate(currentRow) + if err != nil { + return nil, err + } + + stringArgOne, ok := argOneEval.(string) + if !ok { + return nil, sql3.NewErrInternalf("unexpected type converion %T", argOneEval) + } + + // this takes a sliding window approach to evaluate substring. + startIndex, err := strconv.Atoi(n.args[1].String()) + if err != nil { + return nil, sql3.NewErrInternalf("unexpected type converion %T", n.args[1]) + } + if startIndex >= len(stringArgOne) { + return "", nil + } + + endIndex := len(stringArgOne) + if len(n.args) > 2 { + ln, err := strconv.Atoi(n.args[2].String()) + if err != nil { + return nil, sql3.NewErrInternalf("unexpected type converion %T", n.args[1]) + } + endIndex = startIndex + ln + } + if endIndex < 0 { + return "", nil + } + + if startIndex < 0 { + startIndex = 0 + } + + if endIndex > len(stringArgOne) { + return stringArgOne[startIndex:], nil + } + + return stringArgOne[startIndex:endIndex], nil +} diff --git a/sql3/test/defs/defs_string_functions.go b/sql3/test/defs/defs_string_functions.go index 4e4fc560c..d186cc887 100644 --- a/sql3/test/defs/defs_string_functions.go +++ b/sql3/test/defs/defs_string_functions.go @@ -17,6 +17,7 @@ var stringScalarFunctionsTests = TableTest{ ), SQLTests: []SQLTest{ { + name: "ReverseString", SQLs: sqls( "select reverse('this')", ), @@ -28,5 +29,83 @@ var stringScalarFunctionsTests = TableTest{ ), Compare: CompareExactUnordered, }, + { + name: "ReverseReverseString", + SQLs: sqls( + "select reverse(reverse('this'))", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("this")), + ), + Compare: CompareExactUnordered, + }, + { + name: "SubstringPositiveIndex", + SQLs: sqls( + "select substring('testing', 1, 3)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("est")), + ), + Compare: CompareExactUnordered, + }, + { + name: "SubstringNegativeIndex", + SQLs: sqls( + "select substring('testing', -10, 14)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("test")), + ), + Compare: CompareExactUnordered, + }, + { + name: "SubstringNoLength", + SQLs: sqls( + "select substring('testing', -5)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("testing")), + ), + Compare: CompareExactUnordered, + }, + { + name: "ReverseSubstring", + SQLs: sqls( + "select reverse(substring('testing', 0))", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("gnitset")), + ), + Compare: CompareExactUnordered, + }, + { + name: "SubstringReverse", + SQLs: sqls( + "select substring(reverse('testing'), 3)", + ), + ExpHdrs: hdrs( + hdr("", fldTypeString), + ), + ExpRows: rows( + row(string("tset")), + ), + Compare: CompareExactUnordered, + }, }, }