Bug fix round up (fb-1841, fb-1819, fb-1867) (#2396)

* check root operator after optimize

* round of bug fixes

(cherry picked from commit 9f042216a7)
This commit is contained in:
pokeeffe-molecula 2023-01-05 20:13:26 -06:00 committed by Joe Friedrich
parent 01cae92d96
commit 4e43b767df
7 changed files with 282 additions and 266 deletions

View file

@ -185,6 +185,8 @@ func (s *Scanner) scanString() (Pos, Token, string) {
ch, _ := s.read()
if ch == -1 {
return pos, ILLEGAL, `'` + s.buf.String()
} else if ch == '\n' {
return pos, UNTERMSTRING, `'` + s.buf.String()
} else if ch == '\'' {
if s.peek() == '\'' { // escaped quote
s.read()

View file

@ -29,6 +29,7 @@ const (
ILLEGAL Token = iota
EOF
SPACE
UNTERMSTRING
literal_beg
IDENT // IDENT
@ -253,9 +254,10 @@ const (
)
var tokens = [...]string{
ILLEGAL: "ILLEGAL",
EOF: "EOF",
SPACE: "SPACE",
ILLEGAL: "ILLEGAL",
EOF: "EOF",
SPACE: "SPACE",
UNTERMSTRING: "unterminated string literal",
IDENT: "IDENT",
VARIABLE: "VARIABLE",

View file

@ -152,7 +152,7 @@ func (p *ExecutionPlanner) compileBulkInsertStatement(stmt *parser.BulkInsertSta
}
}
return NewPlanOpBulkInsert(p, tableName, options), nil
return NewPlanOpQuery(p, NewPlanOpBulkInsert(p, tableName, options), p.sql), nil
}
// analyzeBulkInsertStatement analyzes a BULK INSERT statement and returns an

View file

@ -68,7 +68,7 @@ func (p *ExecutionPlanner) compileInsertStatement(stmt *parser.InsertStatement)
insertValues = append(insertValues, tupleValues)
}
return NewPlanOpInsert(p, tableName, targetColumns, insertValues), nil
return NewPlanOpQuery(p, NewPlanOpInsert(p, tableName, targetColumns, insertValues), p.sql), nil
}
// analyzeInsertStatement analyzes an INSERT statement and returns and error if

View file

@ -460,274 +460,280 @@ func (i *bulkInsertSourceNDJsonRowIter) Next(ctx context.Context) (types.Row, er
}
}
if i.reader.Scan() {
if err := i.reader.Err(); err != nil {
return nil, err
}
jsonValue := i.reader.Text()
// now we do the mapping to the output row
result := make([]interface{}, len(i.options.mapExpressions))
// parse the json
v := interface{}(nil)
err := json.Unmarshal([]byte(jsonValue), &v)
if err != nil {
return nil, sql3.NewErrParsingJSON(0, 0, jsonValue, err.Error())
}
// type check against the output type of the map operation
for idx, expr := range i.pathExpressions {
evalValue, err := expr(ctx, v)
if err != nil {
if i.options.allowMissingValues && strings.HasPrefix(err.Error(), "unknown key") {
evalValue = nil
} else {
return nil, sql3.NewErrEvaluatingJSONPathExpr(0, 0, i.mapExpressionResults[idx], jsonValue, err.Error())
}
for {
if i.reader.Scan() {
if err := i.reader.Err(); err != nil {
return nil, err
}
// if nil (null) then return nil
if evalValue == nil {
result[idx] = nil
jsonValue := i.reader.Text()
jsonValue = strings.TrimSpace(jsonValue)
if len(jsonValue) == 0 {
continue
}
mapColumn := i.options.mapExpressions[idx]
switch mapColumn.colType.(type) {
case *parser.DataTypeID, *parser.DataTypeInt:
// now we do the mapping to the output row
result := make([]interface{}, len(i.options.mapExpressions))
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int
if v == float64(int64(v)) {
result[idx] = int64(v)
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
// parse the json
v := interface{}(nil)
err := json.Unmarshal([]byte(jsonValue), &v)
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = intVal
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeIDSet:
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int, and then turn that into an idset
if v == float64(int64(v)) {
result[idx] = []int64{int64(v)}
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
setValue := make([]int64, 0)
for _, i := range v {
switch v := i.(type) {
case float64:
if v == float64(int64(v)) {
setValue = append(setValue, int64(v))
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
setValue = append(setValue, int64(intVal))
default:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
}
result[idx] = setValue
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = []int64{int64(intVal)}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeStringSet:
switch v := evalValue.(type) {
case float64:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case []interface{}:
setValue := make([]string, 0)
for _, i := range v {
f, ok := i.(string)
if !ok {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
setValue = append(setValue, f)
}
result[idx] = setValue
case string:
result[idx] = []string{v}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeTimestamp:
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int
if v == float64(int64(v)) {
result[idx] = time.UnixMilli(int64(v)).UTC()
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
if tm, err := time.ParseInLocation(time.RFC3339Nano, v, time.UTC); err == nil {
result[idx] = tm
} else if tm, err := time.ParseInLocation(time.RFC3339, v, time.UTC); err == nil {
result[idx] = tm
} else if tm, err := time.ParseInLocation("2006-01-02", v, time.UTC); err == nil {
result[idx] = tm
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeString:
switch v := evalValue.(type) {
case float64:
// if a whole number make it an int
if v == float64(int64(v)) {
result[idx] = fmt.Sprintf("%d", int64(v))
} else {
result[idx] = fmt.Sprintf("%f", v)
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
result[idx] = v
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeBool:
switch v := evalValue.(type) {
case float64:
// if a whole number make it an int, and convert to a bool
if v == float64(int64(v)) {
result[idx] = v > 0
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case bool:
result[idx] = v
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeDecimal:
switch v := evalValue.(type) {
case float64:
result[idx] = pql.FromFloat64(v)
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
// try to parse from a string
dv, err := pql.ParseDecimal(v)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = dv
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", mapColumn.colType)
if err != nil {
return nil, sql3.NewErrParsingJSON(0, 0, jsonValue, err.Error())
}
// type check against the output type of the map operation
for idx, expr := range i.pathExpressions {
evalValue, err := expr(ctx, v)
if err != nil {
if i.options.allowMissingValues && (strings.HasPrefix(err.Error(), "unknown key") || strings.HasPrefix(err.Error(), "unknown parameter")) {
evalValue = nil
} else {
return nil, sql3.NewErrEvaluatingJSONPathExpr(0, 0, i.mapExpressionResults[idx], jsonValue, err.Error())
}
}
// if nil (null) then return nil
if evalValue == nil {
result[idx] = nil
continue
}
mapColumn := i.options.mapExpressions[idx]
switch mapColumn.colType.(type) {
case *parser.DataTypeID, *parser.DataTypeInt:
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int
if v == float64(int64(v)) {
result[idx] = int64(v)
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = intVal
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeIDSet:
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int, and then turn that into an idset
if v == float64(int64(v)) {
result[idx] = []int64{int64(v)}
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
setValue := make([]int64, 0)
for _, i := range v {
switch v := i.(type) {
case float64:
if v == float64(int64(v)) {
setValue = append(setValue, int64(v))
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
setValue = append(setValue, int64(intVal))
default:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
}
result[idx] = setValue
case string:
intVal, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = []int64{int64(intVal)}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeStringSet:
switch v := evalValue.(type) {
case float64:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case []interface{}:
setValue := make([]string, 0)
for _, i := range v {
f, ok := i.(string)
if !ok {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
setValue = append(setValue, f)
}
result[idx] = setValue
case string:
result[idx] = []string{v}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeTimestamp:
switch v := evalValue.(type) {
case float64:
// if v is a whole number then make it an int
if v == float64(int64(v)) {
result[idx] = time.UnixMilli(int64(v)).UTC()
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
if tm, err := time.ParseInLocation(time.RFC3339Nano, v, time.UTC); err == nil {
result[idx] = tm
} else if tm, err := time.ParseInLocation(time.RFC3339, v, time.UTC); err == nil {
result[idx] = tm
} else if tm, err := time.ParseInLocation("2006-01-02", v, time.UTC); err == nil {
result[idx] = tm
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeString:
switch v := evalValue.(type) {
case float64:
// if a whole number make it an int
if v == float64(int64(v)) {
result[idx] = fmt.Sprintf("%d", int64(v))
} else {
result[idx] = fmt.Sprintf("%f", v)
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
result[idx] = v
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeBool:
switch v := evalValue.(type) {
case float64:
// if a whole number make it an int, and convert to a bool
if v == float64(int64(v)) {
result[idx] = v > 0
} else {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case bool:
result[idx] = v
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
case *parser.DataTypeDecimal:
switch v := evalValue.(type) {
case float64:
result[idx] = pql.FromFloat64(v)
case []interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case string:
// try to parse from a string
dv, err := pql.ParseDecimal(v)
if err != nil {
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
}
result[idx] = dv
case bool:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
case interface{}:
return nil, sql3.NewErrTypeConversionOnMap(0, 0, v, mapColumn.colType.TypeDescription())
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", evalValue)
}
default:
return nil, sql3.NewErrInternalf("unhandled type '%T'", mapColumn.colType)
}
}
return result, nil
}
return result, nil
return nil, types.ErrNoMoreRows
}
return nil, types.ErrNoMoreRows
}
func (i *bulkInsertSourceNDJsonRowIter) Close(ctx context.Context) {

View file

@ -104,6 +104,12 @@ func (p *ExecutionPlanner) optimizePlan(ctx context.Context, plan types.PlanOper
"--------------------------------------------------------------------------------",
)
// check that result is a PlanOpQuery
_, ok := result.(*PlanOpQuery)
if !ok {
return nil, sql3.NewErrInternalf("unexpected root operator type '%T'", result)
}
return result, nil
}

View file

@ -1783,7 +1783,7 @@ func TestPlanner_BulkInsert(t *testing.T) {
'petalWidth' DECIMAL,
'species' STRING)
from
'{"id": 1, "sepalLength": "5.1", "sepalWidth": "3.5", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
x'{"id": 1, "sepalLength": "5.1", "sepalWidth": "3.5", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
{"id": 2, "sepalLength": "4.9", "sepalWidth": "3.0", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
{"id": 3, "sepalLength": "4.7", "sepalWidth": "3.2", "petalLength": "1.3", "petalWidth": "0.2", "species": "setosa"}'
with
@ -1802,7 +1802,7 @@ func TestPlanner_BulkInsert(t *testing.T) {
'petalWidth' DECIMAL(2),
'species' STRING)
from
'{"id": 1, "sepalLength": "5.1", "sepalWidth": "3.5", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
x'{"id": 1, "sepalLength": "5.1", "sepalWidth": "3.5", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
{"id": 2, "sepalLength": "4.9", "sepalWidth": "3.0", "petalLength": "1.4", "petalWidth": "0.2", "species": "setosa"}
{"id": 3, "sepalLength": "4.7", "sepalWidth": "3.2", "petalLength": "1.3", "petalWidth": "0.2", "species": "setosa"}'
with
@ -1940,7 +1940,7 @@ func TestPlanner_BulkInsert(t *testing.T) {
@5,
@6,
@7)
FROM '{"id_col": "3", "string_col": "TEST", "decimal_col": "1.12", "bool_col": false, "time_col": "2013-07-15T01:18:46Z", "stringset_col": "stringset1","ideset_col": 1}
FROM x'{"id_col": "3", "string_col": "TEST", "decimal_col": "1.12", "bool_col": false, "time_col": "2013-07-15T01:18:46Z", "stringset_col": "stringset1","ideset_col": 1}
{"id_col": "4", "string_col": "TEST2", "decimal_col": "1.12", "bool_col": false, "time_col": "2013-07-15T01:18:46Z", "stringset_col": ["stringset1","stringset3"],"ideset_col": [1,2]}
{"id_col": "5", "string_col": "TEST", "int_col": "321", "decimal_col": "12.1", "bool_col": 1, "time_col": "2014-07-15T01:18:46Z", "stringset_col": "stringset2","ideset_col": [1,3]}'
with