From f08d3e2834b13ab1b2754c4349b67db62fd2052d Mon Sep 17 00:00:00 2001 From: Samir Patel <48686912+54mir@users.noreply.github.com> Date: Thu, 16 Jun 2022 22:21:24 -0500 Subject: [PATCH] fix parens around where clause bug (#2121) Bug: if a sql query had a where clause within parens, the entire clause would be ignored; and instead of it translating to a pql intersection, it would become an All(). This occured b/c the parser library mapped such an expresstion to a sqlparser.ParenExpr, and we did not have this as a condition in a type switch. So instead of treating a ParenExpr as nothing, we now recurse into it. --- server/grpc_test.go | 43 ++++++++++++++++++++++++++++++++++++++++--- sql/mask.go | 2 ++ 2 files changed, 42 insertions(+), 3 deletions(-) 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