diff --git a/server/grpc_test.go b/server/grpc_test.go index 57344bd47..51604142d 100644 --- a/server/grpc_test.go +++ b/server/grpc_test.go @@ -499,9 +499,10 @@ func TestQuerySQL(t *testing.T) { defer tearDownFunc() tests := []struct { - sql string - exp tableResponse - eq func(tableResponse, tableResponse) error + sql string + exp tableResponse + eq func(tableResponse, tableResponse) error + wantErr string }{ { // Extract(Limit(All(), limit=100, offset=0),Rows(age)) @@ -885,6 +886,26 @@ func TestQuerySQL(t *testing.T) { }, eq: equalUnordered, }, + { + //Extract(Intersect(Row(timestamp>"2017-09-02T12:32:00Z"),Row(timestamp<"2019-09-02T12:32:00Z")),Rows(age), Rows(height)) + //Testing the parenthesis around where clause + sql: "select age, height from grouper where (timestamp > '2017-09-02T12:32:00Z' and timestamp < '2019-09-02T12:32:00Z')", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"height", "int64"}, + }, + rows: []row{ + {[]columnResponse{int64(31), int64(110)}}, + }, + }, + eq: equalUnordered, + }, + { + //Testing empty parenthesis around where clause + sql: "select age, height from grouper where ()", + wantErr: "parsing sql", + }, { //Distinct(Row(timestamp>"2019-09-02T12:32:00Z"), index='grouper',field='age') sql: "select distinct age from grouper where timestamp > '2019-09-02T12:32:00Z'", @@ -997,6 +1018,14 @@ func TestQuerySQL(t *testing.T) { for i, test := range tests { t.Run("test-"+strconv.Itoa(i), func(t *testing.T) { resp, err := gh.QuerySQLUnary(ctx, &pb.QuerySQLRequest{Sql: test.sql}) + if test.wantErr != "" { + if err == nil { + t.Errorf("expected error %q, got nil", test.wantErr) + } else if !strings.Contains(err.Error(), test.wantErr) { + t.Errorf("expected error %q, got %q", test.wantErr, err) + } + return + } if err != nil { t.Fatalf("sql: %s, error: %v", test.sql, err) } else { @@ -1020,6 +1049,14 @@ func TestQuerySQL(t *testing.T) { } mock := &mockPilosa_QuerySQLServer{ctx: context.Background()} err := gh.QuerySQL(&pb.QuerySQLRequest{Sql: test.sql}, mock) + if test.wantErr != "" { + if err == nil { + t.Errorf("expected error %q, got nil", test.wantErr) + } else if !strings.Contains(err.Error(), test.wantErr) { + t.Errorf("expected error %q, got %q", test.wantErr, err) + } + return + } if err != nil { t.Fatalf("sql: %s, error: %v", test.sql, err) } else { diff --git a/sql/mask.go b/sql/mask.go index 8161eacb7..e0c11c4d2 100644 --- a/sql/mask.go +++ b/sql/mask.go @@ -380,6 +380,8 @@ func generateWhereMask(e sqlparser.Expr) []wherePart { wp = append(wp, generateWhereMask(expr.Expr)...) case *sqlparser.IsExpr: wp = append(wp, WherePartFieldCondition) + case *sqlparser.ParenExpr: + wp = append(wp, generateWhereMask(expr.Expr)...) } return wp