Merge pull request #1953 from molecula/sup-156

[SUP-156] Add SQL SELECT mapping test
This commit is contained in:
Ben Johnson 2022-02-28 12:55:35 -07:00 committed by GitHub
commit 0b5a0d449e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 64 additions and 10 deletions

View file

@ -3,10 +3,13 @@ package sql_test
import (
"context"
"math"
"testing"
"github.com/molecula/featurebase/v3"
"github.com/molecula/featurebase/v3/sql"
"github.com/molecula/featurebase/v3/test"
"vitess.io/vitess/go/vt/sqlparser"
)
func TestHandler(t *testing.T) {
@ -28,3 +31,52 @@ func TestHandler(t *testing.T) {
}
}
func TestSelectHandler_MapSelect(t *testing.T) {
cluster := test.MustRunCluster(t, 1)
defer cluster.Close()
api := cluster.GetNode(0).API
if _, err := api.CreateIndex(context.Background(), "i", pilosa.IndexOptions{}); err != nil {
t.Fatal(err)
} else if _, err = api.CreateField(context.Background(), "i", "bytes", pilosa.OptFieldTypeInt(math.MinInt64, math.MaxInt64)); err != nil {
t.Fatal(err)
} else if _, err = api.CreateField(context.Background(), "i", "duration_time", pilosa.OptFieldTypeInt(math.MinInt64, math.MaxInt64)); err != nil {
t.Fatal(err)
} else if _, err = api.CreateField(context.Background(), "i", "timestamp", pilosa.OptFieldTypeTimestamp(pilosa.DefaultEpoch, pilosa.TimeUnitSeconds)); err != nil {
t.Fatal(err)
}
for _, tt := range []struct {
name string
input string
output string
}{
{
name: "WhereTimestamp",
input: `SELECT * FROM i WHERE timestamp>"2000-01-01T00:00:00Z"`,
output: `Extract(Row(timestamp>"2000-01-01T00:00:00Z"),Rows(bytes),Rows(duration_time),Rows(timestamp))`,
},
{
name: "WhereTimestampWithSpaces",
input: `SELECT * FROM i WHERE timestamp > "2000-01-01T00:00:00Z"`,
output: `Extract(Row(timestamp>"2000-01-01T00:00:00Z"),Rows(bytes),Rows(duration_time),Rows(timestamp))`,
},
} {
t.Run(tt.name, func(t *testing.T) {
query, err := sql.NewMapper().MapSQL(tt.input)
if err != nil {
t.Fatal(err)
}
h := sql.NewSelectHandler(api)
mr, err := h.MapSelect(context.Background(), query.Statement.(*sqlparser.Select), query.Mask)
if err != nil {
t.Fatal(err)
} else if got, want := mr.Query, tt.output; got != want {
t.Fatalf("unexpected pql\npql: %s\nwant: %s", got, want)
}
})
}
}

View file

@ -15,32 +15,32 @@ const timeFormat = "2006-01-02T15:04"
// LT creates a less than query.
func LT(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s<%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s<%s)", fieldName, formatValue(value))
}
// LTE creates a less than or equal query.
func LTE(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s<=%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s<=%s)", fieldName, formatValue(value))
}
// GT creates a greater than query.
func GT(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s>%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s>%s)", fieldName, formatValue(value))
}
// GTE creates a greater than or equal query.
func GTE(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s>=%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s>=%s)", fieldName, formatValue(value))
}
// Equals creates an equals query.
func Equals(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s=%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s=%s)", fieldName, formatValue(value))
}
// NotEquals creates a not equals query.
func NotEquals(fieldName string, value interface{}) string {
return fmt.Sprintf("Row(%s!=%s)", fieldName, intOrFloat(value))
return fmt.Sprintf("Row(%s!=%s)", fieldName, formatValue(value))
}
// NotNull creates a not equal to null query.
@ -94,7 +94,7 @@ func Like(fieldName string, pattern string) string {
// Between creates a between query.
func Between(fieldName string, a interface{}, b interface{}) string {
return fmt.Sprintf("Row(%s >< [%s,%s])", fieldName, intOrFloat(a), intOrFloat(b))
return fmt.Sprintf("Row(%s >< [%s,%s])", fieldName, formatValue(a), formatValue(b))
}
// Distinct creates a Distinct query.
@ -269,8 +269,10 @@ func formatIDKey(idKey interface{}) (string, error) {
}
}
func intOrFloat(value interface{}) string {
func formatValue(value interface{}) string {
switch value.(type) {
case string:
return fmt.Sprintf("%q", value)
case float64, float32:
// In order to test expected values, we set the precision
// to 8. TODO: It's likely we'll need to address this

View file

@ -34,14 +34,14 @@ func (s *SelectHandler) Handle(ctx context.Context, mapped *MappedSQL) (pproto.T
if !ok {
return nil, fmt.Errorf("statement is not type select: %T", mapped.Statement)
}
mr, err := s.mapSelect(ctx, stmt, mapped.Mask)
mr, err := s.MapSelect(ctx, stmt, mapped.Mask)
if err != nil {
return nil, errors.Wrap(err, "mapping select")
}
return s.execMappingResult(ctx, mr, mapped.SQL)
}
func (s *SelectHandler) mapSelect(ctx context.Context, selectStmt *sqlparser.Select, qm QueryMask) (*MappingResult, error) {
func (s *SelectHandler) MapSelect(ctx context.Context, selectStmt *sqlparser.Select, qm QueryMask) (*MappingResult, error) {
// Get the handler for this query mask.
hndlr := s.router.handler(qm)
if hndlr == nil {