Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 13 additions & 2 deletions parser/ast.go
Original file line number Diff line number Diff line change
Expand Up @@ -5267,8 +5267,14 @@ func (f *WindowFrameParam) Accept(visitor ASTVisitor) error {
}

type SelectQuery struct {
SelectPos Pos
StatementEnd Pos
SelectPos Pos
StatementEnd Pos
// InnerQuery makes this node a parenthesized group wrapping that query:
// (SELECT 1 UNION ALL SELECT 2). When set, every other clause field is
// empty except Settings, Format and the set-operation fields below, which
// bind clauses that follow the closing ')'. SelectPos and StatementEnd
// span the parentheses and any trailing clause.
InnerQuery *SelectQuery `json:",omitempty"`
With *WithClause
Top *TopClause
HasDistinct bool
Expand Down Expand Up @@ -5303,6 +5309,11 @@ func (s *SelectQuery) End() Pos {
func (s *SelectQuery) Accept(visitor ASTVisitor) error {
visitor.Enter(s)
defer visitor.Leave(s)
if s.InnerQuery != nil {
if err := s.InnerQuery.Accept(visitor); err != nil {
return err
}
}
if s.With != nil {
if err := s.With.Accept(visitor); err != nil {
return err
Expand Down
19 changes: 19 additions & 0 deletions parser/format.go
Original file line number Diff line number Diff line change
Expand Up @@ -2286,6 +2286,21 @@ func (s *SelectItem) FormatSQL(formatter *Formatter) {
}

func (s *SelectQuery) FormatSQL(formatter *Formatter) {
if s.InnerQuery != nil {
formatter.WriteByte('(')
formatter.WriteExpr(s.InnerQuery)
formatter.WriteByte(')')
if s.Settings != nil {
formatter.Break()
formatter.WriteExpr(s.Settings)
}
if s.Format != nil {
formatter.Break()
formatter.WriteExpr(s.Format)
}
s.formatSetOperation(formatter)
return
}
if s.With != nil {
formatter.WriteString("WITH")
formatter.Indent()
Expand Down Expand Up @@ -2368,6 +2383,10 @@ func (s *SelectQuery) FormatSQL(formatter *Formatter) {
formatter.Break()
formatter.WriteExpr(s.Format)
}
s.formatSetOperation(formatter)
}

func (s *SelectQuery) formatSetOperation(formatter *Formatter) {
if s.UnionAll != nil {
formatter.Break()
formatter.WriteString("UNION ALL")
Expand Down
85 changes: 70 additions & 15 deletions parser/parser_query.go
Original file line number Diff line number Diff line change
Expand Up @@ -1053,48 +1053,100 @@ func (p *Parser) parseSelectQuery(_ Pos) (*SelectQuery, error) {
return nil, fmt.Errorf("expected SELECT, WITH or (, got %s", p.currentTokenKind())
}

hasParen := p.tryConsumeTokenKind(TokenKindLParen) != nil
selectStmt, err := p.parseSelectStmt(p.Pos())
if err != nil {
var selectStmt *SelectQuery
var err error
if lparen := p.tryConsumeTokenKind(TokenKindLParen); lparen != nil {
inner, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
}

rparenPos := p.Pos()
if err := p.expectTokenKind(TokenKindRParen); err != nil {
return nil, err
}

selectStmt = &SelectQuery{
SelectPos: lparen.Pos,
StatementEnd: rparenPos + 1,
InnerQuery: inner,
}

settings, err := p.tryParseSettingsClause(p.Pos())
if err != nil {
return nil, err
}
if settings != nil {
selectStmt.Settings = settings
selectStmt.StatementEnd = settings.End()
}

format, err := p.tryParseFormat(p.Pos())
if err != nil {
return nil, err
}
if format != nil {
selectStmt.Format = format
selectStmt.StatementEnd = format.End()
}

// ClickHouse allows a set operator after ')' only when no SETTINGS
// or FORMAT was consumed: (SELECT 1) SETTINGS a=1 UNION ALL SELECT 2
// is a syntax error there.
if settings != nil || format != nil {
return selectStmt, nil
}
} else {
selectStmt, err = p.parseSelectStmt(p.Pos())
if err != nil {
return nil, err
}
}

if err := p.parseSetOperation(selectStmt); err != nil {
return nil, err
}

return selectStmt, nil
}

// parseSetOperation binds a trailing UNION|EXCEPT|INTERSECT to selectStmt.
// The right operand consumes the rest of the chain by recursing into
// parseSelectQuery, so at most one operator is bound per call.
func (p *Parser) parseSetOperation(selectStmt *SelectQuery) error {
switch {
case p.tryConsumeKeywords(KeywordUnion):
switch {
case p.tryConsumeKeywords(KeywordAll):
unionAllExpr, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
return err
}
selectStmt.UnionAll = unionAllExpr
case p.tryConsumeKeywords(KeywordDistinct):
unionDistinctExpr, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
return err
}
selectStmt.UnionDistinct = unionDistinctExpr
default:
return nil, fmt.Errorf("expected ALL or DISTINCT, got %s", p.currentTokenKind())
return fmt.Errorf("expected ALL or DISTINCT, got %s", p.currentTokenKind())
}
case p.tryConsumeKeywords(KeywordExcept):
exceptExpr, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
return err
}
selectStmt.Except = exceptExpr
case p.tryConsumeKeywords(KeywordIntersect):
intersectExpr, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
return err
}
selectStmt.Intersect = intersectExpr
}
if hasParen {
if err := p.expectTokenKind(TokenKindRParen); err != nil {
return nil, err
}
}
return selectStmt, nil

return nil
}

func (p *Parser) parseSelectStmt(pos Pos) (*SelectQuery, error) { // nolint: funlen
Expand Down Expand Up @@ -1265,11 +1317,14 @@ func (p *Parser) parseCTEStmt(pos Pos) (*CTEStmt, error) {
if err := p.expectKeyword(KeywordAs); err != nil {
return nil, err
}
if p.matchTokenKind(TokenKindLParen) {
if p.tryConsumeTokenKind(TokenKindLParen) != nil {
selectQuery, err := p.parseSelectQuery(p.Pos())
if err != nil {
return nil, err
}
if err := p.expectTokenKind(TokenKindRParen); err != nil {
return nil, err
}
return &CTEStmt{
CTEPos: pos,
Expr: expr,
Expand Down
2 changes: 1 addition & 1 deletion parser/parser_table.go
Original file line number Diff line number Diff line change
Expand Up @@ -1498,7 +1498,7 @@ func (p *Parser) parseStmt(pos Pos) (Expr, error) {
p.matchKeyword(KeywordTruncate),
p.matchKeyword(KeywordRename):
expr, err = p.parseDDL(pos)
case p.matchKeyword(KeywordSelect), p.matchKeyword(KeywordWith):
case p.matchKeyword(KeywordSelect), p.matchKeyword(KeywordWith), p.matchTokenKind(TokenKindLParen):
expr, err = p.parseSelectQuery(pos)
case p.matchKeyword(KeywordDelete):
expr, err = p.parseDeleteClause(pos)
Expand Down
55 changes: 55 additions & 0 deletions parser/parser_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -237,10 +237,65 @@ func TestParser_InvalidSyntax(t *testing.T) {
"SELECT a GLOBAL",
"SELECT a REGEXP",
"SELECT * FROM t WHERE a AND",
// A parenthesized select must still be closed and UNION still needs
// ALL or DISTINCT
"(SELECT 1",
"(SELECT 1) UNION SELECT 2",
// ClickHouse rejects a set operator once SETTINGS is bound to a
// parenthesized group
"(SELECT 1) SETTINGS max_threads=1 UNION ALL SELECT 2",
}
for _, sql := range invalidSQLs {
parser := NewParser(sql)
_, err := parser.ParseStmts()
require.Error(t, err, "Expected error for SQL: %s", sql)
}
}

func TestParser_ParenthesizedSetOperationOperands(t *testing.T) {
// A parenthesized operand becomes a group node, so the operator after
// ')' binds to the whole group instead of leaking into its chain.
stmts, err := NewParser("(SELECT 1 UNION DISTINCT SELECT 2) UNION ALL SELECT 3").ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)

group, ok := stmts[0].(*SelectQuery)
require.True(t, ok)
require.NotNil(t, group.InnerQuery)
require.NotNil(t, group.InnerQuery.UnionDistinct)
require.NotNil(t, group.UnionAll)
require.Nil(t, group.InnerQuery.UnionDistinct.UnionAll)

stmts, err = NewParser("SELECT a FROM ((SELECT 1 AS a) UNION ALL (SELECT 2 AS a))").ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)

outer, ok := stmts[0].(*SelectQuery)
require.True(t, ok)
joinTable, ok := outer.From.Expr.(*JoinTableExpr)
require.True(t, ok)
subQuery, ok := joinTable.Table.Expr.(*SubQuery)
require.True(t, ok)
require.NotNil(t, subQuery.Select.InnerQuery)
require.NotNil(t, subQuery.Select.UnionAll)
require.NotNil(t, subQuery.Select.UnionAll.InnerQuery)

// Grouping survives the round trip: ClickHouse gives INTERSECT higher
// precedence than UNION, so dropping the parens would change semantics.
sql := "(SELECT 1 UNION ALL SELECT 2) INTERSECT SELECT 2"
stmts, err = NewParser(sql).ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)
require.Equal(t, sql, Format(stmts[0]))

// SETTINGS and FORMAT after ')' bind to the group.
stmts, err = NewParser("(SELECT 1) SETTINGS max_threads=1 FORMAT JSONEachRow").ParseStmts()
require.NoError(t, err)
require.Len(t, stmts, 1)

group, ok = stmts[0].(*SelectQuery)
require.True(t, ok)
require.NotNil(t, group.InnerQuery)
require.NotNil(t, group.Settings)
require.NotNil(t, group.Format)
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,102 @@
-- Origin SQL:
SELECT a FROM ((SELECT 1 AS a) UNION ALL (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) UNION DISTINCT SELECT 2 AS a);
SELECT a FROM ((SELECT 1 AS a) EXCEPT (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) INTERSECT (SELECT 2 AS a));
(SELECT 1 AS a) UNION ALL SELECT 2 AS a;
(SELECT 1 UNION ALL SELECT 2) UNION ALL SELECT 3;
SELECT 1 UNION ALL (SELECT 2) UNION ALL SELECT 3;
SELECT a FROM (((SELECT 1 AS a)));
(SELECT 1 UNION ALL SELECT 2) INTERSECT SELECT 2;
SELECT 1 INTERSECT (SELECT 2 UNION ALL SELECT 1);
(SELECT 1) SETTINGS max_threads=1;
(SELECT 1 UNION ALL SELECT 2) SETTINGS max_threads=1 FORMAT JSONEachRow;


-- Beautify SQL:
SELECT
a
FROM
((SELECT
1 AS a)
UNION ALL
(SELECT
2 AS a));
SELECT
a
FROM
((SELECT
1 AS a)
UNION DISTINCT
SELECT
2 AS a);
SELECT
a
FROM
((SELECT
1 AS a)
EXCEPT
(SELECT
2 AS a));
SELECT
a
FROM
((SELECT
1 AS a)
INTERSECT
(SELECT
2 AS a));
(SELECT
1 AS a)
UNION ALL
SELECT
2 AS a;
(SELECT
1
UNION ALL
SELECT
2)
UNION ALL
SELECT
3;
SELECT
1
UNION ALL
(SELECT
2)
UNION ALL
SELECT
3;
SELECT
a
FROM
(((SELECT
1 AS a)));
(SELECT
1
UNION ALL
SELECT
2)
INTERSECT
SELECT
2;
SELECT
1
INTERSECT
(SELECT
2
UNION ALL
SELECT
1);
(SELECT
1)
SETTINGS
max_threads=1;
(SELECT
1
UNION ALL
SELECT
2)
SETTINGS
max_threads=1
FORMAT JSONEachRow;
28 changes: 28 additions & 0 deletions parser/testdata/query/format/select_with_parenthesized_union.sql
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
-- Origin SQL:
SELECT a FROM ((SELECT 1 AS a) UNION ALL (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) UNION DISTINCT SELECT 2 AS a);
SELECT a FROM ((SELECT 1 AS a) EXCEPT (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) INTERSECT (SELECT 2 AS a));
(SELECT 1 AS a) UNION ALL SELECT 2 AS a;
(SELECT 1 UNION ALL SELECT 2) UNION ALL SELECT 3;
SELECT 1 UNION ALL (SELECT 2) UNION ALL SELECT 3;
SELECT a FROM (((SELECT 1 AS a)));
(SELECT 1 UNION ALL SELECT 2) INTERSECT SELECT 2;
SELECT 1 INTERSECT (SELECT 2 UNION ALL SELECT 1);
(SELECT 1) SETTINGS max_threads=1;
(SELECT 1 UNION ALL SELECT 2) SETTINGS max_threads=1 FORMAT JSONEachRow;


-- Format SQL:
SELECT a FROM ((SELECT 1 AS a) UNION ALL (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) UNION DISTINCT SELECT 2 AS a);
SELECT a FROM ((SELECT 1 AS a) EXCEPT (SELECT 2 AS a));
SELECT a FROM ((SELECT 1 AS a) INTERSECT (SELECT 2 AS a));
(SELECT 1 AS a) UNION ALL SELECT 2 AS a;
(SELECT 1 UNION ALL SELECT 2) UNION ALL SELECT 3;
SELECT 1 UNION ALL (SELECT 2) UNION ALL SELECT 3;
SELECT a FROM (((SELECT 1 AS a)));
(SELECT 1 UNION ALL SELECT 2) INTERSECT SELECT 2;
SELECT 1 INTERSECT (SELECT 2 UNION ALL SELECT 1);
(SELECT 1) SETTINGS max_threads=1;
(SELECT 1 UNION ALL SELECT 2) SETTINGS max_threads=1 FORMAT JSONEachRow;
Loading
Loading