bqsqlparse 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
bqsqlparse/parser.py ADDED
@@ -0,0 +1,1715 @@
1
+ """Recursive-descent parser for the BigQuery (GoogleSQL) dialect.
2
+
3
+ Supported statements:
4
+ SELECT (incl. WITH / RECURSIVE, set operations, ORDER BY, LIMIT/OFFSET)
5
+ CREATE [OR REPLACE] [TEMP] TABLE / VIEW / MATERIALIZED VIEW ... AS SELECT
6
+ INSERT, UPDATE, DELETE, MERGE
7
+
8
+ Supported constructs include: STRUCT & ARRAY literals and types, UNNEST
9
+ (WITH OFFSET), PIVOT / UNPIVOT, TABLESAMPLE, FOR SYSTEM_TIME AS OF, QUALIFY,
10
+ named windows, GROUP BY ROLLUP/CUBE/GROUPING SETS/ALL, SELECT * EXCEPT /
11
+ REPLACE, aggregate modifiers (DISTINCT / IGNORE NULLS / ORDER BY / LIMIT /
12
+ HAVING MAX), CAST FORMAT, EXTRACT, INTERVAL, typed literals, JSON subscripts,
13
+ array subscripts with OFFSET/ORDINAL/SAFE_OFFSET/SAFE_ORDINAL, named function
14
+ arguments (=>), query parameters, LIKE ANY/ALL, IS [NOT] DISTINCT FROM.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ from typing import List, Optional
20
+
21
+ from .errors import ParseError, UnsupportedStatementError
22
+ from .lexer import (
23
+ BYTES, EOF, IDENT, KEYWORD, NUMBER, OP, PARAM, QIDENT, STRING,
24
+ Token, tokenize,
25
+ )
26
+ from .nodes import (
27
+ ArrayExpr, ArraySubquery, AssertStmt, Assignment, Between, BinaryOp,
28
+ BreakContinueStmt, CallStmt, Case, Cast, ColumnDef, ColumnRef,
29
+ CreateFunction, CreateProcedure, CreateTableAsSelect, CTE, DeclareStmt,
30
+ DeleteStmt, DropStmt, ExecuteImmediate, ExistsSubquery, Extract,
31
+ FieldAccess, ForInStmt, FrameBound, FuncCall, GroupBy, HavingModifier,
32
+ IfBranch, IfStmt, InExpr, InsertStmt, IntervalExpr, IsExpr, Join,
33
+ LikeExpr, Literal, LoopStmt, MergeStmt, MergeWhen, NamedArg, NamedWindow,
34
+ Node, OrderItem, Param, PivotAgg, PivotRef, PivotValue, Query, RaiseStmt,
35
+ RepeatStmt, ReplaceItem, ReturnStmt, RoutineParam, ScalarSubquery,
36
+ ScriptBlock, Select, SelectItem, SetOp, SetStmt, Star, StructExpr,
37
+ StructFieldType, StructFieldValue, SubqueryRef, Subscript, TableFuncRef,
38
+ TableRef, TableSample, TransactionStmt, TruncateStmt, TypedLiteral,
39
+ TypeNode, UnaryOp, UnnestRef, UnpivotGroup, UnpivotRef, UpdateStmt,
40
+ WhenClause, WhileStmt, WindowFrame, WindowSpec,
41
+ )
42
+
43
+ _COMPARISON_OPS = ("=", "!=", "<>", "<", "<=", ">", ">=")
44
+ _TYPED_LITERAL_NAMES = {
45
+ "DATE", "DATETIME", "TIME", "TIMESTAMP", "NUMERIC", "BIGNUMERIC", "JSON",
46
+ }
47
+ _CALLABLE_KEYWORDS = {"IF", "LEFT", "RIGHT", "GROUPING", "COLLATE", "RANGE"}
48
+ _NILADIC_FUNCTIONS = {
49
+ "CURRENT_DATE", "CURRENT_DATETIME", "CURRENT_TIME", "CURRENT_TIMESTAMP",
50
+ }
51
+ _CLAUSE_STARTERS = {
52
+ "FROM", "WHERE", "GROUP", "HAVING", "QUALIFY", "WINDOW", "ORDER",
53
+ "LIMIT", "UNION", "INTERSECT", "EXCEPT", "INTO",
54
+ }
55
+ _SUBSCRIPT_MODES = {"OFFSET", "ORDINAL", "SAFE_OFFSET", "SAFE_ORDINAL"}
56
+
57
+
58
+ class Parser:
59
+ def __init__(self, sql: str):
60
+ self.sql = sql
61
+ self.tokens = tokenize(sql)
62
+ self.i = 0
63
+
64
+ # ------------------------------------------------------------------ utils
65
+
66
+ def peek(self, k: int = 0) -> Token:
67
+ j = min(self.i + k, len(self.tokens) - 1)
68
+ return self.tokens[j]
69
+
70
+ def advance(self) -> Token:
71
+ tok = self.tokens[self.i]
72
+ if tok.type != EOF:
73
+ self.i += 1
74
+ return tok
75
+
76
+ def error(self, msg: str, tok: Optional[Token] = None):
77
+ tok = tok or self.peek()
78
+ got = tok.value if tok.type != EOF else "<end of input>"
79
+ raise ParseError(f"{msg} (got {got!r})", tok.line, tok.col)
80
+
81
+ def at_op(self, *ops: str, k: int = 0) -> bool:
82
+ tok = self.peek(k)
83
+ return tok.type == OP and tok.value in ops
84
+
85
+ def accept_op(self, *ops: str) -> Optional[Token]:
86
+ if self.at_op(*ops):
87
+ return self.advance()
88
+ return None
89
+
90
+ def expect_op(self, op: str) -> Token:
91
+ if not self.at_op(op):
92
+ self.error(f"Expected {op!r}")
93
+ return self.advance()
94
+
95
+ def at_kw(self, *words: str, k: int = 0) -> bool:
96
+ tok = self.peek(k)
97
+ return tok.type in (KEYWORD, IDENT) and tok.value.upper() in words
98
+
99
+ def accept_kw(self, *words: str) -> Optional[Token]:
100
+ if self.at_kw(*words):
101
+ return self.advance()
102
+ return None
103
+
104
+ def expect_kw(self, *words: str) -> Token:
105
+ if not self.at_kw(*words):
106
+ self.error(f"Expected {' or '.join(words)}")
107
+ return self.advance()
108
+
109
+ def _is_ident(self, k: int = 0) -> bool:
110
+ return self.peek(k).type in (IDENT, QIDENT)
111
+
112
+ def parse_ident(self) -> str:
113
+ tok = self.peek()
114
+ if tok.type in (IDENT, QIDENT):
115
+ return self.advance().value
116
+ self.error("Expected identifier")
117
+
118
+ def _parse_name_after_dot(self) -> str:
119
+ tok = self.peek()
120
+ if tok.type in (IDENT, QIDENT, KEYWORD):
121
+ return self.advance().value
122
+ self.error("Expected name after '.'")
123
+
124
+ def _expect_gt(self):
125
+ """Consume one '>' — splits a '>>' token for nested generic types."""
126
+ tok = self.peek()
127
+ if tok.type == OP and tok.value == ">":
128
+ self.advance()
129
+ elif tok.type == OP and tok.value == ">>":
130
+ tok.value = ">" # consume one of the two
131
+ else:
132
+ self.error("Expected '>'")
133
+
134
+ def _starts_query(self, k: int = 0) -> bool:
135
+ while self.at_op("(", k=k):
136
+ k += 1
137
+ return self.at_kw("SELECT", "WITH", k=k)
138
+
139
+ # ------------------------------------------------------------- statements
140
+
141
+ def parse_statements(self) -> List[Node]:
142
+ stmts: List[Node] = []
143
+ while True:
144
+ while self.accept_op(";"):
145
+ pass
146
+ if self.peek().type == EOF:
147
+ break
148
+ stmts.append(self.parse_statement())
149
+ return stmts
150
+
151
+ def parse_statement(self) -> Node:
152
+ if self._starts_query():
153
+ return self.parse_query()
154
+ if self.at_kw("CREATE"):
155
+ return self.parse_create()
156
+ if self.at_kw("INSERT"):
157
+ return self.parse_insert()
158
+ if self.at_kw("UPDATE"):
159
+ return self.parse_update()
160
+ if self.at_kw("DELETE"):
161
+ return self.parse_delete()
162
+ if self.at_kw("MERGE"):
163
+ return self.parse_merge()
164
+ # ------------------------- scripting (procedural language) ----------
165
+ if self._is_ident() and self.at_op(":", k=1): # label: LOOP/WHILE/...
166
+ label = self.advance().value
167
+ self.advance()
168
+ return self._parse_labeled(label)
169
+ if self.at_kw("DECLARE"):
170
+ return self.parse_declare()
171
+ if self.at_kw("SET"):
172
+ return self.parse_set_stmt()
173
+ if self.at_kw("BEGIN"):
174
+ return self.parse_begin()
175
+ if self.at_kw("IF"):
176
+ return self.parse_if()
177
+ if self.at_kw("LOOP"):
178
+ return self.parse_loop()
179
+ if self.at_kw("WHILE"):
180
+ return self.parse_while()
181
+ if self.at_kw("REPEAT"):
182
+ return self.parse_repeat()
183
+ if self.at_kw("FOR"):
184
+ return self.parse_for_in()
185
+ if self.at_kw("BREAK", "LEAVE", "CONTINUE", "ITERATE"):
186
+ kind = self.advance().value.upper()
187
+ label = self.advance().value if self._is_ident() else None
188
+ return BreakContinueStmt(kind, label)
189
+ if self.at_kw("CALL"):
190
+ self.advance()
191
+ name = self._split_path(self._parse_dotted_path())
192
+ args = self._paren_expr_list() if self.at_op("(") else []
193
+ return CallStmt(name, args)
194
+ if self.at_kw("RETURN"):
195
+ self.advance()
196
+ return ReturnStmt()
197
+ if self.at_kw("RAISE"):
198
+ self.advance()
199
+ message = None
200
+ if self.accept_kw("USING"):
201
+ self.expect_kw("MESSAGE")
202
+ self.expect_op("=")
203
+ message = self.parse_expr()
204
+ return RaiseStmt(message)
205
+ if self.at_kw("EXECUTE") and self.at_kw("IMMEDIATE", k=1):
206
+ return self.parse_execute_immediate()
207
+ if self.at_kw("ASSERT"):
208
+ self.advance()
209
+ cond = self.parse_expr()
210
+ message = None
211
+ if self.accept_kw("AS"):
212
+ message = self.parse_expr()
213
+ return AssertStmt(cond, message)
214
+ if self.at_kw("COMMIT", "ROLLBACK"):
215
+ kind = self.advance().value.upper()
216
+ self.accept_kw("TRANSACTION")
217
+ return TransactionStmt(kind)
218
+ if self.at_kw("TRUNCATE"):
219
+ self.advance()
220
+ self.expect_kw("TABLE")
221
+ return TruncateStmt(self._split_path(self._parse_dotted_path()))
222
+ if self.at_kw("DROP"):
223
+ return self.parse_drop()
224
+ tok = self.peek()
225
+ raise UnsupportedStatementError(
226
+ f"Unsupported statement starting with {tok.value!r}", tok.line, tok.col
227
+ )
228
+
229
+ # ------------------------------------------------------------------ query
230
+
231
+ def parse_query(self) -> Query:
232
+ ctes: List[CTE] = []
233
+ recursive = False
234
+ if self.accept_kw("WITH"):
235
+ recursive = bool(self.accept_kw("RECURSIVE"))
236
+ while True:
237
+ name = self.parse_ident()
238
+ self.expect_kw("AS")
239
+ self.expect_op("(")
240
+ q = self.parse_query()
241
+ self.expect_op(")")
242
+ ctes.append(CTE(name, q))
243
+ if not self.accept_op(","):
244
+ break
245
+ body = self.parse_set_expr()
246
+ order_by = self.parse_order_by_opt()
247
+ limit = offset = None
248
+ if self.accept_kw("LIMIT"):
249
+ limit = self.parse_expr()
250
+ if self.accept_kw("OFFSET"):
251
+ offset = self.parse_expr()
252
+ return Query(body=body, ctes=ctes, recursive=recursive,
253
+ order_by=order_by, limit=limit, offset=offset)
254
+
255
+ def parse_order_by_opt(self) -> List[OrderItem]:
256
+ if not self.at_kw("ORDER"):
257
+ return []
258
+ self.advance()
259
+ self.expect_kw("BY")
260
+ return self.parse_order_items()
261
+
262
+ def parse_order_items(self) -> List[OrderItem]:
263
+ items = []
264
+ while True:
265
+ expr = self.parse_expr()
266
+ desc = None
267
+ if self.accept_kw("ASC"):
268
+ desc = False
269
+ elif self.accept_kw("DESC"):
270
+ desc = True
271
+ nulls = None
272
+ if self.accept_kw("NULLS"):
273
+ nulls = self.expect_kw("FIRST", "LAST").value.upper()
274
+ items.append(OrderItem(expr, desc, nulls))
275
+ if not self.accept_op(","):
276
+ break
277
+ return items
278
+
279
+ def parse_set_expr(self) -> Node:
280
+ left = self.parse_query_primary()
281
+ while True:
282
+ if self.at_kw("UNION") or self.at_kw("INTERSECT"):
283
+ op = self.advance().value.upper()
284
+ elif self.at_kw("EXCEPT") and self.at_kw("ALL", "DISTINCT", k=1):
285
+ op = self.advance().value.upper()
286
+ else:
287
+ break
288
+ all_ = self.expect_kw("ALL", "DISTINCT").value.upper() == "ALL"
289
+ right = self.parse_query_primary()
290
+ left = SetOp(op, all_, left, right)
291
+ return left
292
+
293
+ def parse_query_primary(self) -> Node:
294
+ if self.at_kw("SELECT"):
295
+ return self.parse_select()
296
+ if self.at_op("("):
297
+ self.advance()
298
+ q = self.parse_query()
299
+ self.expect_op(")")
300
+ return q
301
+ self.error("Expected SELECT or '('")
302
+
303
+ def parse_select(self) -> Select:
304
+ self.expect_kw("SELECT")
305
+ distinct = False
306
+ if self.accept_kw("DISTINCT"):
307
+ distinct = True
308
+ else:
309
+ self.accept_kw("ALL")
310
+ as_mode = None
311
+ if self.accept_kw("AS"):
312
+ as_mode = self.expect_kw("STRUCT", "VALUE").value.upper()
313
+
314
+ items: List[Node] = []
315
+ while True:
316
+ items.append(self.parse_select_item())
317
+ if not self.accept_op(","):
318
+ break
319
+ # tolerate BigQuery trailing commas
320
+ if self._at_clause_boundary():
321
+ break
322
+
323
+ from_ = where = group_by = having = qualify = None
324
+ windows: List[NamedWindow] = []
325
+ if self.accept_kw("FROM"):
326
+ from_ = self.parse_from()
327
+ if self.accept_kw("WHERE"):
328
+ where = self.parse_expr()
329
+ if self.at_kw("GROUP"):
330
+ self.advance()
331
+ self.expect_kw("BY")
332
+ group_by = self.parse_group_by()
333
+ if self.accept_kw("HAVING"):
334
+ having = self.parse_expr()
335
+ if self.accept_kw("QUALIFY"):
336
+ qualify = self.parse_expr()
337
+ if self.accept_kw("WINDOW"):
338
+ while True:
339
+ wname = self.parse_ident()
340
+ self.expect_kw("AS")
341
+ self.expect_op("(")
342
+ spec = self.parse_window_spec_body()
343
+ self.expect_op(")")
344
+ windows.append(NamedWindow(wname, spec))
345
+ if not self.accept_op(","):
346
+ break
347
+ return Select(items=items, distinct=distinct, as_mode=as_mode,
348
+ from_=from_, where=where, group_by=group_by,
349
+ having=having, qualify=qualify, windows=windows)
350
+
351
+ def _at_clause_boundary(self) -> bool:
352
+ tok = self.peek()
353
+ if tok.type == EOF or (tok.type == OP and tok.value in (")", ";")):
354
+ return True
355
+ return tok.type == KEYWORD and tok.value in _CLAUSE_STARTERS
356
+
357
+ def parse_select_item(self) -> Node:
358
+ if self.at_op("*"):
359
+ self.advance()
360
+ return self._parse_star_modifiers([])
361
+ expr = self.parse_expr()
362
+ if self.at_op(".") and self.at_op("*", k=1):
363
+ if not isinstance(expr, ColumnRef):
364
+ self.error("Expected path before '.*'")
365
+ self.advance()
366
+ self.advance()
367
+ return self._parse_star_modifiers(expr.path)
368
+ alias = None
369
+ if self.accept_kw("AS"):
370
+ alias = self.parse_ident()
371
+ elif self._is_ident():
372
+ alias = self.advance().value
373
+ return SelectItem(expr, alias)
374
+
375
+ def _parse_star_modifiers(self, prefix: List[str]) -> Star:
376
+ except_: List[str] = []
377
+ replace: List[ReplaceItem] = []
378
+ if self.at_kw("EXCEPT") and self.at_op("(", k=1):
379
+ self.advance()
380
+ self.advance()
381
+ while True:
382
+ except_.append(self.parse_ident())
383
+ if not self.accept_op(","):
384
+ break
385
+ self.expect_op(")")
386
+ if self.at_kw("REPLACE") and self.at_op("(", k=1):
387
+ self.advance()
388
+ self.advance()
389
+ while True:
390
+ e = self.parse_expr()
391
+ self.expect_kw("AS")
392
+ name = self.parse_ident()
393
+ replace.append(ReplaceItem(e, name))
394
+ if not self.accept_op(","):
395
+ break
396
+ self.expect_op(")")
397
+ return Star(prefix=prefix, except_=except_, replace=replace)
398
+
399
+ def parse_group_by(self) -> GroupBy:
400
+ if self.accept_kw("ALL"):
401
+ return GroupBy("all")
402
+ if self.at_kw("ROLLUP") and self.at_op("(", k=1):
403
+ self.advance()
404
+ return GroupBy("rollup", exprs=self._paren_expr_list())
405
+ if self.at_kw("CUBE") and self.at_op("(", k=1):
406
+ self.advance()
407
+ return GroupBy("cube", exprs=self._paren_expr_list())
408
+ if self.at_kw("GROUPING") and self.at_kw("SETS", k=1):
409
+ self.advance()
410
+ self.advance()
411
+ self.expect_op("(")
412
+ sets: List[List[Node]] = []
413
+ while True:
414
+ if self.at_kw("ROLLUP", "CUBE") and self.at_op("(", k=1):
415
+ self.advance()
416
+ sets.append(self._paren_expr_list())
417
+ elif self.accept_op("("):
418
+ inner: List[Node] = []
419
+ if not self.at_op(")"):
420
+ while True:
421
+ inner.append(self.parse_expr())
422
+ if not self.accept_op(","):
423
+ break
424
+ self.expect_op(")")
425
+ sets.append(inner)
426
+ else:
427
+ sets.append([self.parse_expr()])
428
+ if not self.accept_op(","):
429
+ break
430
+ self.expect_op(")")
431
+ return GroupBy("sets", sets=sets)
432
+ exprs = []
433
+ while True:
434
+ exprs.append(self.parse_expr())
435
+ if not self.accept_op(","):
436
+ break
437
+ return GroupBy("exprs", exprs=exprs)
438
+
439
+ def _paren_expr_list(self) -> List[Node]:
440
+ self.expect_op("(")
441
+ exprs = []
442
+ if not self.at_op(")"):
443
+ while True:
444
+ exprs.append(self.parse_expr())
445
+ if not self.accept_op(","):
446
+ break
447
+ self.expect_op(")")
448
+ return exprs
449
+
450
+ # ------------------------------------------------------------------- FROM
451
+
452
+ def parse_from(self) -> Node:
453
+ left = self.parse_join_expr()
454
+ while self.accept_op(","):
455
+ right = self.parse_join_expr()
456
+ left = Join("CROSS", left, right)
457
+ return left
458
+
459
+ def parse_join_expr(self) -> Node:
460
+ left = self.parse_table_primary()
461
+ while True:
462
+ kind = None
463
+ if self.at_kw("CROSS") and self.at_kw("JOIN", k=1):
464
+ self.advance()
465
+ self.advance()
466
+ kind = "CROSS"
467
+ elif self.at_kw("INNER"):
468
+ self.advance()
469
+ self.expect_kw("JOIN")
470
+ kind = "INNER"
471
+ elif self.at_kw("LEFT", "RIGHT", "FULL") and self.at_kw("OUTER", "JOIN", k=1):
472
+ kind = self.advance().value.upper()
473
+ self.accept_kw("OUTER")
474
+ self.expect_kw("JOIN")
475
+ elif self.at_kw("JOIN"):
476
+ self.advance()
477
+ kind = "INNER"
478
+ else:
479
+ break
480
+ right = self.parse_table_primary()
481
+ on = None
482
+ using: List[str] = []
483
+ if kind != "CROSS":
484
+ if self.accept_kw("ON"):
485
+ on = self.parse_expr()
486
+ elif self.accept_kw("USING"):
487
+ self.expect_op("(")
488
+ while True:
489
+ using.append(self.parse_ident())
490
+ if not self.accept_op(","):
491
+ break
492
+ self.expect_op(")")
493
+ left = Join(kind, left, right, on=on, using=using)
494
+ return left
495
+
496
+ def parse_table_primary(self) -> Node:
497
+ node = self._parse_table_primary_base()
498
+ while True:
499
+ if self.at_kw("PIVOT") and self.at_op("(", k=1):
500
+ node = self._parse_pivot(node)
501
+ elif self.at_kw("UNPIVOT") and (
502
+ self.at_op("(", k=1) or self.at_kw("INCLUDE", "EXCLUDE", k=1)
503
+ ):
504
+ node = self._parse_unpivot(node)
505
+ elif self.at_kw("TABLESAMPLE"):
506
+ self.advance()
507
+ method = self.expect_kw("SYSTEM", "BERNOULLI", "RESERVOIR").value.upper()
508
+ self.expect_op("(")
509
+ value = self.parse_expr()
510
+ unit = "PERCENT"
511
+ if self.at_kw("PERCENT", "ROWS"):
512
+ unit = self.advance().value.upper()
513
+ self.expect_op(")")
514
+ if isinstance(node, TableRef):
515
+ node.sample = TableSample(method, value, unit)
516
+ else:
517
+ break
518
+ return node
519
+
520
+ def _parse_table_alias(self) -> Optional[str]:
521
+ if self.accept_kw("AS"):
522
+ return self.parse_ident()
523
+ if self._is_ident():
524
+ up = self.peek().value.upper()
525
+ if up in ("PIVOT", "UNPIVOT") and self.at_op("(", k=1):
526
+ return None
527
+ if up in ("TABLESAMPLE",):
528
+ return None
529
+ return self.advance().value
530
+ return None
531
+
532
+ def _parse_table_primary_base(self) -> Node:
533
+ if self.at_op("("):
534
+ if self._starts_query():
535
+ self.advance()
536
+ q = self.parse_query()
537
+ self.expect_op(")")
538
+ alias = self._parse_table_alias()
539
+ return SubqueryRef(q, alias)
540
+ self.advance()
541
+ inner = self.parse_from()
542
+ self.expect_op(")")
543
+ return inner
544
+ if self.at_kw("UNNEST"):
545
+ self.advance()
546
+ self.expect_op("(")
547
+ expr = self.parse_expr()
548
+ self.expect_op(")")
549
+ alias = self._parse_table_alias()
550
+ with_offset = False
551
+ offset_alias = None
552
+ if self.at_kw("WITH") and self.at_kw("OFFSET", k=1):
553
+ self.advance()
554
+ self.advance()
555
+ with_offset = True
556
+ if self.accept_kw("AS"):
557
+ offset_alias = self.parse_ident()
558
+ elif self._is_ident():
559
+ offset_alias = self.advance().value
560
+ return UnnestRef(expr, alias, with_offset, offset_alias)
561
+ # table path (possibly a table-valued function)
562
+ path = [self.parse_ident()]
563
+ while self.at_op(".") and not self.at_op("*", k=1):
564
+ self.advance()
565
+ path.append(self._parse_name_after_dot())
566
+ if self.at_op("("):
567
+ args = self._paren_expr_list()
568
+ alias = self._parse_table_alias()
569
+ return TableFuncRef(path, args, alias)
570
+ # a quoted `project.dataset.table` arrives as a single token — split it
571
+ parts: List[str] = []
572
+ for p in path:
573
+ parts.extend(p.split(".") if "." in p else [p])
574
+ system_time = None
575
+ if self.at_kw("FOR") and self.at_kw("SYSTEM_TIME", k=1):
576
+ self.advance()
577
+ self.advance()
578
+ self.expect_kw("AS")
579
+ self.expect_kw("OF")
580
+ system_time = self.parse_expr()
581
+ alias = self._parse_table_alias()
582
+ return TableRef(parts, alias, system_time)
583
+
584
+ def _parse_pivot(self, input_: Node) -> PivotRef:
585
+ self.advance() # PIVOT
586
+ self.expect_op("(")
587
+ aggs: List[PivotAgg] = []
588
+ while True:
589
+ expr = self.parse_expr()
590
+ if not isinstance(expr, FuncCall):
591
+ self.error("PIVOT expects aggregate function calls")
592
+ agg_alias = None
593
+ if self.accept_kw("AS"):
594
+ agg_alias = self.parse_ident()
595
+ elif self._is_ident() and not self.at_kw("FOR"):
596
+ agg_alias = self.advance().value
597
+ aggs.append(PivotAgg(expr, agg_alias))
598
+ if not self.accept_op(","):
599
+ break
600
+ self.expect_kw("FOR")
601
+ for_col = ColumnRef(self._parse_dotted_path())
602
+ self.expect_kw("IN")
603
+ self.expect_op("(")
604
+ in_values: List[PivotValue] = []
605
+ while True:
606
+ v = self.parse_expr()
607
+ v_alias = None
608
+ if self.accept_kw("AS"):
609
+ v_alias = self.parse_ident()
610
+ elif self._is_ident():
611
+ v_alias = self.advance().value
612
+ in_values.append(PivotValue(v, v_alias))
613
+ if not self.accept_op(","):
614
+ break
615
+ self.expect_op(")")
616
+ self.expect_op(")")
617
+ alias = self._parse_table_alias()
618
+ return PivotRef(input_, aggs, for_col, in_values, alias)
619
+
620
+ def _parse_unpivot(self, input_: Node) -> UnpivotRef:
621
+ self.advance() # UNPIVOT
622
+ include_nulls = None
623
+ if self.accept_kw("INCLUDE"):
624
+ self.expect_kw("NULLS")
625
+ include_nulls = True
626
+ elif self.accept_kw("EXCLUDE"):
627
+ self.expect_kw("NULLS")
628
+ include_nulls = False
629
+ self.expect_op("(")
630
+ value_columns: List[str] = []
631
+ multi = False
632
+ if self.accept_op("("):
633
+ multi = True
634
+ while True:
635
+ value_columns.append(self.parse_ident())
636
+ if not self.accept_op(","):
637
+ break
638
+ self.expect_op(")")
639
+ else:
640
+ value_columns.append(self.parse_ident())
641
+ self.expect_kw("FOR")
642
+ name_column = self.parse_ident()
643
+ self.expect_kw("IN")
644
+ self.expect_op("(")
645
+ groups: List[UnpivotGroup] = []
646
+ while True:
647
+ cols: List[List[str]] = []
648
+ if multi:
649
+ self.expect_op("(")
650
+ while True:
651
+ cols.append(self._parse_dotted_path())
652
+ if not self.accept_op(","):
653
+ break
654
+ self.expect_op(")")
655
+ else:
656
+ cols.append(self._parse_dotted_path())
657
+ label = None
658
+ if self.accept_kw("AS"):
659
+ label = self.parse_expr()
660
+ elif self.peek().type in (STRING, NUMBER):
661
+ label = self.parse_expr()
662
+ groups.append(UnpivotGroup(cols, label))
663
+ if not self.accept_op(","):
664
+ break
665
+ self.expect_op(")")
666
+ self.expect_op(")")
667
+ alias = self._parse_table_alias()
668
+ return UnpivotRef(input_, value_columns, name_column, groups,
669
+ include_nulls, alias)
670
+
671
+ def _parse_dotted_path(self) -> List[str]:
672
+ path = [self.parse_ident()]
673
+ while self.at_op(".") and not self.at_op("*", k=1):
674
+ self.advance()
675
+ path.append(self._parse_name_after_dot())
676
+ return path
677
+
678
+ # ------------------------------------------------------------ expressions
679
+
680
+ def parse_expr(self) -> Node:
681
+ return self.parse_or()
682
+
683
+ def parse_or(self) -> Node:
684
+ left = self.parse_and()
685
+ while self.at_kw("OR") and self.peek().type == KEYWORD:
686
+ self.advance()
687
+ left = BinaryOp("OR", left, self.parse_and())
688
+ return left
689
+
690
+ def parse_and(self) -> Node:
691
+ left = self.parse_not()
692
+ while self.at_kw("AND") and self.peek().type == KEYWORD:
693
+ self.advance()
694
+ left = BinaryOp("AND", left, self.parse_not())
695
+ return left
696
+
697
+ def parse_not(self) -> Node:
698
+ if self.peek().type == KEYWORD and self.peek().value == "NOT" \
699
+ and not self.at_kw("LIKE", "BETWEEN", "IN", k=1):
700
+ self.advance()
701
+ return UnaryOp("NOT", self.parse_not())
702
+ return self.parse_comparison()
703
+
704
+ def parse_comparison(self) -> Node:
705
+ left = self.parse_bitor()
706
+ while True:
707
+ negated = False
708
+ if self.peek().type == KEYWORD and self.peek().value == "NOT" \
709
+ and self.at_kw("LIKE", "BETWEEN", "IN", k=1):
710
+ self.advance()
711
+ negated = True
712
+ tok = self.peek()
713
+ if tok.type == OP and tok.value in _COMPARISON_OPS:
714
+ self.advance()
715
+ left = BinaryOp(tok.value, left, self.parse_bitor())
716
+ elif self.at_kw("LIKE"):
717
+ self.advance()
718
+ quantifier = None
719
+ if self.at_kw("ANY", "ALL", "SOME"):
720
+ quantifier = self.advance().value.upper()
721
+ patterns = self._paren_expr_list()
722
+ else:
723
+ patterns = [self.parse_bitor()]
724
+ left = LikeExpr(left, patterns, quantifier, negated)
725
+ elif self.at_kw("BETWEEN"):
726
+ self.advance()
727
+ low = self.parse_bitor()
728
+ self.expect_kw("AND")
729
+ high = self.parse_bitor()
730
+ left = Between(left, low, high, negated)
731
+ elif self.at_kw("IN"):
732
+ self.advance()
733
+ if self.at_kw("UNNEST"):
734
+ self.advance()
735
+ self.expect_op("(")
736
+ arr = self.parse_expr()
737
+ self.expect_op(")")
738
+ left = InExpr(left, unnest=arr, negated=negated)
739
+ else:
740
+ self.expect_op("(")
741
+ if self._starts_query():
742
+ q = self.parse_query()
743
+ self.expect_op(")")
744
+ left = InExpr(left, query=q, negated=negated)
745
+ else:
746
+ values = [self.parse_expr()]
747
+ while self.accept_op(","):
748
+ values.append(self.parse_expr())
749
+ self.expect_op(")")
750
+ left = InExpr(left, values=values, negated=negated)
751
+ elif self.at_kw("IS"):
752
+ self.advance()
753
+ neg = bool(self.accept_kw("NOT"))
754
+ if self.accept_kw("DISTINCT"):
755
+ self.expect_kw("FROM")
756
+ op = "IS NOT DISTINCT FROM" if neg else "IS DISTINCT FROM"
757
+ left = BinaryOp(op, left, self.parse_bitor())
758
+ else:
759
+ val = self.expect_kw("NULL", "TRUE", "FALSE", "UNKNOWN")
760
+ left = IsExpr(left, val.value.upper(), neg)
761
+ else:
762
+ break
763
+ return left
764
+
765
+ def parse_bitor(self) -> Node:
766
+ left = self.parse_bitxor()
767
+ while self.at_op("|"):
768
+ self.advance()
769
+ left = BinaryOp("|", left, self.parse_bitxor())
770
+ return left
771
+
772
+ def parse_bitxor(self) -> Node:
773
+ left = self.parse_bitand()
774
+ while self.at_op("^"):
775
+ self.advance()
776
+ left = BinaryOp("^", left, self.parse_bitand())
777
+ return left
778
+
779
+ def parse_bitand(self) -> Node:
780
+ left = self.parse_shift()
781
+ while self.at_op("&"):
782
+ self.advance()
783
+ left = BinaryOp("&", left, self.parse_shift())
784
+ return left
785
+
786
+ def parse_shift(self) -> Node:
787
+ left = self.parse_additive()
788
+ while self.at_op("<<", ">>"):
789
+ op = self.advance().value
790
+ left = BinaryOp(op, left, self.parse_additive())
791
+ return left
792
+
793
+ def parse_additive(self) -> Node:
794
+ left = self.parse_multiplicative()
795
+ while self.at_op("+", "-"):
796
+ op = self.advance().value
797
+ left = BinaryOp(op, left, self.parse_multiplicative())
798
+ return left
799
+
800
+ def parse_multiplicative(self) -> Node:
801
+ left = self.parse_unary()
802
+ while self.at_op("*", "/", "||"):
803
+ op = self.advance().value
804
+ left = BinaryOp(op, left, self.parse_unary())
805
+ return left
806
+
807
+ def parse_unary(self) -> Node:
808
+ if self.at_op("-", "+", "~"):
809
+ op = self.advance().value
810
+ return UnaryOp(op, self.parse_unary())
811
+ return self.parse_postfix()
812
+
813
+ def parse_postfix(self) -> Node:
814
+ expr = self.parse_primary()
815
+ while True:
816
+ if self.at_op(".") and not self.at_op("*", k=1):
817
+ self.advance()
818
+ name = self._parse_name_after_dot()
819
+ if isinstance(expr, ColumnRef):
820
+ expr = ColumnRef(expr.path + [name])
821
+ else:
822
+ expr = FieldAccess(expr, name)
823
+ elif self.at_op("["):
824
+ self.advance()
825
+ mode = None
826
+ if self.at_kw(*_SUBSCRIPT_MODES) and self.at_op("(", k=1):
827
+ mode = self.advance().value.upper()
828
+ self.expect_op("(")
829
+ index = self.parse_expr()
830
+ self.expect_op(")")
831
+ else:
832
+ index = self.parse_expr()
833
+ self.expect_op("]")
834
+ expr = Subscript(expr, index, mode)
835
+ else:
836
+ break
837
+ return expr
838
+
839
+ def parse_primary(self) -> Node:
840
+ tok = self.peek()
841
+
842
+ if tok.type == NUMBER:
843
+ self.advance()
844
+ return Literal(tok.value, "number")
845
+ if tok.type == STRING:
846
+ self.advance()
847
+ return Literal(tok.value, "string")
848
+ if tok.type == BYTES:
849
+ self.advance()
850
+ return Literal(tok.value, "bytes")
851
+ if tok.type == PARAM:
852
+ self.advance()
853
+ return Param(tok.value)
854
+
855
+ if tok.type == KEYWORD:
856
+ v = tok.value
857
+ if v in ("TRUE", "FALSE"):
858
+ self.advance()
859
+ return Literal(v == "TRUE", "bool")
860
+ if v == "NULL":
861
+ self.advance()
862
+ return Literal(None, "null")
863
+ if v == "DEFAULT":
864
+ self.advance()
865
+ return Literal(None, "default")
866
+ if v == "CASE":
867
+ return self.parse_case()
868
+ if v == "CAST":
869
+ return self.parse_cast(safe=False)
870
+ if v == "EXTRACT":
871
+ return self.parse_extract()
872
+ if v == "EXISTS" and self.at_op("(", k=1):
873
+ self.advance()
874
+ self.expect_op("(")
875
+ q = self.parse_query()
876
+ self.expect_op(")")
877
+ return ExistsSubquery(q)
878
+ if v == "ARRAY":
879
+ return self.parse_array()
880
+ if v == "STRUCT":
881
+ return self.parse_struct()
882
+ if v == "INTERVAL":
883
+ return self.parse_interval()
884
+ if v == "NOT":
885
+ self.advance()
886
+ return UnaryOp("NOT", self.parse_not())
887
+ if v in _CALLABLE_KEYWORDS and self.at_op("(", k=1):
888
+ self.advance()
889
+ return self.parse_func_call([v])
890
+
891
+ if self.at_op("["): # bare array literal: [1, 2, 3]
892
+ self.advance()
893
+ elements: List[Node] = []
894
+ if not self.at_op("]"):
895
+ while True:
896
+ elements.append(self.parse_expr())
897
+ if not self.accept_op(","):
898
+ break
899
+ self.expect_op("]")
900
+ return ArrayExpr(elements)
901
+
902
+ if self.at_op("("):
903
+ if self._starts_query():
904
+ self.advance()
905
+ q = self.parse_query()
906
+ self.expect_op(")")
907
+ return ScalarSubquery(q)
908
+ self.advance()
909
+ expr = self.parse_expr()
910
+ if self.at_op(","):
911
+ fields = [StructFieldValue(expr)]
912
+ while self.accept_op(","):
913
+ fields.append(StructFieldValue(self.parse_expr()))
914
+ self.expect_op(")")
915
+ return StructExpr(fields)
916
+ self.expect_op(")")
917
+ return expr
918
+
919
+ # typed literals: DATE '...', JSON '...', etc.
920
+ if tok.type == IDENT and tok.value.upper() in _TYPED_LITERAL_NAMES \
921
+ and self.peek(1).type == STRING:
922
+ self.advance()
923
+ lit = self.advance()
924
+ return TypedLiteral(tok.value.upper(), lit.value)
925
+
926
+ if tok.type in (IDENT, QIDENT):
927
+ path = [self.advance().value]
928
+ while self.at_op(".") and not self.at_op("*", k=1) \
929
+ and self.peek(1).type in (IDENT, QIDENT, KEYWORD):
930
+ self.advance()
931
+ path.append(self._parse_name_after_dot())
932
+ if self.at_op("("):
933
+ if path[-1].upper() == "SAFE_CAST" and len(path) == 1:
934
+ return self.parse_cast(safe=True, consumed_name=True)
935
+ return self.parse_func_call(path)
936
+ if len(path) == 1 and path[0].upper() in _NILADIC_FUNCTIONS:
937
+ return FuncCall([path[0].upper()])
938
+ return ColumnRef(path)
939
+
940
+ self.error("Unexpected token in expression")
941
+
942
+ def parse_case(self) -> Case:
943
+ self.expect_kw("CASE")
944
+ operand = None
945
+ if not self.at_kw("WHEN"):
946
+ operand = self.parse_expr()
947
+ whens: List[WhenClause] = []
948
+ while self.accept_kw("WHEN"):
949
+ cond = self.parse_expr()
950
+ self.expect_kw("THEN")
951
+ whens.append(WhenClause(cond, self.parse_expr()))
952
+ else_ = None
953
+ if self.accept_kw("ELSE"):
954
+ else_ = self.parse_expr()
955
+ self.expect_kw("END")
956
+ return Case(operand, whens, else_)
957
+
958
+ def parse_cast(self, safe: bool, consumed_name: bool = False) -> Cast:
959
+ if not consumed_name:
960
+ self.advance() # CAST
961
+ self.expect_op("(")
962
+ expr = self.parse_expr()
963
+ self.expect_kw("AS")
964
+ to_type = self.parse_type()
965
+ fmt = None
966
+ if self.accept_kw("FORMAT"):
967
+ fmt = self.parse_expr()
968
+ self.expect_op(")")
969
+ return Cast(expr, to_type, safe, fmt)
970
+
971
+ def parse_extract(self) -> Extract:
972
+ self.expect_kw("EXTRACT")
973
+ self.expect_op("(")
974
+ part_tok = self.advance()
975
+ part = part_tok.value.upper()
976
+ if self.at_op("("): # WEEK(MONDAY)
977
+ self.advance()
978
+ day = self.advance().value.upper()
979
+ self.expect_op(")")
980
+ part = f"{part}({day})"
981
+ self.expect_kw("FROM")
982
+ expr = self.parse_expr()
983
+ tz = None
984
+ if self.accept_kw("AT"):
985
+ self.expect_kw("TIME")
986
+ self.expect_kw("ZONE")
987
+ tz = self.parse_expr()
988
+ self.expect_op(")")
989
+ return Extract(part, expr, tz)
990
+
991
+ def parse_array(self) -> Node:
992
+ self.expect_kw("ARRAY")
993
+ elem_type = None
994
+ if self.at_op("<"):
995
+ self.advance()
996
+ elem_type = self.parse_type()
997
+ self._expect_gt()
998
+ if self.at_op("("):
999
+ self.advance()
1000
+ q = self.parse_query()
1001
+ self.expect_op(")")
1002
+ return ArraySubquery(q)
1003
+ self.expect_op("[")
1004
+ elements: List[Node] = []
1005
+ if not self.at_op("]"):
1006
+ while True:
1007
+ elements.append(self.parse_expr())
1008
+ if not self.accept_op(","):
1009
+ break
1010
+ self.expect_op("]")
1011
+ return ArrayExpr(elements, elem_type)
1012
+
1013
+ def parse_struct(self) -> StructExpr:
1014
+ self.expect_kw("STRUCT")
1015
+ type_fields = None
1016
+ if self.at_op("<"):
1017
+ self.advance()
1018
+ type_fields = self._parse_struct_type_fields()
1019
+ self.expect_op("(")
1020
+ fields: List[StructFieldValue] = []
1021
+ if not self.at_op(")"):
1022
+ while True:
1023
+ e = self.parse_expr()
1024
+ name = None
1025
+ if self.accept_kw("AS"):
1026
+ name = self.parse_ident()
1027
+ fields.append(StructFieldValue(e, name))
1028
+ if not self.accept_op(","):
1029
+ break
1030
+ self.expect_op(")")
1031
+ return StructExpr(fields, type_fields)
1032
+
1033
+ def parse_interval(self) -> IntervalExpr:
1034
+ self.expect_kw("INTERVAL")
1035
+ value = self.parse_additive()
1036
+ unit = self.advance().value.upper()
1037
+ to_unit = None
1038
+ if self.accept_kw("TO"):
1039
+ to_unit = self.advance().value.upper()
1040
+ return IntervalExpr(value, unit, to_unit)
1041
+
1042
+ def parse_func_call(self, path: List[str]) -> FuncCall:
1043
+ self.expect_op("(")
1044
+ distinct = bool(self.accept_kw("DISTINCT"))
1045
+ args: List[Node] = []
1046
+ nulls = None
1047
+ order_by: List[OrderItem] = []
1048
+ limit = None
1049
+ having = None
1050
+ if not self.at_op(")"):
1051
+ if self.at_op("*"):
1052
+ self.advance()
1053
+ args.append(Star())
1054
+ else:
1055
+ while True:
1056
+ if self._is_ident() and self.at_op("=>", k=1):
1057
+ name = self.advance().value
1058
+ self.advance()
1059
+ args.append(NamedArg(name, self.parse_expr()))
1060
+ else:
1061
+ args.append(self.parse_expr())
1062
+ if not self.accept_op(","):
1063
+ break
1064
+ # aggregate modifiers, in any order, before ')'
1065
+ while True:
1066
+ if self.at_kw("IGNORE", "RESPECT") and self.at_kw("NULLS", k=1):
1067
+ nulls = self.advance().value.upper()
1068
+ self.advance()
1069
+ elif self.at_kw("HAVING") and self.at_kw("MAX", "MIN", k=1):
1070
+ self.advance()
1071
+ kind = self.advance().value.upper()
1072
+ having = HavingModifier(kind, self.parse_expr())
1073
+ elif self.at_kw("ORDER"):
1074
+ self.advance()
1075
+ self.expect_kw("BY")
1076
+ order_by = self.parse_order_items()
1077
+ elif self.at_kw("LIMIT"):
1078
+ self.advance()
1079
+ limit = self.parse_expr()
1080
+ else:
1081
+ break
1082
+ self.expect_op(")")
1083
+ over = None
1084
+ if self.at_kw("OVER"):
1085
+ self.advance()
1086
+ if self.at_op("("):
1087
+ self.advance()
1088
+ over = self.parse_window_spec_body()
1089
+ self.expect_op(")")
1090
+ else:
1091
+ over = WindowSpec(name=self.parse_ident(), is_ref=True)
1092
+ return FuncCall(path, args, distinct, nulls, order_by, limit, having, over)
1093
+
1094
+ def parse_window_spec_body(self) -> WindowSpec:
1095
+ name = None
1096
+ if self._is_ident() and not self.at_kw("PARTITION", "ORDER", "ROWS", "RANGE", "GROUPS"):
1097
+ name = self.advance().value
1098
+ partition_by: List[Node] = []
1099
+ if self.at_kw("PARTITION"):
1100
+ self.advance()
1101
+ self.expect_kw("BY")
1102
+ while True:
1103
+ partition_by.append(self.parse_expr())
1104
+ if not self.accept_op(","):
1105
+ break
1106
+ order_by = self.parse_order_by_opt()
1107
+ frame = None
1108
+ if self.at_kw("ROWS", "RANGE", "GROUPS"):
1109
+ unit = self.advance().value.upper()
1110
+ if self.accept_kw("BETWEEN"):
1111
+ start = self.parse_frame_bound()
1112
+ self.expect_kw("AND")
1113
+ end = self.parse_frame_bound()
1114
+ frame = WindowFrame(unit, start, end)
1115
+ else:
1116
+ frame = WindowFrame(unit, self.parse_frame_bound())
1117
+ return WindowSpec(name, partition_by, order_by, frame)
1118
+
1119
+ def parse_frame_bound(self) -> FrameBound:
1120
+ if self.accept_kw("UNBOUNDED"):
1121
+ kind = self.expect_kw("PRECEDING", "FOLLOWING").value.upper()
1122
+ return FrameBound(f"UNBOUNDED {kind}")
1123
+ if self.accept_kw("CURRENT"):
1124
+ self.expect_kw("ROW")
1125
+ return FrameBound("CURRENT ROW")
1126
+ value = self.parse_expr()
1127
+ kind = self.expect_kw("PRECEDING", "FOLLOWING").value.upper()
1128
+ return FrameBound(kind, value)
1129
+
1130
+ # ------------------------------------------------------------------ types
1131
+
1132
+ def parse_type(self) -> TypeNode:
1133
+ if self.at_kw("ARRAY"):
1134
+ self.advance()
1135
+ self.expect_op("<")
1136
+ elem = self.parse_type()
1137
+ self._expect_gt()
1138
+ return TypeNode("ARRAY", element=elem)
1139
+ if self.at_kw("STRUCT"):
1140
+ self.advance()
1141
+ fields: List[StructFieldType] = []
1142
+ if self.at_op("<"):
1143
+ self.advance()
1144
+ fields = self._parse_struct_type_fields()
1145
+ return TypeNode("STRUCT", fields_=fields)
1146
+ if self.at_kw("RANGE"):
1147
+ self.advance()
1148
+ self.expect_op("<")
1149
+ elem = self.parse_type()
1150
+ self._expect_gt()
1151
+ return TypeNode("RANGE", element=elem)
1152
+ if self.at_kw("INTERVAL"):
1153
+ self.advance()
1154
+ return TypeNode("INTERVAL")
1155
+ if self.at_kw("ANY"):
1156
+ self.advance()
1157
+ self.expect_kw("TYPE")
1158
+ return TypeNode("ANY TYPE")
1159
+ tok = self.peek()
1160
+ if tok.type not in (IDENT, QIDENT):
1161
+ self.error("Expected type name")
1162
+ name = self.advance().value.upper()
1163
+ params: List[str] = []
1164
+ if self.at_op("("):
1165
+ self.advance()
1166
+ while True:
1167
+ p = self.advance()
1168
+ params.append(p.value)
1169
+ if not self.accept_op(","):
1170
+ break
1171
+ self.expect_op(")")
1172
+ return TypeNode(name, params=params)
1173
+
1174
+ def _parse_struct_type_fields(self) -> List[StructFieldType]:
1175
+ fields: List[StructFieldType] = []
1176
+ while True:
1177
+ fname: Optional[str] = None
1178
+ tok = self.peek()
1179
+ if tok.type in (IDENT, QIDENT) and (
1180
+ self.peek(1).type in (IDENT, QIDENT)
1181
+ or self.at_kw("ARRAY", "STRUCT", "RANGE", "INTERVAL", k=1)
1182
+ or self.at_op("<", k=1)
1183
+ ):
1184
+ # "name type" — but "name <" only if name isn't itself generic
1185
+ fname = self.advance().value
1186
+ fields.append(StructFieldType(fname, self.parse_type()))
1187
+ if not self.accept_op(","):
1188
+ break
1189
+ self._expect_gt()
1190
+ return fields
1191
+
1192
+ # ------------------------------------------------------------- scripting
1193
+
1194
+ _SCRIPT_TERMINATORS = ("END", "ELSEIF", "ELSE", "EXCEPTION", "UNTIL", "WHEN")
1195
+
1196
+ def parse_statement_list(self) -> List[Node]:
1197
+ stmts: List[Node] = []
1198
+ while True:
1199
+ while self.accept_op(";"):
1200
+ pass
1201
+ if self.peek().type == EOF or self.at_kw(*self._SCRIPT_TERMINATORS):
1202
+ break
1203
+ stmts.append(self.parse_statement())
1204
+ return stmts
1205
+
1206
+ def _parse_labeled(self, label: str) -> Node:
1207
+ if self.at_kw("LOOP"):
1208
+ stmt = self.parse_loop()
1209
+ elif self.at_kw("WHILE"):
1210
+ stmt = self.parse_while()
1211
+ elif self.at_kw("REPEAT"):
1212
+ stmt = self.parse_repeat()
1213
+ elif self.at_kw("FOR"):
1214
+ stmt = self.parse_for_in()
1215
+ elif self.at_kw("BEGIN"):
1216
+ stmt = self.parse_begin()
1217
+ else:
1218
+ self.error("Expected LOOP, WHILE, REPEAT, FOR or BEGIN after label")
1219
+ stmt.label = label
1220
+ self._accept_trailing_label(stmt)
1221
+ return stmt
1222
+
1223
+ def _accept_trailing_label(self, label_holder: Node):
1224
+ if getattr(label_holder, "label", None) and self._is_ident() \
1225
+ and self.peek().value.lower() == label_holder.label.lower():
1226
+ self.advance()
1227
+
1228
+ def parse_declare(self) -> DeclareStmt:
1229
+ self.expect_kw("DECLARE")
1230
+ names = [self.parse_ident()]
1231
+ while self.accept_op(","):
1232
+ names.append(self.parse_ident())
1233
+ var_type = None
1234
+ default = None
1235
+ if self.accept_kw("DEFAULT"):
1236
+ default = self.parse_expr()
1237
+ elif not self.at_op(";") and self.peek().type != EOF:
1238
+ var_type = self.parse_type()
1239
+ if self.accept_kw("DEFAULT"):
1240
+ default = self.parse_expr()
1241
+ return DeclareStmt(names, var_type, default)
1242
+
1243
+ def parse_set_stmt(self) -> SetStmt:
1244
+ self.expect_kw("SET")
1245
+ targets: List[str] = []
1246
+ if self.accept_op("("):
1247
+ while True:
1248
+ targets.append(self.parse_ident())
1249
+ if not self.accept_op(","):
1250
+ break
1251
+ self.expect_op(")")
1252
+ elif self.peek().type == PARAM:
1253
+ targets.append(self.advance().value)
1254
+ else:
1255
+ targets.append(self.parse_ident())
1256
+ self.expect_op("=")
1257
+ return SetStmt(targets, self.parse_expr())
1258
+
1259
+ def parse_begin(self) -> Node:
1260
+ self.expect_kw("BEGIN")
1261
+ if self.accept_kw("TRANSACTION"):
1262
+ return TransactionStmt("BEGIN")
1263
+ if self.at_op(";") or self.peek().type == EOF:
1264
+ return TransactionStmt("BEGIN")
1265
+ statements = self.parse_statement_list()
1266
+ exception_statements: List[Node] = []
1267
+ has_handler = False
1268
+ if self.accept_kw("EXCEPTION"):
1269
+ self.expect_kw("WHEN")
1270
+ self.expect_kw("ERROR")
1271
+ self.expect_kw("THEN")
1272
+ has_handler = True
1273
+ exception_statements = self.parse_statement_list()
1274
+ self.expect_kw("END")
1275
+ block = ScriptBlock(statements, exception_statements, has_handler)
1276
+ self._accept_trailing_label(block)
1277
+ return block
1278
+
1279
+ def parse_if(self) -> IfStmt:
1280
+ self.expect_kw("IF")
1281
+ cond = self.parse_expr()
1282
+ self.expect_kw("THEN")
1283
+ branches = [IfBranch(cond, self.parse_statement_list())]
1284
+ while self.at_kw("ELSEIF"):
1285
+ self.advance()
1286
+ cond = self.parse_expr()
1287
+ self.expect_kw("THEN")
1288
+ branches.append(IfBranch(cond, self.parse_statement_list()))
1289
+ else_statements: List[Node] = []
1290
+ if self.accept_kw("ELSE"):
1291
+ else_statements = self.parse_statement_list()
1292
+ self.expect_kw("END")
1293
+ self.expect_kw("IF")
1294
+ return IfStmt(branches, else_statements)
1295
+
1296
+ def parse_loop(self) -> LoopStmt:
1297
+ self.expect_kw("LOOP")
1298
+ stmts = self.parse_statement_list()
1299
+ self.expect_kw("END")
1300
+ self.expect_kw("LOOP")
1301
+ loop = LoopStmt(stmts)
1302
+ self._accept_trailing_label(loop)
1303
+ return loop
1304
+
1305
+ def parse_while(self) -> WhileStmt:
1306
+ self.expect_kw("WHILE")
1307
+ cond = self.parse_expr()
1308
+ self.expect_kw("DO")
1309
+ stmts = self.parse_statement_list()
1310
+ self.expect_kw("END")
1311
+ self.expect_kw("WHILE")
1312
+ loop = WhileStmt(cond, stmts)
1313
+ self._accept_trailing_label(loop)
1314
+ return loop
1315
+
1316
+ def parse_repeat(self) -> RepeatStmt:
1317
+ self.expect_kw("REPEAT")
1318
+ stmts = self.parse_statement_list()
1319
+ self.expect_kw("UNTIL")
1320
+ until = self.parse_expr()
1321
+ self.expect_kw("END")
1322
+ self.expect_kw("REPEAT")
1323
+ loop = RepeatStmt(stmts, until)
1324
+ self._accept_trailing_label(loop)
1325
+ return loop
1326
+
1327
+ def parse_for_in(self) -> ForInStmt:
1328
+ self.expect_kw("FOR")
1329
+ var = self.parse_ident()
1330
+ self.expect_kw("IN")
1331
+ self.expect_op("(")
1332
+ query = self.parse_query()
1333
+ self.expect_op(")")
1334
+ self.expect_kw("DO")
1335
+ stmts = self.parse_statement_list()
1336
+ self.expect_kw("END")
1337
+ self.expect_kw("FOR")
1338
+ loop = ForInStmt(var, query, stmts)
1339
+ self._accept_trailing_label(loop)
1340
+ return loop
1341
+
1342
+ def parse_execute_immediate(self) -> ExecuteImmediate:
1343
+ self.expect_kw("EXECUTE")
1344
+ self.expect_kw("IMMEDIATE")
1345
+ sql_expr = self.parse_expr()
1346
+ into: List[str] = []
1347
+ using: List = []
1348
+ if self.accept_kw("INTO"):
1349
+ while True:
1350
+ into.append(self.parse_ident())
1351
+ if not self.accept_op(","):
1352
+ break
1353
+ if self.accept_kw("USING"):
1354
+ while True:
1355
+ e = self.parse_expr()
1356
+ alias = None
1357
+ if self.accept_kw("AS"):
1358
+ alias = self.parse_ident()
1359
+ using.append((e, alias))
1360
+ if not self.accept_op(","):
1361
+ break
1362
+ return ExecuteImmediate(sql_expr, into, using)
1363
+
1364
+ def parse_drop(self) -> DropStmt:
1365
+ self.expect_kw("DROP")
1366
+ kind = self.advance().value.upper()
1367
+ if kind in ("MATERIALIZED", "EXTERNAL", "SNAPSHOT", "SEARCH", "VECTOR"):
1368
+ kind += " " + self.advance().value.upper()
1369
+ elif kind == "TABLE" and self.at_kw("FUNCTION"):
1370
+ kind += " " + self.advance().value.upper()
1371
+ if_exists = False
1372
+ if self.at_kw("IF") and self.at_kw("EXISTS", k=1):
1373
+ self.advance()
1374
+ self.advance()
1375
+ if_exists = True
1376
+ return DropStmt(kind, self._split_path(self._parse_dotted_path()), if_exists)
1377
+
1378
+ # ------------------------------------------------------------- DDL / DML
1379
+
1380
+ def _skip_balanced_parens(self) -> str:
1381
+ start_tok = self.expect_op("(")
1382
+ depth = 1
1383
+ start = start_tok.pos
1384
+ while depth > 0:
1385
+ tok = self.advance()
1386
+ if tok.type == EOF:
1387
+ self.error("Unbalanced parentheses")
1388
+ if tok.type == OP and tok.value == "(":
1389
+ depth += 1
1390
+ elif tok.type == OP and tok.value == ")":
1391
+ depth -= 1
1392
+ end = tok.pos + 1
1393
+ return self.sql[start:end]
1394
+
1395
+ def parse_create(self) -> Node:
1396
+ self.expect_kw("CREATE")
1397
+ replace = False
1398
+ if self.accept_kw("OR"):
1399
+ self.expect_kw("REPLACE")
1400
+ replace = True
1401
+ temp = bool(self.accept_kw("TEMP", "TEMPORARY"))
1402
+ if self.at_kw("PROCEDURE"):
1403
+ return self._parse_create_procedure(replace)
1404
+ if self.at_kw("AGGREGATE") and self.at_kw("FUNCTION", k=1):
1405
+ self.advance()
1406
+ self.advance()
1407
+ return self._parse_create_function(replace, temp, aggregate=True)
1408
+ if self.at_kw("FUNCTION"):
1409
+ self.advance()
1410
+ return self._parse_create_function(replace, temp)
1411
+ if self.at_kw("TABLE") and self.at_kw("FUNCTION", k=1):
1412
+ self.advance()
1413
+ self.advance()
1414
+ return self._parse_create_function(replace, temp, table_function=True)
1415
+ kind = "TABLE"
1416
+ if self.accept_kw("MATERIALIZED"):
1417
+ self.expect_kw("VIEW")
1418
+ kind = "MATERIALIZED VIEW"
1419
+ elif self.accept_kw("VIEW"):
1420
+ kind = "VIEW"
1421
+ else:
1422
+ self.expect_kw("TABLE")
1423
+ if_not_exists = False
1424
+ if self.at_kw("IF") and self.at_kw("NOT", k=1):
1425
+ self.advance()
1426
+ self.advance()
1427
+ self.expect_kw("EXISTS")
1428
+ if_not_exists = True
1429
+ name = self._parse_dotted_path()
1430
+ parts: List[str] = []
1431
+ for p in name:
1432
+ parts.extend(p.split(".") if "." in p else [p])
1433
+
1434
+ columns: List[ColumnDef] = []
1435
+ if self.at_op("("):
1436
+ self.advance()
1437
+ while True:
1438
+ cname = self.parse_ident()
1439
+ ctype = None
1440
+ if kind == "TABLE" or self.peek().type in (IDENT, QIDENT) \
1441
+ or self.at_kw("ARRAY", "STRUCT", "RANGE", "INTERVAL"):
1442
+ if not self.at_op(",", ")"):
1443
+ ctype = self.parse_type()
1444
+ # skip column attributes (NOT NULL, OPTIONS(...), DEFAULT ...)
1445
+ depth = 0
1446
+ while not (depth == 0 and self.at_op(",", ")")):
1447
+ tok = self.advance()
1448
+ if tok.type == EOF:
1449
+ self.error("Unterminated column definition list")
1450
+ if tok.type == OP and tok.value in "([":
1451
+ depth += 1
1452
+ elif tok.type == OP and tok.value in ")]":
1453
+ depth -= 1
1454
+ columns.append(ColumnDef(cname, ctype))
1455
+ if not self.accept_op(","):
1456
+ break
1457
+ self.expect_op(")")
1458
+
1459
+ partition_by = None
1460
+ cluster_by: List[Node] = []
1461
+ options_sql = None
1462
+ while True:
1463
+ if self.at_kw("PARTITION"):
1464
+ self.advance()
1465
+ self.expect_kw("BY")
1466
+ partition_by = self.parse_expr()
1467
+ elif self.at_kw("CLUSTER"):
1468
+ self.advance()
1469
+ self.expect_kw("BY")
1470
+ while True:
1471
+ cluster_by.append(self.parse_expr())
1472
+ if not self.accept_op(","):
1473
+ break
1474
+ elif self.at_kw("OPTIONS"):
1475
+ self.advance()
1476
+ options_sql = self._skip_balanced_parens()
1477
+ else:
1478
+ break
1479
+ query = None
1480
+ if self.accept_kw("AS"):
1481
+ query = self.parse_query()
1482
+ return CreateTableAsSelect(kind, parts, query, replace, temp,
1483
+ if_not_exists, columns, partition_by,
1484
+ cluster_by, options_sql)
1485
+
1486
+ def _parse_if_not_exists(self) -> bool:
1487
+ if self.at_kw("IF") and self.at_kw("NOT", k=1):
1488
+ self.advance()
1489
+ self.advance()
1490
+ self.expect_kw("EXISTS")
1491
+ return True
1492
+ return False
1493
+
1494
+ def _parse_routine_params(self, allow_modes: bool) -> List[RoutineParam]:
1495
+ params: List[RoutineParam] = []
1496
+ self.expect_op("(")
1497
+ if not self.at_op(")"):
1498
+ while True:
1499
+ mode = None
1500
+ if allow_modes and self.at_kw("IN", "OUT", "INOUT") \
1501
+ and self._is_ident(k=1) and not self.at_op(",", ")", k=2):
1502
+ mode = self.advance().value.upper()
1503
+ pname = self.parse_ident()
1504
+ ptype = None
1505
+ if not self.at_op(",", ")"):
1506
+ ptype = self.parse_type()
1507
+ params.append(RoutineParam(pname, ptype, mode))
1508
+ if not self.accept_op(","):
1509
+ break
1510
+ self.expect_op(")")
1511
+ return params
1512
+
1513
+ def _parse_create_procedure(self, replace: bool) -> CreateProcedure:
1514
+ self.expect_kw("PROCEDURE")
1515
+ if_not_exists = self._parse_if_not_exists()
1516
+ name = self._split_path(self._parse_dotted_path())
1517
+ params = self._parse_routine_params(allow_modes=True)
1518
+ options_sql = None
1519
+ if self.at_kw("OPTIONS"):
1520
+ self.advance()
1521
+ options_sql = self._skip_balanced_parens()
1522
+ body = self.parse_begin()
1523
+ if not isinstance(body, ScriptBlock):
1524
+ self.error("Expected BEGIN ... END procedure body")
1525
+ return CreateProcedure(name, params, body, replace, if_not_exists,
1526
+ options_sql)
1527
+
1528
+ def _parse_create_function(self, replace: bool, temp: bool,
1529
+ table_function: bool = False,
1530
+ aggregate: bool = False) -> CreateFunction:
1531
+ if_not_exists = self._parse_if_not_exists()
1532
+ name = self._split_path(self._parse_dotted_path())
1533
+ params = self._parse_routine_params(allow_modes=False)
1534
+ fn = CreateFunction(name, params, replace=replace, temp=temp,
1535
+ if_not_exists=if_not_exists,
1536
+ table_function=table_function, aggregate=aggregate)
1537
+ while True:
1538
+ if self.accept_kw("RETURNS"):
1539
+ if self.at_kw("TABLE"):
1540
+ self.advance()
1541
+ self.expect_op("<")
1542
+ fn.returns_table = self._parse_struct_type_fields()
1543
+ else:
1544
+ fn.returns = self.parse_type()
1545
+ elif self.at_kw("DETERMINISTIC"):
1546
+ self.advance()
1547
+ fn.deterministic = True
1548
+ elif self.at_kw("NOT") and self.at_kw("DETERMINISTIC", k=1):
1549
+ self.advance()
1550
+ self.advance()
1551
+ fn.deterministic = False
1552
+ elif self.accept_kw("LANGUAGE"):
1553
+ fn.language = self.parse_ident()
1554
+ elif self.at_kw("OPTIONS"):
1555
+ self.advance()
1556
+ fn.options_sql = self._skip_balanced_parens()
1557
+ elif self.accept_kw("AS"):
1558
+ if self.at_op("("):
1559
+ self.advance()
1560
+ if self._starts_query():
1561
+ fn.body_query = self.parse_query()
1562
+ else:
1563
+ fn.body_expr = self.parse_expr()
1564
+ self.expect_op(")")
1565
+ else:
1566
+ tok = self.peek()
1567
+ if tok.type not in (STRING, BYTES):
1568
+ self.error("Expected function body")
1569
+ fn.body_string = self.advance().value
1570
+ else:
1571
+ break
1572
+ return fn
1573
+
1574
+ def parse_insert(self) -> InsertStmt:
1575
+ self.expect_kw("INSERT")
1576
+ self.accept_kw("INTO")
1577
+ table = self._split_path(self._parse_dotted_path())
1578
+ columns: List[str] = []
1579
+ if self.at_op("(") and not self._starts_query():
1580
+ self.advance()
1581
+ while True:
1582
+ columns.append(self.parse_ident())
1583
+ if not self.accept_op(","):
1584
+ break
1585
+ self.expect_op(")")
1586
+ if self.at_kw("VALUES"):
1587
+ self.advance()
1588
+ values: List[List[Node]] = []
1589
+ while True:
1590
+ row = self._paren_expr_list()
1591
+ values.append(row)
1592
+ if not self.accept_op(","):
1593
+ break
1594
+ return InsertStmt(table, columns, values=values)
1595
+ query = self.parse_query()
1596
+ return InsertStmt(table, columns, query=query)
1597
+
1598
+ def parse_update(self) -> UpdateStmt:
1599
+ self.expect_kw("UPDATE")
1600
+ table = self._split_path(self._parse_dotted_path())
1601
+ alias = None
1602
+ if self.accept_kw("AS"):
1603
+ alias = self.parse_ident()
1604
+ elif self._is_ident() and not self.at_kw("SET"):
1605
+ alias = self.advance().value
1606
+ self.expect_kw("SET")
1607
+ assignments = self._parse_assignments()
1608
+ from_ = None
1609
+ if self.accept_kw("FROM"):
1610
+ from_ = self.parse_from()
1611
+ where = None
1612
+ if self.accept_kw("WHERE"):
1613
+ where = self.parse_expr()
1614
+ return UpdateStmt(table, assignments, alias, from_, where)
1615
+
1616
+ def _parse_assignments(self) -> List[Assignment]:
1617
+ assignments: List[Assignment] = []
1618
+ while True:
1619
+ target = self._parse_dotted_path()
1620
+ self.expect_op("=")
1621
+ assignments.append(Assignment(target, self.parse_expr()))
1622
+ if not self.accept_op(","):
1623
+ break
1624
+ return assignments
1625
+
1626
+ def parse_delete(self) -> DeleteStmt:
1627
+ self.expect_kw("DELETE")
1628
+ self.accept_kw("FROM")
1629
+ table = self._split_path(self._parse_dotted_path())
1630
+ alias = None
1631
+ if self.accept_kw("AS"):
1632
+ alias = self.parse_ident()
1633
+ elif self._is_ident() and not self.at_kw("WHERE"):
1634
+ alias = self.advance().value
1635
+ where = None
1636
+ if self.accept_kw("WHERE"):
1637
+ where = self.parse_expr()
1638
+ return DeleteStmt(table, alias, where)
1639
+
1640
+ def parse_merge(self) -> MergeStmt:
1641
+ self.expect_kw("MERGE")
1642
+ self.accept_kw("INTO")
1643
+ target = self._split_path(self._parse_dotted_path())
1644
+ alias = None
1645
+ if self.accept_kw("AS"):
1646
+ alias = self.parse_ident()
1647
+ elif self._is_ident() and not self.at_kw("USING"):
1648
+ alias = self.advance().value
1649
+ self.expect_kw("USING")
1650
+ source = self.parse_table_primary()
1651
+ self.expect_kw("ON")
1652
+ on = self.parse_expr()
1653
+ whens: List[MergeWhen] = []
1654
+ while self.at_kw("WHEN"):
1655
+ self.advance()
1656
+ if self.accept_kw("MATCHED"):
1657
+ match_kind = "MATCHED"
1658
+ else:
1659
+ self.expect_kw("NOT")
1660
+ self.expect_kw("MATCHED")
1661
+ match_kind = "NOT_MATCHED"
1662
+ if self.accept_kw("BY"):
1663
+ which = self.expect_kw("TARGET", "SOURCE").value.upper()
1664
+ if which == "SOURCE":
1665
+ match_kind = "NOT_MATCHED_BY_SOURCE"
1666
+ condition = None
1667
+ if self.accept_kw("AND"):
1668
+ condition = self.parse_expr()
1669
+ self.expect_kw("THEN")
1670
+ if self.accept_kw("UPDATE"):
1671
+ self.expect_kw("SET")
1672
+ whens.append(MergeWhen(match_kind, "UPDATE", condition,
1673
+ assignments=self._parse_assignments()))
1674
+ elif self.accept_kw("DELETE"):
1675
+ whens.append(MergeWhen(match_kind, "DELETE", condition))
1676
+ else:
1677
+ self.expect_kw("INSERT")
1678
+ if self.accept_kw("ROW"):
1679
+ whens.append(MergeWhen(match_kind, "INSERT_ROW", condition))
1680
+ else:
1681
+ cols: List[str] = []
1682
+ if self.at_op("("):
1683
+ self.advance()
1684
+ while True:
1685
+ cols.append(self.parse_ident())
1686
+ if not self.accept_op(","):
1687
+ break
1688
+ self.expect_op(")")
1689
+ self.expect_kw("VALUES")
1690
+ vals = self._paren_expr_list()
1691
+ whens.append(MergeWhen(match_kind, "INSERT", condition,
1692
+ columns=cols, values=vals))
1693
+ return MergeStmt(target, source, on, whens, alias)
1694
+
1695
+ @staticmethod
1696
+ def _split_path(path: List[str]) -> List[str]:
1697
+ parts: List[str] = []
1698
+ for p in path:
1699
+ parts.extend(p.split(".") if "." in p else [p])
1700
+ return parts
1701
+
1702
+
1703
+ def parse(sql: str) -> List[Node]:
1704
+ """Parse one or more ``;``-separated statements into AST nodes."""
1705
+ return Parser(sql).parse_statements()
1706
+
1707
+
1708
+ def parse_one(sql: str) -> Node:
1709
+ """Parse exactly one statement and return its AST."""
1710
+ stmts = parse(sql)
1711
+ if not stmts:
1712
+ raise ParseError("No statement found")
1713
+ if len(stmts) > 1:
1714
+ raise ParseError(f"Expected a single statement, found {len(stmts)}")
1715
+ return stmts[0]