sqlglot.transforms
1from __future__ import annotations 2 3import typing as t 4 5from sqlglot import expressions as exp 6from sqlglot.errors import UnsupportedError 7from sqlglot.helper import find_new_name, name_sequence, seq_get 8 9 10if t.TYPE_CHECKING: 11 from sqlglot._typing import E 12 from sqlglot.generator import Generator 13 14 15class SqlHandler(t.Protocol): 16 def __call__(self, expression: exp.Expr, *args: t.Any, **kwargs: t.Any) -> str: ... 17 18 19def preprocess( 20 transforms: list[t.Callable[[exp.Expr], exp.Expr]], 21 generator: t.Callable[[Generator, exp.Expr], str] | None = None, 22) -> t.Callable[[Generator, exp.Expr], str]: 23 """ 24 Creates a new transform by chaining a sequence of transformations and converts the resulting 25 expression to SQL, using either the "_sql" method corresponding to the resulting expression, 26 or the appropriate `Generator.TRANSFORMS` function (when applicable -- see below). 27 28 Args: 29 transforms: sequence of transform functions. These will be called in order. 30 31 Returns: 32 Function that can be used as a generator transform. 33 """ 34 35 def _to_sql(self: Generator, expression: exp.Expr) -> str: 36 expression_type = type(expression) 37 38 try: 39 expression = transforms[0](expression) 40 for transform in transforms[1:]: 41 expression = transform(expression) 42 except UnsupportedError as unsupported_error: 43 self.unsupported(str(unsupported_error)) 44 45 if generator: 46 return generator(self, expression) 47 48 _sql_handler: SqlHandler | None = getattr(self, expression.key + "_sql", None) 49 if _sql_handler: 50 return _sql_handler(expression) 51 52 transforms_handler = self.TRANSFORMS.get(type(expression)) 53 if transforms_handler: 54 if expression_type is type(expression): 55 if isinstance(expression, exp.Func): 56 return self.function_fallback_sql(expression) 57 58 # Ensures we don't enter an infinite loop. This can happen when the original expression 59 # has the same type as the final expression and there's no _sql method available for it, 60 # because then it'd re-enter _to_sql. 61 raise ValueError( 62 f"Expr type {expression.__class__.__name__} requires a _sql method in order to be transformed." 63 ) 64 65 return transforms_handler(self, expression) 66 67 raise ValueError(f"Unsupported expression type {expression.__class__.__name__}.") 68 69 return _to_sql 70 71 72def unnest_generate_date_array_using_recursive_cte(expression: exp.Expr) -> exp.Expr: 73 if isinstance(expression, exp.Select): 74 count = 0 75 recursive_ctes: list[exp.Expr] = [] 76 77 for unnest in expression.find_all(exp.Unnest): 78 if ( 79 not isinstance(unnest.parent, (exp.From, exp.Join)) 80 or len(unnest.expressions) != 1 81 or not isinstance(unnest.expressions[0], exp.GenerateDateArray) 82 ): 83 continue 84 85 generate_date_array = unnest.expressions[0] 86 start: exp.Expr | None = generate_date_array.args.get("start") 87 end: exp.Expr | None = generate_date_array.args.get("end") 88 step: exp.Expr | None = generate_date_array.args.get("step") 89 90 if not start or not end or not isinstance(step, exp.Interval): 91 continue 92 93 alias: exp.TableAlias | None = unnest.args.get("alias") 94 column_name: str = ( 95 alias.columns[0] if isinstance(alias, exp.TableAlias) else "date_value" 96 ) 97 98 start = exp.cast(start, "date") 99 date_add = exp.func( 100 "date_add", column_name, exp.Literal.number(step.name), step.args.get("unit") 101 ) 102 cast_date_add = exp.cast(date_add, "date") 103 104 cte_name = "_generated_dates" + (f"_{count}" if count else "") 105 106 base_query = exp.select(start.as_(column_name)) 107 recursive_query = ( 108 exp.select(cast_date_add) 109 .from_(cte_name) 110 .where(cast_date_add <= exp.cast(end, "date")) 111 ) 112 cte_query = base_query.union(recursive_query, distinct=False) 113 114 generate_dates_query = exp.select(column_name).from_(cte_name) 115 unnest.replace(generate_dates_query.subquery(cte_name)) 116 117 recursive_ctes.append( 118 exp.alias_(exp.CTE(this=cte_query), cte_name, table=[column_name]) 119 ) 120 count += 1 121 122 if recursive_ctes: 123 with_expression: exp.With = expression.args.get("with_") or exp.With() 124 with_expression.set("recursive", True) 125 with_expression.set("expressions", [*recursive_ctes, *with_expression.expressions]) 126 expression.set("with_", with_expression) 127 128 return expression 129 130 131def unnest_generate_series(expression: exp.Expr) -> exp.Expr: 132 """Unnests GENERATE_SERIES or SEQUENCE table references.""" 133 this = expression.this 134 if isinstance(expression, exp.Table) and isinstance(this, exp.GenerateSeries): 135 unnest = exp.Unnest(expressions=[this]) 136 if expression.alias: 137 return exp.alias_(unnest, alias="_u", table=[expression.alias], copy=False) 138 139 return unnest 140 141 return expression 142 143 144def eliminate_distinct_on(expression: exp.Expr) -> exp.Expr: 145 """ 146 Convert SELECT DISTINCT ON statements to a subquery with a window function. 147 148 This is useful for dialects that don't support SELECT DISTINCT ON but support window functions. 149 150 Args: 151 expression: the expression that will be transformed. 152 153 Returns: 154 The transformed expression. 155 """ 156 if ( 157 isinstance(expression, exp.Select) 158 and expression.args.get("distinct") 159 and isinstance(expression.args["distinct"].args.get("on"), exp.Tuple) 160 ): 161 row_number_window_alias = find_new_name(expression.named_selects, "_row_number") 162 163 distinct_cols = expression.args["distinct"].pop().args["on"].expressions 164 window = exp.Window(this=exp.RowNumber(), partition_by=distinct_cols) 165 166 order: exp.Order | None = expression.args.get("order") 167 if order: 168 window.set("order", order.pop()) 169 else: 170 window.set("order", exp.Order(expressions=[c.copy() for c in distinct_cols])) 171 172 expression.select(exp.alias_(window, row_number_window_alias), copy=False) 173 174 # We add aliases to the projections so that we can safely reference them in the outer query 175 new_selects: list[exp.Expr] = [] 176 taken_names = {row_number_window_alias} 177 for select in expression.selects[:-1]: 178 if select.is_star: 179 new_selects = [exp.Star()] 180 break 181 182 if not isinstance(select, exp.Alias): 183 alias = find_new_name(taken_names, select.output_name or "_col") 184 quoted: bool | None = ( 185 select.this.args.get("quoted") if isinstance(select, exp.Column) else None 186 ) 187 select = select.replace(exp.alias_(select, alias, quoted=quoted)) 188 189 taken_names.add(select.output_name) 190 new_selects.append(select.args["alias"]) 191 192 return ( 193 exp.select(*new_selects, copy=False) 194 .from_(expression.subquery("_t", copy=False), copy=False) 195 .where(exp.column(row_number_window_alias).eq(1), copy=False) 196 ) 197 198 return expression 199 200 201def eliminate_qualify(expression: exp.Expr) -> exp.Expr: 202 """ 203 Convert SELECT statements that contain the QUALIFY clause into subqueries, filtered equivalently. 204 205 The idea behind this transformation can be seen in Snowflake's documentation for QUALIFY: 206 https://docs.snowflake.com/en/sql-reference/constructs/qualify 207 208 Some dialects don't support window functions in the WHERE clause, so we need to include them as 209 projections in the subquery, in order to refer to them in the outer filter using aliases. Also, 210 if a column is referenced in the QUALIFY clause but is not selected, we need to include it too, 211 otherwise we won't be able to refer to it in the outer query's WHERE clause. Finally, if a 212 newly aliased projection is referenced in the QUALIFY clause, it will be replaced by the 213 corresponding expression to avoid creating invalid column references. 214 """ 215 if isinstance(expression, exp.Select) and expression.args.get("qualify"): 216 taken = set(expression.named_selects) 217 for select in expression.selects: 218 if not select.alias_or_name: 219 alias = find_new_name(taken, "_c") 220 select.replace(exp.alias_(select, alias)) 221 taken.add(alias) 222 223 def _select_alias_or_name(select: exp.Expr) -> str | exp.Column: 224 alias_or_name = select.alias_or_name 225 identifier = select.args.get("alias") or select.this 226 if isinstance(identifier, exp.Identifier): 227 return exp.column(alias_or_name, quoted=identifier.args.get("quoted")) 228 return alias_or_name 229 230 outer_selects = exp.select(*map(_select_alias_or_name, expression.selects)) 231 qualify_filters: exp.Expr = expression.args["qualify"].pop().this 232 expression_by_alias: dict[str, exp.Expr] = { 233 select.alias: select.this 234 for select in expression.selects 235 if isinstance(select, exp.Alias) 236 } 237 238 select_candidates = (exp.Window,) if expression.is_star else (exp.Window, exp.Column) 239 for select_candidate in list(qualify_filters.find_all(*select_candidates)): 240 if isinstance(select_candidate, exp.Window): 241 if expression_by_alias: 242 for column in select_candidate.find_all(exp.Column): 243 expr = expression_by_alias.get(column.name) 244 if expr: 245 column.replace(expr) 246 247 alias = find_new_name(expression.named_selects, "_w") 248 expression.select(exp.alias_(select_candidate, alias), copy=False) 249 column = exp.column(alias) 250 251 if isinstance(select_candidate.parent, exp.Qualify): 252 qualify_filters = column 253 else: 254 select_candidate.replace(column) 255 elif select_candidate.name not in expression.named_selects and not ( 256 select_candidate.find_ancestor(exp.Window) 257 ): 258 # A column that is only read by a window function is computed by the window 259 # itself, so projecting it in the subquery is redundant and can even produce 260 # invalid SQL, e.g. when the query is grouped 261 expression.select(select_candidate.copy(), copy=False) 262 263 return outer_selects.from_(expression.subquery(alias="_t", copy=False), copy=False).where( 264 qualify_filters, copy=False 265 ) 266 267 return expression 268 269 270def remove_precision_parameterized_types(expression: exp.Expr) -> exp.Expr: 271 """ 272 Some dialects only allow the precision for parameterized types to be defined in the DDL and not in 273 other expressions. This transforms removes the precision from parameterized types in expressions. 274 """ 275 for node in expression.find_all(exp.DataType): 276 node.set( 277 "expressions", [e for e in node.expressions if not isinstance(e, exp.DataTypeParam)] 278 ) 279 280 return expression 281 282 283def unqualify_unnest(expression: exp.Expr) -> exp.Expr: 284 """Remove references to unnest table aliases, added by the optimizer's qualify_columns step.""" 285 from sqlglot.optimizer.scope import find_all_in_scope 286 287 if isinstance(expression, exp.Select): 288 unnest_aliases = { 289 unnest.alias 290 for unnest in find_all_in_scope(expression, exp.Unnest) 291 if isinstance(unnest.parent, (exp.From, exp.Join)) 292 } 293 if unnest_aliases: 294 for column in expression.find_all(exp.Column): 295 leftmost_part = column.parts[0] 296 if leftmost_part.arg_key != "this" and leftmost_part.this in unnest_aliases: 297 leftmost_part.pop() 298 299 return expression 300 301 302def unnest_to_explode( 303 expression: exp.Expr, 304 unnest_using_arrays_zip: bool = True, 305) -> exp.Expr: 306 """Convert cross join unnest into lateral view explode.""" 307 308 def _unnest_zip_exprs( 309 u: exp.Unnest, unnest_exprs: list[exp.Expr], has_multi_expr: bool 310 ) -> list[exp.Expr]: 311 if has_multi_expr: 312 if not unnest_using_arrays_zip: 313 raise UnsupportedError("Cannot transpile UNNEST with multiple input arrays") 314 315 # Use INLINE(ARRAYS_ZIP(...)) for multiple expressions 316 zip_exprs: list[exp.Expr] = [exp.Anonymous(this="ARRAYS_ZIP", expressions=unnest_exprs)] 317 u.set("expressions", zip_exprs) 318 return zip_exprs 319 return unnest_exprs 320 321 def _udtf_type(u: exp.Unnest, has_multi_expr: bool) -> type[exp.Func]: 322 if u.args.get("offset"): 323 return exp.Posexplode 324 return exp.Inline if has_multi_expr else exp.Explode 325 326 if isinstance(expression, exp.Select): 327 from_ = expression.args.get("from_") 328 329 if from_ and isinstance(from_.this, exp.Unnest): 330 unnest: exp.Unnest = from_.this 331 alias: exp.TableAlias | None = unnest.args.get("alias") 332 exprs: list[exp.Expr] = unnest.expressions 333 has_multi_expr = len(exprs) > 1 334 this, *_ = _unnest_zip_exprs(unnest, exprs, has_multi_expr) 335 336 columns: list[exp.Identifier] = alias.columns if alias else [] 337 offset: exp.Expr | None = unnest.args.get("offset") 338 if offset: 339 columns.insert( 340 0, offset if isinstance(offset, exp.Identifier) else exp.to_identifier("pos") 341 ) 342 343 unnest.replace( 344 exp.Table( 345 this=_udtf_type(unnest, has_multi_expr)(this=this), 346 alias=exp.TableAlias(this=alias.this, columns=columns) if alias else None, 347 ) 348 ) 349 350 joins: list[exp.Join] = expression.args.get("joins") or [] 351 for join in list(joins): 352 join_expr = join.this 353 354 is_lateral = isinstance(join_expr, exp.Lateral) 355 356 unnest = join_expr.this if is_lateral else join_expr 357 358 if isinstance(unnest, exp.Unnest): 359 if is_lateral: 360 alias = join_expr.args.get("alias") 361 else: 362 alias = unnest.args.get("alias") 363 364 if alias is None: 365 raise UnsupportedError( 366 "CROSS JOIN UNNEST to LATERAL VIEW EXPLODE transformation requires an alias" 367 ) 368 369 exprs = unnest.expressions 370 # The number of unnest.expressions will be changed by _unnest_zip_exprs, we need to record it here 371 has_multi_expr = len(exprs) > 1 372 exprs = _unnest_zip_exprs(unnest, exprs, has_multi_expr) 373 374 joins.remove(join) 375 376 alias_cols: list[exp.Identifier] = alias.columns 377 378 # # Handle UNNEST to LATERAL VIEW EXPLODE: Exception is raised when there are 0 or > 2 aliases 379 # Spark LATERAL VIEW EXPLODE requires single alias for array/struct and two for Map type column unlike unnest in trino/presto which can take an arbitrary amount. 380 # Refs: https://spark.apache.org/docs/latest/sql-ref-syntax-qry-select-lateral-view.html 381 382 if not has_multi_expr and len(alias_cols) not in (1, 2): 383 raise UnsupportedError( 384 "CROSS JOIN UNNEST to LATERAL VIEW EXPLODE transformation requires explicit column aliases" 385 ) 386 387 offset = unnest.args.get("offset") 388 if offset: 389 alias_cols.insert( 390 0, 391 offset if isinstance(offset, exp.Identifier) else exp.to_identifier("pos"), 392 ) 393 394 for e, column in zip(exprs, alias_cols): 395 expression.append( 396 "laterals", 397 exp.Lateral( 398 this=_udtf_type(unnest, has_multi_expr)(this=e), 399 view=True, 400 alias=exp.TableAlias(this=alias.this, columns=alias_cols), 401 ), 402 ) 403 404 return expression 405 406 407def explode_projection_to_unnest( 408 index_offset: int = 0, 409 unnest_map: bool = False, 410) -> t.Callable[[exp.Expr], exp.Expr]: 411 """Convert explode/posexplode projections into unnests.""" 412 413 def _explode_projection_to_unnest(expression: exp.Expr) -> exp.Expr: 414 if isinstance(expression, exp.Select): 415 from sqlglot.optimizer.scope import Scope 416 417 taken_select_names = set(expression.named_selects) 418 taken_source_names = {name for name, _ in Scope(expression).references} 419 420 def new_name(names: set[str], name: str) -> str: 421 name = find_new_name(names, name) 422 names.add(name) 423 return name 424 425 arrays: list[exp.Condition] = [] 426 series_alias = new_name(taken_select_names, "pos") 427 series = exp.alias_( 428 exp.Unnest( 429 expressions=[exp.GenerateSeries(start=exp.Literal.number(index_offset))] 430 ), 431 new_name(taken_source_names, "_u"), 432 table=[series_alias], 433 ) 434 435 # we use list here because expression.selects is mutated inside the loop 436 for select in list(expression.selects): 437 explode = select.find(exp.Explode) 438 439 if explode: 440 if ( 441 unnest_map 442 and type(explode) is exp.Explode 443 and explode.this.is_type(exp.DType.MAP) 444 and (select is explode or isinstance(select, exp.Aliases)) 445 ): 446 map_key_alias: t.Any 447 map_value_alias: t.Any 448 if isinstance(select, exp.Aliases): 449 map_key_alias, map_value_alias = select.aliases 450 else: 451 map_key_alias = new_name(taken_select_names, "key") 452 map_value_alias = new_name(taken_select_names, "value") 453 map_unnest_source = new_name(taken_source_names, "_u") 454 455 map_key_select = select.replace( 456 exp.column(map_key_alias, table=map_unnest_source).as_(map_key_alias) 457 ) 458 459 expressions = expression.expressions 460 expressions.insert( 461 expressions.index(map_key_select) + 1, 462 exp.column(map_value_alias, table=map_unnest_source).as_( 463 map_value_alias 464 ), 465 ) 466 expression.set("expressions", expressions) 467 468 unnest = exp.alias_( 469 exp.Unnest(expressions=[explode.this.copy()]), 470 map_unnest_source, 471 table=[map_key_alias, map_value_alias], 472 ) 473 if expression.args.get("from_"): 474 expression.join(unnest, copy=False, join_type="CROSS") 475 else: 476 expression.from_(unnest, copy=False) 477 478 continue 479 480 pos_alias: t.Any = "" 481 explode_alias: t.Any = "" 482 483 if isinstance(select, exp.Alias): 484 explode_alias = select.args["alias"] 485 alias: exp.Expr = select 486 elif isinstance(select, exp.Aliases): 487 pos_alias = select.aliases[0] 488 explode_alias = select.aliases[1] 489 alias = select.replace(exp.alias_(select.this, "", copy=False)) 490 else: 491 alias = select.replace(exp.alias_(select, "")) 492 explode = alias.find(exp.Explode) 493 assert explode 494 495 is_posexplode = isinstance(explode, exp.Posexplode) 496 explode_arg = explode.this 497 498 if isinstance(explode, exp.ExplodeOuter): 499 bracket = explode_arg[0] 500 bracket.set("safe", True) 501 bracket.set("offset", True) 502 explode_arg = exp.func( 503 "IF", 504 exp.func( 505 "ARRAY_SIZE", exp.func("COALESCE", explode_arg, exp.Array()) 506 ).eq(0), 507 exp.array(bracket, copy=False), 508 explode_arg, 509 ) 510 511 # This ensures that we won't use [POS]EXPLODE's argument as a new selection 512 if isinstance(explode_arg, exp.Column): 513 taken_select_names.add(explode_arg.output_name) 514 515 unnest_source_alias = new_name(taken_source_names, "_u") 516 517 if not explode_alias: 518 explode_alias = new_name(taken_select_names, "col") 519 520 if is_posexplode: 521 pos_alias = new_name(taken_select_names, "pos") 522 523 if not pos_alias: 524 pos_alias = new_name(taken_select_names, "pos") 525 526 alias.set("alias", exp.to_identifier(explode_alias)) 527 528 series_table_alias = series.args["alias"].this 529 column = exp.If( 530 this=exp.column(series_alias, table=series_table_alias).eq( 531 exp.column(pos_alias, table=unnest_source_alias) 532 ), 533 true=exp.column(explode_alias, table=unnest_source_alias), 534 ) 535 536 explode.replace(column) 537 538 if is_posexplode: 539 expressions = expression.expressions 540 expressions.insert( 541 expressions.index(alias) + 1, 542 exp.If( 543 this=exp.column(series_alias, table=series_table_alias).eq( 544 exp.column(pos_alias, table=unnest_source_alias) 545 ), 546 true=exp.column(pos_alias, table=unnest_source_alias), 547 ).as_(pos_alias), 548 ) 549 expression.set("expressions", expressions) 550 551 if not arrays: 552 if expression.args.get("from_"): 553 expression.join(series, copy=False, join_type="CROSS") 554 else: 555 expression.from_(series, copy=False) 556 557 size: exp.Condition = exp.ArraySize(this=explode_arg.copy()) 558 arrays.append(size) 559 560 # trino doesn't support left join unnest with on conditions 561 # if it did, this would be much simpler 562 expression.join( 563 exp.alias_( 564 exp.Unnest( 565 expressions=[explode_arg.copy()], 566 offset=exp.to_identifier(pos_alias), 567 ), 568 unnest_source_alias, 569 table=[explode_alias], 570 ), 571 join_type="CROSS", 572 copy=False, 573 ) 574 575 if index_offset != 1: 576 size = size - 1 577 578 expression.where( 579 exp.column(series_alias, table=series_table_alias) 580 .eq(exp.column(pos_alias, table=unnest_source_alias)) 581 .or_( 582 (exp.column(series_alias, table=series_table_alias) > size).and_( 583 exp.column(pos_alias, table=unnest_source_alias).eq(size) 584 ) 585 ), 586 copy=False, 587 ) 588 589 if arrays: 590 end: exp.Condition = exp.Greatest(this=arrays[0], expressions=arrays[1:]) 591 592 if index_offset != 1: 593 end = end - (1 - index_offset) 594 series.expressions[0].set("end", end) 595 596 return expression 597 598 return _explode_projection_to_unnest 599 600 601def add_within_group_for_percentiles(expression: exp.Expr) -> exp.Expr: 602 """Transforms percentiles by adding a WITHIN GROUP clause to them.""" 603 if ( 604 isinstance(expression, exp.PERCENTILES) 605 and not isinstance(expression.parent, exp.WithinGroup) 606 and expression.expression 607 ): 608 column = expression.this.pop() 609 expression.set("this", expression.expression.pop()) 610 order = exp.Order(expressions=[exp.Ordered(this=column)]) 611 expression = exp.WithinGroup(this=expression, expression=order) 612 613 return expression 614 615 616def remove_within_group_for_percentiles(expression: exp.Expr) -> exp.Expr: 617 """Transforms percentiles by getting rid of their corresponding WITHIN GROUP clause.""" 618 if ( 619 isinstance(expression, exp.WithinGroup) 620 and isinstance(expression.this, exp.PERCENTILES) 621 and isinstance(expression.expression, exp.Order) 622 ): 623 quantile = expression.this.this 624 input_value = t.cast(exp.Ordered, expression.find(exp.Ordered)).this 625 return expression.replace(exp.ApproxQuantile(this=input_value, quantile=quantile)) 626 627 return expression 628 629 630def add_recursive_cte_column_names(expression: exp.Expr) -> exp.Expr: 631 """Uses projection output names in recursive CTE definitions to define the CTEs' columns.""" 632 if isinstance(expression, exp.With) and expression.recursive: 633 next_name = name_sequence("_c_") 634 635 for cte in expression.expressions: 636 if not cte.args["alias"].columns: 637 query = cte.this 638 if isinstance(query, exp.SetOperation): 639 query = query.this 640 641 cte.args["alias"].set( 642 "columns", 643 [exp.to_identifier(s.alias_or_name or next_name()) for s in query.selects], 644 ) 645 646 return expression 647 648 649def epoch_cast_to_ts(expression: exp.Expr) -> exp.Expr: 650 """Replace 'epoch' in casts by the equivalent date literal.""" 651 if ( 652 isinstance(expression, (exp.Cast, exp.TryCast)) 653 and expression.name.lower() == "epoch" 654 and expression.to.this in exp.DataType.TEMPORAL_TYPES 655 ): 656 expression.this.replace(exp.Literal.string("1970-01-01 00:00:00")) 657 658 return expression 659 660 661def eliminate_semi_and_anti_joins(expression: exp.Expr) -> exp.Expr: 662 """Convert SEMI and ANTI joins into equivalent forms that use EXIST instead.""" 663 if isinstance(expression, exp.Select): 664 for join in list[exp.Join](expression.args.get("joins") or []): 665 on: exp.Expr | None = join.args.get("on") 666 if on and join.kind in ("SEMI", "ANTI"): 667 subquery = exp.select("1").from_(join.this).where(on) 668 exists: exp.Exists | exp.Not = exp.Exists(this=subquery) 669 if join.kind == "ANTI": 670 exists = exists.not_(copy=False) 671 672 join.pop() 673 expression.where(exists, copy=False) 674 675 return expression 676 677 678def eliminate_full_outer_join(expression: exp.Expr) -> exp.Expr: 679 """ 680 Converts a query with a FULL OUTER join to a union of identical queries that 681 use LEFT/RIGHT OUTER joins instead. This transformation currently only works 682 for queries that have a single FULL OUTER join. 683 """ 684 if isinstance(expression, exp.Select): 685 full_outer_joins: list[tuple[int, exp.Join]] = [ 686 (index, join) 687 for index, join in enumerate[exp.Join](expression.args.get("joins") or []) 688 if join.side == "FULL" 689 ] 690 691 if len(full_outer_joins) == 1: 692 expression_copy = expression.copy() 693 index, full_outer_join = full_outer_joins[0] 694 695 tables = (expression.args["from_"].alias_or_name, full_outer_join.alias_or_name) 696 join_conditions = full_outer_join.args.get("on") or exp.and_( 697 *[ 698 exp.column(col, tables[0]).eq(exp.column(col, tables[1])) 699 for col in t.cast(list[exp.Identifier], full_outer_join.args.get("using")) 700 ] 701 ) 702 703 full_outer_join.set("side", "left") 704 anti_join_clause = ( 705 exp.select("1").from_(expression.args["from_"]).where(join_conditions) 706 ) 707 expression_copy.args["joins"][index].set("side", "right") 708 expression_copy = expression_copy.where(exp.Exists(this=anti_join_clause).not_()) 709 710 union = exp.union(expression, expression_copy, copy=False, distinct=False) 711 for arg in ("with_", "order", "limit", "offset"): 712 value = expression.args.get(arg) 713 if value: 714 expression.set(arg, None) 715 expression_copy.set(arg, None) 716 union.set(arg, value) 717 return union 718 719 return expression 720 721 722def move_ctes_to_top_level(expression: E) -> E: 723 """ 724 Some dialects (e.g. Hive, T-SQL, Spark prior to version 3) only allow CTEs to be 725 defined at the top-level, so for example queries like: 726 727 SELECT * FROM (WITH t(c) AS (SELECT 1) SELECT * FROM t) AS subq 728 729 are invalid in those dialects. This transformation can be used to ensure all CTEs are 730 moved to the top level so that the final SQL code is valid from a syntax standpoint. 731 732 TODO: handle name clashes whilst moving CTEs (it can get quite tricky & costly). 733 """ 734 top_level_with: exp.With | None = expression.args.get("with_") 735 for inner_with in expression.find_all(exp.With): 736 if inner_with.parent is expression: 737 continue 738 739 if not top_level_with: 740 top_level_with = inner_with.pop() 741 expression.set("with_", top_level_with) 742 else: 743 if inner_with.recursive: 744 top_level_with.set("recursive", True) 745 746 parent_cte = inner_with.find_ancestor(exp.CTE) 747 inner_with.pop() 748 749 if parent_cte: 750 i = top_level_with.expressions.index(parent_cte) 751 top_level_with.expressions[i:i] = inner_with.expressions 752 top_level_with.set("expressions", top_level_with.expressions) 753 else: 754 top_level_with.set( 755 "expressions", top_level_with.expressions + inner_with.expressions 756 ) 757 758 return expression 759 760 761def ensure_bools(expression: exp.Expr) -> exp.Expr: 762 """Converts numeric values used in conditions into explicit boolean expressions.""" 763 from sqlglot.optimizer.canonicalize import ensure_bools 764 765 def _ensure_bool(node: exp.Expr) -> None: 766 if ( 767 node.is_number 768 or ( 769 not isinstance(node, exp.SubqueryPredicate) 770 and node.is_type(exp.DType.UNKNOWN, *exp.DataType.NUMERIC_TYPES) 771 ) 772 or (isinstance(node, exp.Column) and not node.type) 773 ): 774 node.replace(node.neq(0)) 775 776 for node in expression.walk(): 777 ensure_bools(node, _ensure_bool) 778 779 return expression 780 781 782def unqualify_columns(expression: exp.Expr) -> exp.Expr: 783 for column in expression.find_all(exp.Column): 784 # We only wanna pop off the table, db, catalog args 785 for part in column.parts[:-1]: 786 part.pop() 787 788 return expression 789 790 791def unqualify_pivot_fields(expression: exp.Expr) -> exp.Expr: 792 """ 793 Some dialects only accept simple column names in a (UN)PIVOT's FOR clause and IN-list 794 (Oracle raises ORA-01748), even though the aggregate itself may stay qualified. 795 796 Example: 797 >>> from sqlglot import parse_one 798 >>> expr = parse_one("SELECT * FROM tbl PIVOT (SUM(tbl.sales) FOR tbl.quarter IN ('Q1', 'Q2'))") 799 >>> print(unqualify_pivot_fields(expr).sql(dialect="spark")) 800 SELECT * FROM tbl PIVOT(SUM(tbl.sales) FOR quarter IN ('Q1', 'Q2')) 801 """ 802 if isinstance(expression, exp.Pivot): 803 expression.set("fields", [unqualify_columns(field) for field in expression.fields]) 804 805 return expression 806 807 808def remove_unique_constraints(expression: exp.Expr) -> exp.Expr: 809 assert isinstance(expression, exp.Create) 810 for constraint in expression.find_all(exp.UniqueColumnConstraint): 811 parent = constraint.parent 812 (parent if isinstance(parent, (exp.ColumnConstraint, exp.Constraint)) else constraint).pop() 813 814 return expression 815 816 817def ctas_with_tmp_tables_to_create_tmp_view( 818 expression: exp.Expr, 819 tmp_storage_provider: t.Callable[[exp.Expr], exp.Expr] = lambda e: e, 820) -> exp.Expr: 821 assert isinstance(expression, exp.Create) 822 properties: exp.Properties | None = expression.args.get("properties") 823 temporary = any( 824 isinstance(prop, exp.TemporaryProperty) 825 for prop in (properties.expressions if properties is not None else []) 826 ) 827 828 # CTAS with temp tables map to CREATE TEMPORARY VIEW 829 if expression.kind == "TABLE" and temporary: 830 if expression.expression: 831 return exp.Create( 832 kind="TEMPORARY VIEW", 833 this=expression.this, 834 expression=expression.expression, 835 ) 836 return tmp_storage_provider(expression) 837 838 return expression 839 840 841def move_schema_columns_to_partitioned_by(expression: exp.Expr) -> exp.Expr: 842 """ 843 In Hive, the PARTITIONED BY property acts as an extension of a table's schema. When the 844 PARTITIONED BY value is an array of column names, they are transformed into a schema. 845 The corresponding columns are removed from the create statement. 846 """ 847 assert isinstance(expression, exp.Create) 848 schema = expression.this 849 is_partitionable = expression.kind in {"TABLE", "VIEW"} 850 851 if isinstance(schema, exp.Schema) and is_partitionable: 852 prop = expression.find(exp.PartitionedByProperty) 853 if prop and prop.this and not isinstance(prop.this, exp.Schema): 854 columns: set[str] = {v.name.upper() for v in prop.this.expressions} 855 schema_exprs: list[exp.Expr] = schema.expressions 856 partitions = [col for col in schema_exprs if col.name.upper() in columns] 857 schema.set("expressions", [e for e in schema_exprs if e not in partitions]) 858 prop.replace(exp.PartitionedByProperty(this=exp.Schema(expressions=partitions))) 859 expression.set("this", schema) 860 861 return expression 862 863 864def move_partitioned_by_to_schema_columns(expression: exp.Expr) -> exp.Expr: 865 """ 866 Spark 3 supports both "HIVEFORMAT" and "DATASOURCE" formats for CREATE TABLE. 867 868 Currently, SQLGlot uses the DATASOURCE format for Spark 3. 869 """ 870 assert isinstance(expression, exp.Create) 871 prop = expression.find(exp.PartitionedByProperty) 872 if ( 873 prop 874 and prop.this 875 and isinstance(prop.this, exp.Schema) 876 and all(isinstance(e, exp.ColumnDef) and e.kind for e in prop.this.expressions) 877 ): 878 prop_this = exp.Tuple( 879 expressions=[exp.to_identifier(e.this) for e in prop.this.expressions] 880 ) 881 schema: exp.Schema = expression.this 882 for e in prop.this.expressions: 883 schema.append("expressions", e) 884 prop.set("this", prop_this) 885 886 return expression 887 888 889def struct_kv_to_alias(expression: exp.Expr) -> exp.Expr: 890 """Converts struct arguments to aliases, e.g. STRUCT(1 AS y).""" 891 if isinstance(expression, exp.Struct): 892 expression.set( 893 "expressions", 894 [ 895 exp.alias_(e.expression, e.this) if isinstance(e, exp.PropertyEQ) else e 896 for e in expression.expressions 897 ], 898 ) 899 900 return expression 901 902 903def eliminate_join_marks(expression: exp.Expr) -> exp.Expr: 904 """https://docs.oracle.com/cd/B19306_01/server.102/b14200/queries006.htm#sthref3178 905 906 1. You cannot specify the (+) operator in a query block that also contains FROM clause join syntax. 907 908 2. The (+) operator can appear only in the WHERE clause or, in the context of left-correlation (that is, when specifying the TABLE clause) in the FROM clause, and can be applied only to a column of a table or view. 909 910 The (+) operator does not produce an outer join if you specify one table in the outer query and the other table in an inner query. 911 912 You cannot use the (+) operator to outer-join a table to itself, although self joins are valid. 913 914 The (+) operator can be applied only to a column, not to an arbitrary expression. However, an arbitrary expression can contain one or more columns marked with the (+) operator. 915 916 A WHERE condition containing the (+) operator cannot be combined with another condition using the OR logical operator. 917 918 A WHERE condition cannot use the IN comparison condition to compare a column marked with the (+) operator with an expression. 919 920 A WHERE condition cannot compare any column marked with the (+) operator with a subquery. 921 922 -- example with WHERE 923 SELECT d.department_name, sum(e.salary) as total_salary 924 FROM departments d, employees e 925 WHERE e.department_id(+) = d.department_id 926 group by department_name 927 928 -- example of left correlation in select 929 SELECT d.department_name, ( 930 SELECT SUM(e.salary) 931 FROM employees e 932 WHERE e.department_id(+) = d.department_id) AS total_salary 933 FROM departments d; 934 935 -- example of left correlation in from 936 SELECT d.department_name, t.total_salary 937 FROM departments d, ( 938 SELECT SUM(e.salary) AS total_salary 939 FROM employees e 940 WHERE e.department_id(+) = d.department_id 941 ) t 942 """ 943 944 from sqlglot.optimizer.scope import traverse_scope 945 from sqlglot.optimizer.normalize import normalize, normalized 946 from collections import defaultdict 947 948 # we go in reverse to check the main query for left correlation 949 for scope in reversed(traverse_scope(expression)): 950 query = scope.expression 951 952 where: exp.Expr | None = query.args.get("where") 953 joins: list[exp.Join] = query.args.get("joins", []) 954 955 if not where or not any(c.args.get("join_mark") for c in where.find_all(exp.Column)): 956 continue 957 958 # knockout: we do not support left correlation (see point 2) 959 assert not scope.is_correlated_subquery, "Correlated queries are not supported" 960 961 # make sure we have AND of ORs to have clear join terms 962 where = normalize(where.this) 963 assert normalized(where), "Cannot normalize JOIN predicates" 964 # dict of {name: list of join AND conditions} 965 joins_ons: defaultdict[str, list[exp.Expr]] = defaultdict(list) 966 for cond in [where] if not isinstance(where, exp.And) else where.flatten(): 967 join_cols = [col for col in cond.find_all(exp.Column) if col.args.get("join_mark")] 968 969 left_join_table = set(col.table for col in join_cols) 970 if not left_join_table: 971 continue 972 973 assert not (len(left_join_table) > 1), ( 974 "Cannot combine JOIN predicates from different tables" 975 ) 976 977 for col in join_cols: 978 col.set("join_mark", False) 979 980 joins_ons[left_join_table.pop()].append(cond) 981 982 old_joins = {join.alias_or_name: join for join in joins} 983 new_joins: dict[str, exp.Join] = {} 984 query_from = query.args["from_"] 985 986 for table, predicates in joins_ons.items(): 987 join_what = old_joins.get(table, query_from).this.copy() 988 new_joins[join_what.alias_or_name] = exp.Join( 989 this=join_what, on=exp.and_(*predicates), kind="LEFT" 990 ) 991 992 for p in predicates: 993 while isinstance(p.parent, exp.Paren): 994 p.parent.replace(p) 995 996 parent = p.parent 997 p.pop() 998 if isinstance(parent, exp.Binary): 999 left = parent.args.get("this") 1000 parent.replace(parent.right if left is None else left) 1001 elif isinstance(parent, exp.Where): 1002 parent.pop() 1003 1004 if query_from.alias_or_name in new_joins: 1005 only_old_joins: set[str] = old_joins.keys() - new_joins.keys() 1006 assert len(only_old_joins) >= 1, ( 1007 "Cannot determine which table to use in the new FROM clause" 1008 ) 1009 1010 new_from_name = list[str](only_old_joins)[0] 1011 query.set("from_", exp.From(this=old_joins[new_from_name].this)) 1012 1013 if new_joins: 1014 for n, j in old_joins.items(): # preserve any other joins 1015 if n not in new_joins and n != query.args["from_"].name: 1016 if not j.kind: 1017 j.set("kind", "CROSS") 1018 new_joins[n] = j 1019 query.set("joins", list(new_joins.values())) 1020 1021 return expression 1022 1023 1024def any_to_exists(expression: exp.Expr) -> exp.Expr: 1025 """ 1026 Transform ANY operator to Spark's EXISTS 1027 1028 For example, 1029 - Postgres: SELECT * FROM tbl WHERE 5 > ANY(tbl.col) 1030 - Spark: SELECT * FROM tbl WHERE EXISTS(tbl.col, x -> x < 5) 1031 1032 Both ANY and EXISTS accept queries but currently only array expressions are supported for this 1033 transformation 1034 """ 1035 if isinstance(expression, exp.Select): 1036 for any_expr in expression.find_all(exp.Any): 1037 this: exp.Expr = any_expr.this 1038 if isinstance(this, exp.Query) or isinstance(any_expr.parent, (exp.Like, exp.ILike)): 1039 continue 1040 1041 binop = any_expr.parent 1042 if isinstance(binop, exp.Binary): 1043 lambda_arg = exp.to_identifier("x") 1044 any_expr.replace(lambda_arg) 1045 lambda_expr = exp.Lambda(this=binop.copy(), expressions=[lambda_arg]) 1046 binop.replace(exp.Exists(this=this.unnest(), expression=lambda_expr)) 1047 1048 return expression 1049 1050 1051def eliminate_window_clause(expression: exp.Expr) -> exp.Expr: 1052 """Eliminates the `WINDOW` query clause by inling each named window.""" 1053 windows: list[exp.Expr] | None = expression.args.get("windows") 1054 if isinstance(expression, exp.Select) and windows is not None: 1055 from sqlglot.optimizer.scope import find_all_in_scope 1056 1057 expression.set("windows", None) 1058 1059 window_expression: dict[str, exp.Expr] = {} 1060 1061 def _inline_inherited_window(window: exp.Expr) -> None: 1062 inherited_window = window_expression.get(window.alias.lower()) 1063 if not inherited_window: 1064 return 1065 1066 window.set("alias", None) 1067 for key in ("partition_by", "order", "spec"): 1068 arg: exp.Expr | None = inherited_window.args.get(key) 1069 if arg is not None: 1070 window.set(key, arg.copy()) 1071 1072 for window in windows: 1073 _inline_inherited_window(window) 1074 window_expression[window.name.lower()] = window 1075 1076 for window in find_all_in_scope(expression, exp.Window): 1077 _inline_inherited_window(window) 1078 1079 return expression 1080 1081 1082def inherit_struct_field_names(expression: exp.Expr) -> exp.Expr: 1083 """ 1084 Inherit field names from the first struct in an array. 1085 1086 BigQuery supports implicitly inheriting names from the first STRUCT in an array: 1087 1088 Example: 1089 ARRAY[ 1090 STRUCT('Alice' AS name, 85 AS score), -- defines names 1091 STRUCT('Bob', 92), -- inherits names 1092 STRUCT('Diana', 95) -- inherits names 1093 ] 1094 1095 This transformation makes the field names explicit on all structs by adding 1096 PropertyEQ nodes, in order to facilitate transpilation to other dialects. 1097 1098 Args: 1099 expression: The expression tree to transform 1100 1101 Returns: 1102 The modified expression with field names inherited in all structs 1103 """ 1104 if ( 1105 isinstance(expression, exp.Array) 1106 and expression.args.get("struct_name_inheritance") 1107 and isinstance(first_item := seq_get(expression.expressions, 0), exp.Struct) 1108 and all(isinstance(fld, exp.PropertyEQ) for fld in first_item.expressions) 1109 ): 1110 field_names: list[exp.Identifier] = [fld.this for fld in first_item.expressions] 1111 1112 # Apply field names to subsequent structs that don't have them 1113 for struct in expression.expressions[1:]: 1114 if not isinstance(struct, exp.Struct) or len(struct.expressions) != len(field_names): 1115 continue 1116 1117 # Convert unnamed expressions to PropertyEQ with inherited names 1118 new_expressions: list[exp.PropertyEQ] = [] 1119 for i, expr in enumerate(struct.expressions): 1120 if not isinstance(expr, exp.PropertyEQ): 1121 # Create PropertyEQ: field_name := value, preserving the type from the inner expression 1122 property_eq = exp.PropertyEQ( 1123 this=field_names[i].copy(), 1124 expression=expr, 1125 ) 1126 property_eq.type = expr.type 1127 new_expressions.append(property_eq) 1128 else: 1129 new_expressions.append(expr) 1130 1131 struct.set("expressions", new_expressions) 1132 1133 return expression
16class SqlHandler(t.Protocol): 17 def __call__(self, expression: exp.Expr, *args: t.Any, **kwargs: t.Any) -> str: ...
Base class for protocol classes.
Protocol classes are defined as::
class Proto(Protocol):
def meth(self) -> int:
...
Such classes are primarily used with static type checkers that recognize structural subtyping (static duck-typing), for example::
class C:
def meth(self) -> int:
return 0
def func(x: Proto) -> int:
return x.meth()
func(C()) # Passes static type check
See PEP 544 for details. Protocol classes decorated with @typing.runtime_checkable act as simple-minded runtime protocols that check only the presence of given attributes, ignoring their type signatures. Protocol classes can be generic, they are defined as::
class GenProto(Protocol[T]):
def meth(self) -> T:
...
1431def _no_init_or_replace_init(self, *args, **kwargs): 1432 cls = type(self) 1433 1434 if cls._is_protocol: 1435 raise TypeError('Protocols cannot be instantiated') 1436 1437 # Already using a custom `__init__`. No need to calculate correct 1438 # `__init__` to call. This can lead to RecursionError. See bpo-45121. 1439 if cls.__init__ is not _no_init_or_replace_init: 1440 return 1441 1442 # Initially, `__init__` of a protocol subclass is set to `_no_init_or_replace_init`. 1443 # The first instantiation of the subclass will call `_no_init_or_replace_init` which 1444 # searches for a proper new `__init__` in the MRO. The new `__init__` 1445 # replaces the subclass' old `__init__` (ie `_no_init_or_replace_init`). Subsequent 1446 # instantiation of the protocol subclass will thus use the new 1447 # `__init__` and no longer call `_no_init_or_replace_init`. 1448 for base in cls.__mro__: 1449 init = base.__dict__.get('__init__', _no_init_or_replace_init) 1450 if init is not _no_init_or_replace_init: 1451 cls.__init__ = init 1452 break 1453 else: 1454 # should not happen 1455 cls.__init__ = object.__init__ 1456 1457 cls.__init__(self, *args, **kwargs)
20def preprocess( 21 transforms: list[t.Callable[[exp.Expr], exp.Expr]], 22 generator: t.Callable[[Generator, exp.Expr], str] | None = None, 23) -> t.Callable[[Generator, exp.Expr], str]: 24 """ 25 Creates a new transform by chaining a sequence of transformations and converts the resulting 26 expression to SQL, using either the "_sql" method corresponding to the resulting expression, 27 or the appropriate `Generator.TRANSFORMS` function (when applicable -- see below). 28 29 Args: 30 transforms: sequence of transform functions. These will be called in order. 31 32 Returns: 33 Function that can be used as a generator transform. 34 """ 35 36 def _to_sql(self: Generator, expression: exp.Expr) -> str: 37 expression_type = type(expression) 38 39 try: 40 expression = transforms[0](expression) 41 for transform in transforms[1:]: 42 expression = transform(expression) 43 except UnsupportedError as unsupported_error: 44 self.unsupported(str(unsupported_error)) 45 46 if generator: 47 return generator(self, expression) 48 49 _sql_handler: SqlHandler | None = getattr(self, expression.key + "_sql", None) 50 if _sql_handler: 51 return _sql_handler(expression) 52 53 transforms_handler = self.TRANSFORMS.get(type(expression)) 54 if transforms_handler: 55 if expression_type is type(expression): 56 if isinstance(expression, exp.Func): 57 return self.function_fallback_sql(expression) 58 59 # Ensures we don't enter an infinite loop. This can happen when the original expression 60 # has the same type as the final expression and there's no _sql method available for it, 61 # because then it'd re-enter _to_sql. 62 raise ValueError( 63 f"Expr type {expression.__class__.__name__} requires a _sql method in order to be transformed." 64 ) 65 66 return transforms_handler(self, expression) 67 68 raise ValueError(f"Unsupported expression type {expression.__class__.__name__}.") 69 70 return _to_sql
Creates a new transform by chaining a sequence of transformations and converts the resulting
expression to SQL, using either the "_sql" method corresponding to the resulting expression,
or the appropriate Generator.TRANSFORMS function (when applicable -- see below).
Arguments:
- transforms: sequence of transform functions. These will be called in order.
Returns:
Function that can be used as a generator transform.
73def unnest_generate_date_array_using_recursive_cte(expression: exp.Expr) -> exp.Expr: 74 if isinstance(expression, exp.Select): 75 count = 0 76 recursive_ctes: list[exp.Expr] = [] 77 78 for unnest in expression.find_all(exp.Unnest): 79 if ( 80 not isinstance(unnest.parent, (exp.From, exp.Join)) 81 or len(unnest.expressions) != 1 82 or not isinstance(unnest.expressions[0], exp.GenerateDateArray) 83 ): 84 continue 85 86 generate_date_array = unnest.expressions[0] 87 start: exp.Expr | None = generate_date_array.args.get("start") 88 end: exp.Expr | None = generate_date_array.args.get("end") 89 step: exp.Expr | None = generate_date_array.args.get("step") 90 91 if not start or not end or not isinstance(step, exp.Interval): 92 continue 93 94 alias: exp.TableAlias | None = unnest.args.get("alias") 95 column_name: str = ( 96 alias.columns[0] if isinstance(alias, exp.TableAlias) else "date_value" 97 ) 98 99 start = exp.cast(start, "date") 100 date_add = exp.func( 101 "date_add", column_name, exp.Literal.number(step.name), step.args.get("unit") 102 ) 103 cast_date_add = exp.cast(date_add, "date") 104 105 cte_name = "_generated_dates" + (f"_{count}" if count else "") 106 107 base_query = exp.select(start.as_(column_name)) 108 recursive_query = ( 109 exp.select(cast_date_add) 110 .from_(cte_name) 111 .where(cast_date_add <= exp.cast(end, "date")) 112 ) 113 cte_query = base_query.union(recursive_query, distinct=False) 114 115 generate_dates_query = exp.select(column_name).from_(cte_name) 116 unnest.replace(generate_dates_query.subquery(cte_name)) 117 118 recursive_ctes.append( 119 exp.alias_(exp.CTE(this=cte_query), cte_name, table=[column_name]) 120 ) 121 count += 1 122 123 if recursive_ctes: 124 with_expression: exp.With = expression.args.get("with_") or exp.With() 125 with_expression.set("recursive", True) 126 with_expression.set("expressions", [*recursive_ctes, *with_expression.expressions]) 127 expression.set("with_", with_expression) 128 129 return expression
132def unnest_generate_series(expression: exp.Expr) -> exp.Expr: 133 """Unnests GENERATE_SERIES or SEQUENCE table references.""" 134 this = expression.this 135 if isinstance(expression, exp.Table) and isinstance(this, exp.GenerateSeries): 136 unnest = exp.Unnest(expressions=[this]) 137 if expression.alias: 138 return exp.alias_(unnest, alias="_u", table=[expression.alias], copy=False) 139 140 return unnest 141 142 return expression
Unnests GENERATE_SERIES or SEQUENCE table references.
145def eliminate_distinct_on(expression: exp.Expr) -> exp.Expr: 146 """ 147 Convert SELECT DISTINCT ON statements to a subquery with a window function. 148 149 This is useful for dialects that don't support SELECT DISTINCT ON but support window functions. 150 151 Args: 152 expression: the expression that will be transformed. 153 154 Returns: 155 The transformed expression. 156 """ 157 if ( 158 isinstance(expression, exp.Select) 159 and expression.args.get("distinct") 160 and isinstance(expression.args["distinct"].args.get("on"), exp.Tuple) 161 ): 162 row_number_window_alias = find_new_name(expression.named_selects, "_row_number") 163 164 distinct_cols = expression.args["distinct"].pop().args["on"].expressions 165 window = exp.Window(this=exp.RowNumber(), partition_by=distinct_cols) 166 167 order: exp.Order | None = expression.args.get("order") 168 if order: 169 window.set("order", order.pop()) 170 else: 171 window.set("order", exp.Order(expressions=[c.copy() for c in distinct_cols])) 172 173 expression.select(exp.alias_(window, row_number_window_alias), copy=False) 174 175 # We add aliases to the projections so that we can safely reference them in the outer query 176 new_selects: list[exp.Expr] = [] 177 taken_names = {row_number_window_alias} 178 for select in expression.selects[:-1]: 179 if select.is_star: 180 new_selects = [exp.Star()] 181 break 182 183 if not isinstance(select, exp.Alias): 184 alias = find_new_name(taken_names, select.output_name or "_col") 185 quoted: bool | None = ( 186 select.this.args.get("quoted") if isinstance(select, exp.Column) else None 187 ) 188 select = select.replace(exp.alias_(select, alias, quoted=quoted)) 189 190 taken_names.add(select.output_name) 191 new_selects.append(select.args["alias"]) 192 193 return ( 194 exp.select(*new_selects, copy=False) 195 .from_(expression.subquery("_t", copy=False), copy=False) 196 .where(exp.column(row_number_window_alias).eq(1), copy=False) 197 ) 198 199 return expression
Convert SELECT DISTINCT ON statements to a subquery with a window function.
This is useful for dialects that don't support SELECT DISTINCT ON but support window functions.
Arguments:
- expression: the expression that will be transformed.
Returns:
The transformed expression.
202def eliminate_qualify(expression: exp.Expr) -> exp.Expr: 203 """ 204 Convert SELECT statements that contain the QUALIFY clause into subqueries, filtered equivalently. 205 206 The idea behind this transformation can be seen in Snowflake's documentation for QUALIFY: 207 https://docs.snowflake.com/en/sql-reference/constructs/qualify 208 209 Some dialects don't support window functions in the WHERE clause, so we need to include them as 210 projections in the subquery, in order to refer to them in the outer filter using aliases. Also, 211 if a column is referenced in the QUALIFY clause but is not selected, we need to include it too, 212 otherwise we won't be able to refer to it in the outer query's WHERE clause. Finally, if a 213 newly aliased projection is referenced in the QUALIFY clause, it will be replaced by the 214 corresponding expression to avoid creating invalid column references. 215 """ 216 if isinstance(expression, exp.Select) and expression.args.get("qualify"): 217 taken = set(expression.named_selects) 218 for select in expression.selects: 219 if not select.alias_or_name: 220 alias = find_new_name(taken, "_c") 221 select.replace(exp.alias_(select, alias)) 222 taken.add(alias) 223 224 def _select_alias_or_name(select: exp.Expr) -> str | exp.Column: 225 alias_or_name = select.alias_or_name 226 identifier = select.args.get("alias") or select.this 227 if isinstance(identifier, exp.Identifier): 228 return exp.column(alias_or_name, quoted=identifier.args.get("quoted")) 229 return alias_or_name 230 231 outer_selects = exp.select(*map(_select_alias_or_name, expression.selects)) 232 qualify_filters: exp.Expr = expression.args["qualify"].pop().this 233 expression_by_alias: dict[str, exp.Expr] = { 234 select.alias: select.this 235 for select in expression.selects 236 if isinstance(select, exp.Alias) 237 } 238 239 select_candidates = (exp.Window,) if expression.is_star else (exp.Window, exp.Column) 240 for select_candidate in list(qualify_filters.find_all(*select_candidates)): 241 if isinstance(select_candidate, exp.Window): 242 if expression_by_alias: 243 for column in select_candidate.find_all(exp.Column): 244 expr = expression_by_alias.get(column.name) 245 if expr: 246 column.replace(expr) 247 248 alias = find_new_name(expression.named_selects, "_w") 249 expression.select(exp.alias_(select_candidate, alias), copy=False) 250 column = exp.column(alias) 251 252 if isinstance(select_candidate.parent, exp.Qualify): 253 qualify_filters = column 254 else: 255 select_candidate.replace(column) 256 elif select_candidate.name not in expression.named_selects and not ( 257 select_candidate.find_ancestor(exp.Window) 258 ): 259 # A column that is only read by a window function is computed by the window 260 # itself, so projecting it in the subquery is redundant and can even produce 261 # invalid SQL, e.g. when the query is grouped 262 expression.select(select_candidate.copy(), copy=False) 263 264 return outer_selects.from_(expression.subquery(alias="_t", copy=False), copy=False).where( 265 qualify_filters, copy=False 266 ) 267 268 return expression
Convert SELECT statements that contain the QUALIFY clause into subqueries, filtered equivalently.
The idea behind this transformation can be seen in Snowflake's documentation for QUALIFY: https://docs.snowflake.com/en/sql-reference/constructs/qualify
Some dialects don't support window functions in the WHERE clause, so we need to include them as projections in the subquery, in order to refer to them in the outer filter using aliases. Also, if a column is referenced in the QUALIFY clause but is not selected, we need to include it too, otherwise we won't be able to refer to it in the outer query's WHERE clause. Finally, if a newly aliased projection is referenced in the QUALIFY clause, it will be replaced by the corresponding expression to avoid creating invalid column references.
271def remove_precision_parameterized_types(expression: exp.Expr) -> exp.Expr: 272 """ 273 Some dialects only allow the precision for parameterized types to be defined in the DDL and not in 274 other expressions. This transforms removes the precision from parameterized types in expressions. 275 """ 276 for node in expression.find_all(exp.DataType): 277 node.set( 278 "expressions", [e for e in node.expressions if not isinstance(e, exp.DataTypeParam)] 279 ) 280 281 return expression
Some dialects only allow the precision for parameterized types to be defined in the DDL and not in other expressions. This transforms removes the precision from parameterized types in expressions.
284def unqualify_unnest(expression: exp.Expr) -> exp.Expr: 285 """Remove references to unnest table aliases, added by the optimizer's qualify_columns step.""" 286 from sqlglot.optimizer.scope import find_all_in_scope 287 288 if isinstance(expression, exp.Select): 289 unnest_aliases = { 290 unnest.alias 291 for unnest in find_all_in_scope(expression, exp.Unnest) 292 if isinstance(unnest.parent, (exp.From, exp.Join)) 293 } 294 if unnest_aliases: 295 for column in expression.find_all(exp.Column): 296 leftmost_part = column.parts[0] 297 if leftmost_part.arg_key != "this" and leftmost_part.this in unnest_aliases: 298 leftmost_part.pop() 299 300 return expression
Remove references to unnest table aliases, added by the optimizer's qualify_columns step.
303def unnest_to_explode( 304 expression: exp.Expr, 305 unnest_using_arrays_zip: bool = True, 306) -> exp.Expr: 307 """Convert cross join unnest into lateral view explode.""" 308 309 def _unnest_zip_exprs( 310 u: exp.Unnest, unnest_exprs: list[exp.Expr], has_multi_expr: bool 311 ) -> list[exp.Expr]: 312 if has_multi_expr: 313 if not unnest_using_arrays_zip: 314 raise UnsupportedError("Cannot transpile UNNEST with multiple input arrays") 315 316 # Use INLINE(ARRAYS_ZIP(...)) for multiple expressions 317 zip_exprs: list[exp.Expr] = [exp.Anonymous(this="ARRAYS_ZIP", expressions=unnest_exprs)] 318 u.set("expressions", zip_exprs) 319 return zip_exprs 320 return unnest_exprs 321 322 def _udtf_type(u: exp.Unnest, has_multi_expr: bool) -> type[exp.Func]: 323 if u.args.get("offset"): 324 return exp.Posexplode 325 return exp.Inline if has_multi_expr else exp.Explode 326 327 if isinstance(expression, exp.Select): 328 from_ = expression.args.get("from_") 329 330 if from_ and isinstance(from_.this, exp.Unnest): 331 unnest: exp.Unnest = from_.this 332 alias: exp.TableAlias | None = unnest.args.get("alias") 333 exprs: list[exp.Expr] = unnest.expressions 334 has_multi_expr = len(exprs) > 1 335 this, *_ = _unnest_zip_exprs(unnest, exprs, has_multi_expr) 336 337 columns: list[exp.Identifier] = alias.columns if alias else [] 338 offset: exp.Expr | None = unnest.args.get("offset") 339 if offset: 340 columns.insert( 341 0, offset if isinstance(offset, exp.Identifier) else exp.to_identifier("pos") 342 ) 343 344 unnest.replace( 345 exp.Table( 346 this=_udtf_type(unnest, has_multi_expr)(this=this), 347 alias=exp.TableAlias(this=alias.this, columns=columns) if alias else None, 348 ) 349 ) 350 351 joins: list[exp.Join] = expression.args.get("joins") or [] 352 for join in list(joins): 353 join_expr = join.this 354 355 is_lateral = isinstance(join_expr, exp.Lateral) 356 357 unnest = join_expr.this if is_lateral else join_expr 358 359 if isinstance(unnest, exp.Unnest): 360 if is_lateral: 361 alias = join_expr.args.get("alias") 362 else: 363 alias = unnest.args.get("alias") 364 365 if alias is None: 366 raise UnsupportedError( 367 "CROSS JOIN UNNEST to LATERAL VIEW EXPLODE transformation requires an alias" 368 ) 369 370 exprs = unnest.expressions 371 # The number of unnest.expressions will be changed by _unnest_zip_exprs, we need to record it here 372 has_multi_expr = len(exprs) > 1 373 exprs = _unnest_zip_exprs(unnest, exprs, has_multi_expr) 374 375 joins.remove(join) 376 377 alias_cols: list[exp.Identifier] = alias.columns 378 379 # # Handle UNNEST to LATERAL VIEW EXPLODE: Exception is raised when there are 0 or > 2 aliases 380 # Spark LATERAL VIEW EXPLODE requires single alias for array/struct and two for Map type column unlike unnest in trino/presto which can take an arbitrary amount. 381 # Refs: https://spark.apache.org/docs/latest/sql-ref-syntax-qry-select-lateral-view.html 382 383 if not has_multi_expr and len(alias_cols) not in (1, 2): 384 raise UnsupportedError( 385 "CROSS JOIN UNNEST to LATERAL VIEW EXPLODE transformation requires explicit column aliases" 386 ) 387 388 offset = unnest.args.get("offset") 389 if offset: 390 alias_cols.insert( 391 0, 392 offset if isinstance(offset, exp.Identifier) else exp.to_identifier("pos"), 393 ) 394 395 for e, column in zip(exprs, alias_cols): 396 expression.append( 397 "laterals", 398 exp.Lateral( 399 this=_udtf_type(unnest, has_multi_expr)(this=e), 400 view=True, 401 alias=exp.TableAlias(this=alias.this, columns=alias_cols), 402 ), 403 ) 404 405 return expression
Convert cross join unnest into lateral view explode.
408def explode_projection_to_unnest( 409 index_offset: int = 0, 410 unnest_map: bool = False, 411) -> t.Callable[[exp.Expr], exp.Expr]: 412 """Convert explode/posexplode projections into unnests.""" 413 414 def _explode_projection_to_unnest(expression: exp.Expr) -> exp.Expr: 415 if isinstance(expression, exp.Select): 416 from sqlglot.optimizer.scope import Scope 417 418 taken_select_names = set(expression.named_selects) 419 taken_source_names = {name for name, _ in Scope(expression).references} 420 421 def new_name(names: set[str], name: str) -> str: 422 name = find_new_name(names, name) 423 names.add(name) 424 return name 425 426 arrays: list[exp.Condition] = [] 427 series_alias = new_name(taken_select_names, "pos") 428 series = exp.alias_( 429 exp.Unnest( 430 expressions=[exp.GenerateSeries(start=exp.Literal.number(index_offset))] 431 ), 432 new_name(taken_source_names, "_u"), 433 table=[series_alias], 434 ) 435 436 # we use list here because expression.selects is mutated inside the loop 437 for select in list(expression.selects): 438 explode = select.find(exp.Explode) 439 440 if explode: 441 if ( 442 unnest_map 443 and type(explode) is exp.Explode 444 and explode.this.is_type(exp.DType.MAP) 445 and (select is explode or isinstance(select, exp.Aliases)) 446 ): 447 map_key_alias: t.Any 448 map_value_alias: t.Any 449 if isinstance(select, exp.Aliases): 450 map_key_alias, map_value_alias = select.aliases 451 else: 452 map_key_alias = new_name(taken_select_names, "key") 453 map_value_alias = new_name(taken_select_names, "value") 454 map_unnest_source = new_name(taken_source_names, "_u") 455 456 map_key_select = select.replace( 457 exp.column(map_key_alias, table=map_unnest_source).as_(map_key_alias) 458 ) 459 460 expressions = expression.expressions 461 expressions.insert( 462 expressions.index(map_key_select) + 1, 463 exp.column(map_value_alias, table=map_unnest_source).as_( 464 map_value_alias 465 ), 466 ) 467 expression.set("expressions", expressions) 468 469 unnest = exp.alias_( 470 exp.Unnest(expressions=[explode.this.copy()]), 471 map_unnest_source, 472 table=[map_key_alias, map_value_alias], 473 ) 474 if expression.args.get("from_"): 475 expression.join(unnest, copy=False, join_type="CROSS") 476 else: 477 expression.from_(unnest, copy=False) 478 479 continue 480 481 pos_alias: t.Any = "" 482 explode_alias: t.Any = "" 483 484 if isinstance(select, exp.Alias): 485 explode_alias = select.args["alias"] 486 alias: exp.Expr = select 487 elif isinstance(select, exp.Aliases): 488 pos_alias = select.aliases[0] 489 explode_alias = select.aliases[1] 490 alias = select.replace(exp.alias_(select.this, "", copy=False)) 491 else: 492 alias = select.replace(exp.alias_(select, "")) 493 explode = alias.find(exp.Explode) 494 assert explode 495 496 is_posexplode = isinstance(explode, exp.Posexplode) 497 explode_arg = explode.this 498 499 if isinstance(explode, exp.ExplodeOuter): 500 bracket = explode_arg[0] 501 bracket.set("safe", True) 502 bracket.set("offset", True) 503 explode_arg = exp.func( 504 "IF", 505 exp.func( 506 "ARRAY_SIZE", exp.func("COALESCE", explode_arg, exp.Array()) 507 ).eq(0), 508 exp.array(bracket, copy=False), 509 explode_arg, 510 ) 511 512 # This ensures that we won't use [POS]EXPLODE's argument as a new selection 513 if isinstance(explode_arg, exp.Column): 514 taken_select_names.add(explode_arg.output_name) 515 516 unnest_source_alias = new_name(taken_source_names, "_u") 517 518 if not explode_alias: 519 explode_alias = new_name(taken_select_names, "col") 520 521 if is_posexplode: 522 pos_alias = new_name(taken_select_names, "pos") 523 524 if not pos_alias: 525 pos_alias = new_name(taken_select_names, "pos") 526 527 alias.set("alias", exp.to_identifier(explode_alias)) 528 529 series_table_alias = series.args["alias"].this 530 column = exp.If( 531 this=exp.column(series_alias, table=series_table_alias).eq( 532 exp.column(pos_alias, table=unnest_source_alias) 533 ), 534 true=exp.column(explode_alias, table=unnest_source_alias), 535 ) 536 537 explode.replace(column) 538 539 if is_posexplode: 540 expressions = expression.expressions 541 expressions.insert( 542 expressions.index(alias) + 1, 543 exp.If( 544 this=exp.column(series_alias, table=series_table_alias).eq( 545 exp.column(pos_alias, table=unnest_source_alias) 546 ), 547 true=exp.column(pos_alias, table=unnest_source_alias), 548 ).as_(pos_alias), 549 ) 550 expression.set("expressions", expressions) 551 552 if not arrays: 553 if expression.args.get("from_"): 554 expression.join(series, copy=False, join_type="CROSS") 555 else: 556 expression.from_(series, copy=False) 557 558 size: exp.Condition = exp.ArraySize(this=explode_arg.copy()) 559 arrays.append(size) 560 561 # trino doesn't support left join unnest with on conditions 562 # if it did, this would be much simpler 563 expression.join( 564 exp.alias_( 565 exp.Unnest( 566 expressions=[explode_arg.copy()], 567 offset=exp.to_identifier(pos_alias), 568 ), 569 unnest_source_alias, 570 table=[explode_alias], 571 ), 572 join_type="CROSS", 573 copy=False, 574 ) 575 576 if index_offset != 1: 577 size = size - 1 578 579 expression.where( 580 exp.column(series_alias, table=series_table_alias) 581 .eq(exp.column(pos_alias, table=unnest_source_alias)) 582 .or_( 583 (exp.column(series_alias, table=series_table_alias) > size).and_( 584 exp.column(pos_alias, table=unnest_source_alias).eq(size) 585 ) 586 ), 587 copy=False, 588 ) 589 590 if arrays: 591 end: exp.Condition = exp.Greatest(this=arrays[0], expressions=arrays[1:]) 592 593 if index_offset != 1: 594 end = end - (1 - index_offset) 595 series.expressions[0].set("end", end) 596 597 return expression 598 599 return _explode_projection_to_unnest
Convert explode/posexplode projections into unnests.
602def add_within_group_for_percentiles(expression: exp.Expr) -> exp.Expr: 603 """Transforms percentiles by adding a WITHIN GROUP clause to them.""" 604 if ( 605 isinstance(expression, exp.PERCENTILES) 606 and not isinstance(expression.parent, exp.WithinGroup) 607 and expression.expression 608 ): 609 column = expression.this.pop() 610 expression.set("this", expression.expression.pop()) 611 order = exp.Order(expressions=[exp.Ordered(this=column)]) 612 expression = exp.WithinGroup(this=expression, expression=order) 613 614 return expression
Transforms percentiles by adding a WITHIN GROUP clause to them.
617def remove_within_group_for_percentiles(expression: exp.Expr) -> exp.Expr: 618 """Transforms percentiles by getting rid of their corresponding WITHIN GROUP clause.""" 619 if ( 620 isinstance(expression, exp.WithinGroup) 621 and isinstance(expression.this, exp.PERCENTILES) 622 and isinstance(expression.expression, exp.Order) 623 ): 624 quantile = expression.this.this 625 input_value = t.cast(exp.Ordered, expression.find(exp.Ordered)).this 626 return expression.replace(exp.ApproxQuantile(this=input_value, quantile=quantile)) 627 628 return expression
Transforms percentiles by getting rid of their corresponding WITHIN GROUP clause.
631def add_recursive_cte_column_names(expression: exp.Expr) -> exp.Expr: 632 """Uses projection output names in recursive CTE definitions to define the CTEs' columns.""" 633 if isinstance(expression, exp.With) and expression.recursive: 634 next_name = name_sequence("_c_") 635 636 for cte in expression.expressions: 637 if not cte.args["alias"].columns: 638 query = cte.this 639 if isinstance(query, exp.SetOperation): 640 query = query.this 641 642 cte.args["alias"].set( 643 "columns", 644 [exp.to_identifier(s.alias_or_name or next_name()) for s in query.selects], 645 ) 646 647 return expression
Uses projection output names in recursive CTE definitions to define the CTEs' columns.
650def epoch_cast_to_ts(expression: exp.Expr) -> exp.Expr: 651 """Replace 'epoch' in casts by the equivalent date literal.""" 652 if ( 653 isinstance(expression, (exp.Cast, exp.TryCast)) 654 and expression.name.lower() == "epoch" 655 and expression.to.this in exp.DataType.TEMPORAL_TYPES 656 ): 657 expression.this.replace(exp.Literal.string("1970-01-01 00:00:00")) 658 659 return expression
Replace 'epoch' in casts by the equivalent date literal.
662def eliminate_semi_and_anti_joins(expression: exp.Expr) -> exp.Expr: 663 """Convert SEMI and ANTI joins into equivalent forms that use EXIST instead.""" 664 if isinstance(expression, exp.Select): 665 for join in list[exp.Join](expression.args.get("joins") or []): 666 on: exp.Expr | None = join.args.get("on") 667 if on and join.kind in ("SEMI", "ANTI"): 668 subquery = exp.select("1").from_(join.this).where(on) 669 exists: exp.Exists | exp.Not = exp.Exists(this=subquery) 670 if join.kind == "ANTI": 671 exists = exists.not_(copy=False) 672 673 join.pop() 674 expression.where(exists, copy=False) 675 676 return expression
Convert SEMI and ANTI joins into equivalent forms that use EXIST instead.
679def eliminate_full_outer_join(expression: exp.Expr) -> exp.Expr: 680 """ 681 Converts a query with a FULL OUTER join to a union of identical queries that 682 use LEFT/RIGHT OUTER joins instead. This transformation currently only works 683 for queries that have a single FULL OUTER join. 684 """ 685 if isinstance(expression, exp.Select): 686 full_outer_joins: list[tuple[int, exp.Join]] = [ 687 (index, join) 688 for index, join in enumerate[exp.Join](expression.args.get("joins") or []) 689 if join.side == "FULL" 690 ] 691 692 if len(full_outer_joins) == 1: 693 expression_copy = expression.copy() 694 index, full_outer_join = full_outer_joins[0] 695 696 tables = (expression.args["from_"].alias_or_name, full_outer_join.alias_or_name) 697 join_conditions = full_outer_join.args.get("on") or exp.and_( 698 *[ 699 exp.column(col, tables[0]).eq(exp.column(col, tables[1])) 700 for col in t.cast(list[exp.Identifier], full_outer_join.args.get("using")) 701 ] 702 ) 703 704 full_outer_join.set("side", "left") 705 anti_join_clause = ( 706 exp.select("1").from_(expression.args["from_"]).where(join_conditions) 707 ) 708 expression_copy.args["joins"][index].set("side", "right") 709 expression_copy = expression_copy.where(exp.Exists(this=anti_join_clause).not_()) 710 711 union = exp.union(expression, expression_copy, copy=False, distinct=False) 712 for arg in ("with_", "order", "limit", "offset"): 713 value = expression.args.get(arg) 714 if value: 715 expression.set(arg, None) 716 expression_copy.set(arg, None) 717 union.set(arg, value) 718 return union 719 720 return expression
Converts a query with a FULL OUTER join to a union of identical queries that use LEFT/RIGHT OUTER joins instead. This transformation currently only works for queries that have a single FULL OUTER join.
723def move_ctes_to_top_level(expression: E) -> E: 724 """ 725 Some dialects (e.g. Hive, T-SQL, Spark prior to version 3) only allow CTEs to be 726 defined at the top-level, so for example queries like: 727 728 SELECT * FROM (WITH t(c) AS (SELECT 1) SELECT * FROM t) AS subq 729 730 are invalid in those dialects. This transformation can be used to ensure all CTEs are 731 moved to the top level so that the final SQL code is valid from a syntax standpoint. 732 733 TODO: handle name clashes whilst moving CTEs (it can get quite tricky & costly). 734 """ 735 top_level_with: exp.With | None = expression.args.get("with_") 736 for inner_with in expression.find_all(exp.With): 737 if inner_with.parent is expression: 738 continue 739 740 if not top_level_with: 741 top_level_with = inner_with.pop() 742 expression.set("with_", top_level_with) 743 else: 744 if inner_with.recursive: 745 top_level_with.set("recursive", True) 746 747 parent_cte = inner_with.find_ancestor(exp.CTE) 748 inner_with.pop() 749 750 if parent_cte: 751 i = top_level_with.expressions.index(parent_cte) 752 top_level_with.expressions[i:i] = inner_with.expressions 753 top_level_with.set("expressions", top_level_with.expressions) 754 else: 755 top_level_with.set( 756 "expressions", top_level_with.expressions + inner_with.expressions 757 ) 758 759 return expression
Some dialects (e.g. Hive, T-SQL, Spark prior to version 3) only allow CTEs to be defined at the top-level, so for example queries like:
SELECT * FROM (WITH t(c) AS (SELECT 1) SELECT * FROM t) AS subq
are invalid in those dialects. This transformation can be used to ensure all CTEs are moved to the top level so that the final SQL code is valid from a syntax standpoint.
TODO: handle name clashes whilst moving CTEs (it can get quite tricky & costly).
762def ensure_bools(expression: exp.Expr) -> exp.Expr: 763 """Converts numeric values used in conditions into explicit boolean expressions.""" 764 from sqlglot.optimizer.canonicalize import ensure_bools 765 766 def _ensure_bool(node: exp.Expr) -> None: 767 if ( 768 node.is_number 769 or ( 770 not isinstance(node, exp.SubqueryPredicate) 771 and node.is_type(exp.DType.UNKNOWN, *exp.DataType.NUMERIC_TYPES) 772 ) 773 or (isinstance(node, exp.Column) and not node.type) 774 ): 775 node.replace(node.neq(0)) 776 777 for node in expression.walk(): 778 ensure_bools(node, _ensure_bool) 779 780 return expression
Converts numeric values used in conditions into explicit boolean expressions.
792def unqualify_pivot_fields(expression: exp.Expr) -> exp.Expr: 793 """ 794 Some dialects only accept simple column names in a (UN)PIVOT's FOR clause and IN-list 795 (Oracle raises ORA-01748), even though the aggregate itself may stay qualified. 796 797 Example: 798 >>> from sqlglot import parse_one 799 >>> expr = parse_one("SELECT * FROM tbl PIVOT (SUM(tbl.sales) FOR tbl.quarter IN ('Q1', 'Q2'))") 800 >>> print(unqualify_pivot_fields(expr).sql(dialect="spark")) 801 SELECT * FROM tbl PIVOT(SUM(tbl.sales) FOR quarter IN ('Q1', 'Q2')) 802 """ 803 if isinstance(expression, exp.Pivot): 804 expression.set("fields", [unqualify_columns(field) for field in expression.fields]) 805 806 return expression
Some dialects only accept simple column names in a (UN)PIVOT's FOR clause and IN-list (Oracle raises ORA-01748), even though the aggregate itself may stay qualified.
Example:
>>> from sqlglot import parse_one >>> expr = parse_one("SELECT * FROM tbl PIVOT (SUM(tbl.sales) FOR tbl.quarter IN ('Q1', 'Q2'))") >>> print(unqualify_pivot_fields(expr).sql(dialect="spark")) SELECT * FROM tbl PIVOT(SUM(tbl.sales) FOR quarter IN ('Q1', 'Q2'))
809def remove_unique_constraints(expression: exp.Expr) -> exp.Expr: 810 assert isinstance(expression, exp.Create) 811 for constraint in expression.find_all(exp.UniqueColumnConstraint): 812 parent = constraint.parent 813 (parent if isinstance(parent, (exp.ColumnConstraint, exp.Constraint)) else constraint).pop() 814 815 return expression
818def ctas_with_tmp_tables_to_create_tmp_view( 819 expression: exp.Expr, 820 tmp_storage_provider: t.Callable[[exp.Expr], exp.Expr] = lambda e: e, 821) -> exp.Expr: 822 assert isinstance(expression, exp.Create) 823 properties: exp.Properties | None = expression.args.get("properties") 824 temporary = any( 825 isinstance(prop, exp.TemporaryProperty) 826 for prop in (properties.expressions if properties is not None else []) 827 ) 828 829 # CTAS with temp tables map to CREATE TEMPORARY VIEW 830 if expression.kind == "TABLE" and temporary: 831 if expression.expression: 832 return exp.Create( 833 kind="TEMPORARY VIEW", 834 this=expression.this, 835 expression=expression.expression, 836 ) 837 return tmp_storage_provider(expression) 838 839 return expression
842def move_schema_columns_to_partitioned_by(expression: exp.Expr) -> exp.Expr: 843 """ 844 In Hive, the PARTITIONED BY property acts as an extension of a table's schema. When the 845 PARTITIONED BY value is an array of column names, they are transformed into a schema. 846 The corresponding columns are removed from the create statement. 847 """ 848 assert isinstance(expression, exp.Create) 849 schema = expression.this 850 is_partitionable = expression.kind in {"TABLE", "VIEW"} 851 852 if isinstance(schema, exp.Schema) and is_partitionable: 853 prop = expression.find(exp.PartitionedByProperty) 854 if prop and prop.this and not isinstance(prop.this, exp.Schema): 855 columns: set[str] = {v.name.upper() for v in prop.this.expressions} 856 schema_exprs: list[exp.Expr] = schema.expressions 857 partitions = [col for col in schema_exprs if col.name.upper() in columns] 858 schema.set("expressions", [e for e in schema_exprs if e not in partitions]) 859 prop.replace(exp.PartitionedByProperty(this=exp.Schema(expressions=partitions))) 860 expression.set("this", schema) 861 862 return expression
In Hive, the PARTITIONED BY property acts as an extension of a table's schema. When the PARTITIONED BY value is an array of column names, they are transformed into a schema. The corresponding columns are removed from the create statement.
865def move_partitioned_by_to_schema_columns(expression: exp.Expr) -> exp.Expr: 866 """ 867 Spark 3 supports both "HIVEFORMAT" and "DATASOURCE" formats for CREATE TABLE. 868 869 Currently, SQLGlot uses the DATASOURCE format for Spark 3. 870 """ 871 assert isinstance(expression, exp.Create) 872 prop = expression.find(exp.PartitionedByProperty) 873 if ( 874 prop 875 and prop.this 876 and isinstance(prop.this, exp.Schema) 877 and all(isinstance(e, exp.ColumnDef) and e.kind for e in prop.this.expressions) 878 ): 879 prop_this = exp.Tuple( 880 expressions=[exp.to_identifier(e.this) for e in prop.this.expressions] 881 ) 882 schema: exp.Schema = expression.this 883 for e in prop.this.expressions: 884 schema.append("expressions", e) 885 prop.set("this", prop_this) 886 887 return expression
Spark 3 supports both "HIVEFORMAT" and "DATASOURCE" formats for CREATE TABLE.
Currently, SQLGlot uses the DATASOURCE format for Spark 3.
890def struct_kv_to_alias(expression: exp.Expr) -> exp.Expr: 891 """Converts struct arguments to aliases, e.g. STRUCT(1 AS y).""" 892 if isinstance(expression, exp.Struct): 893 expression.set( 894 "expressions", 895 [ 896 exp.alias_(e.expression, e.this) if isinstance(e, exp.PropertyEQ) else e 897 for e in expression.expressions 898 ], 899 ) 900 901 return expression
Converts struct arguments to aliases, e.g. STRUCT(1 AS y).
904def eliminate_join_marks(expression: exp.Expr) -> exp.Expr: 905 """https://docs.oracle.com/cd/B19306_01/server.102/b14200/queries006.htm#sthref3178 906 907 1. You cannot specify the (+) operator in a query block that also contains FROM clause join syntax. 908 909 2. The (+) operator can appear only in the WHERE clause or, in the context of left-correlation (that is, when specifying the TABLE clause) in the FROM clause, and can be applied only to a column of a table or view. 910 911 The (+) operator does not produce an outer join if you specify one table in the outer query and the other table in an inner query. 912 913 You cannot use the (+) operator to outer-join a table to itself, although self joins are valid. 914 915 The (+) operator can be applied only to a column, not to an arbitrary expression. However, an arbitrary expression can contain one or more columns marked with the (+) operator. 916 917 A WHERE condition containing the (+) operator cannot be combined with another condition using the OR logical operator. 918 919 A WHERE condition cannot use the IN comparison condition to compare a column marked with the (+) operator with an expression. 920 921 A WHERE condition cannot compare any column marked with the (+) operator with a subquery. 922 923 -- example with WHERE 924 SELECT d.department_name, sum(e.salary) as total_salary 925 FROM departments d, employees e 926 WHERE e.department_id(+) = d.department_id 927 group by department_name 928 929 -- example of left correlation in select 930 SELECT d.department_name, ( 931 SELECT SUM(e.salary) 932 FROM employees e 933 WHERE e.department_id(+) = d.department_id) AS total_salary 934 FROM departments d; 935 936 -- example of left correlation in from 937 SELECT d.department_name, t.total_salary 938 FROM departments d, ( 939 SELECT SUM(e.salary) AS total_salary 940 FROM employees e 941 WHERE e.department_id(+) = d.department_id 942 ) t 943 """ 944 945 from sqlglot.optimizer.scope import traverse_scope 946 from sqlglot.optimizer.normalize import normalize, normalized 947 from collections import defaultdict 948 949 # we go in reverse to check the main query for left correlation 950 for scope in reversed(traverse_scope(expression)): 951 query = scope.expression 952 953 where: exp.Expr | None = query.args.get("where") 954 joins: list[exp.Join] = query.args.get("joins", []) 955 956 if not where or not any(c.args.get("join_mark") for c in where.find_all(exp.Column)): 957 continue 958 959 # knockout: we do not support left correlation (see point 2) 960 assert not scope.is_correlated_subquery, "Correlated queries are not supported" 961 962 # make sure we have AND of ORs to have clear join terms 963 where = normalize(where.this) 964 assert normalized(where), "Cannot normalize JOIN predicates" 965 # dict of {name: list of join AND conditions} 966 joins_ons: defaultdict[str, list[exp.Expr]] = defaultdict(list) 967 for cond in [where] if not isinstance(where, exp.And) else where.flatten(): 968 join_cols = [col for col in cond.find_all(exp.Column) if col.args.get("join_mark")] 969 970 left_join_table = set(col.table for col in join_cols) 971 if not left_join_table: 972 continue 973 974 assert not (len(left_join_table) > 1), ( 975 "Cannot combine JOIN predicates from different tables" 976 ) 977 978 for col in join_cols: 979 col.set("join_mark", False) 980 981 joins_ons[left_join_table.pop()].append(cond) 982 983 old_joins = {join.alias_or_name: join for join in joins} 984 new_joins: dict[str, exp.Join] = {} 985 query_from = query.args["from_"] 986 987 for table, predicates in joins_ons.items(): 988 join_what = old_joins.get(table, query_from).this.copy() 989 new_joins[join_what.alias_or_name] = exp.Join( 990 this=join_what, on=exp.and_(*predicates), kind="LEFT" 991 ) 992 993 for p in predicates: 994 while isinstance(p.parent, exp.Paren): 995 p.parent.replace(p) 996 997 parent = p.parent 998 p.pop() 999 if isinstance(parent, exp.Binary): 1000 left = parent.args.get("this") 1001 parent.replace(parent.right if left is None else left) 1002 elif isinstance(parent, exp.Where): 1003 parent.pop() 1004 1005 if query_from.alias_or_name in new_joins: 1006 only_old_joins: set[str] = old_joins.keys() - new_joins.keys() 1007 assert len(only_old_joins) >= 1, ( 1008 "Cannot determine which table to use in the new FROM clause" 1009 ) 1010 1011 new_from_name = list[str](only_old_joins)[0] 1012 query.set("from_", exp.From(this=old_joins[new_from_name].this)) 1013 1014 if new_joins: 1015 for n, j in old_joins.items(): # preserve any other joins 1016 if n not in new_joins and n != query.args["from_"].name: 1017 if not j.kind: 1018 j.set("kind", "CROSS") 1019 new_joins[n] = j 1020 query.set("joins", list(new_joins.values())) 1021 1022 return expression
https://docs.oracle.com/cd/B19306_01/server.102/b14200/queries006.htm#sthref3178
You cannot specify the (+) operator in a query block that also contains FROM clause join syntax.
The (+) operator can appear only in the WHERE clause or, in the context of left-correlation (that is, when specifying the TABLE clause) in the FROM clause, and can be applied only to a column of a table or view.
The (+) operator does not produce an outer join if you specify one table in the outer query and the other table in an inner query.
You cannot use the (+) operator to outer-join a table to itself, although self joins are valid.
The (+) operator can be applied only to a column, not to an arbitrary expression. However, an arbitrary expression can contain one or more columns marked with the (+) operator.
A WHERE condition containing the (+) operator cannot be combined with another condition using the OR logical operator.
A WHERE condition cannot use the IN comparison condition to compare a column marked with the (+) operator with an expression.
A WHERE condition cannot compare any column marked with the (+) operator with a subquery.
-- example with WHERE SELECT d.department_name, sum(e.salary) as total_salary FROM departments d, employees e WHERE e.department_id(+) = d.department_id group by department_name
-- example of left correlation in select SELECT d.department_name, ( SELECT SUM(e.salary) FROM employees e WHERE e.department_id(+) = d.department_id) AS total_salary FROM departments d;
-- example of left correlation in from SELECT d.department_name, t.total_salary FROM departments d, ( SELECT SUM(e.salary) AS total_salary FROM employees e WHERE e.department_id(+) = d.department_id ) t
1025def any_to_exists(expression: exp.Expr) -> exp.Expr: 1026 """ 1027 Transform ANY operator to Spark's EXISTS 1028 1029 For example, 1030 - Postgres: SELECT * FROM tbl WHERE 5 > ANY(tbl.col) 1031 - Spark: SELECT * FROM tbl WHERE EXISTS(tbl.col, x -> x < 5) 1032 1033 Both ANY and EXISTS accept queries but currently only array expressions are supported for this 1034 transformation 1035 """ 1036 if isinstance(expression, exp.Select): 1037 for any_expr in expression.find_all(exp.Any): 1038 this: exp.Expr = any_expr.this 1039 if isinstance(this, exp.Query) or isinstance(any_expr.parent, (exp.Like, exp.ILike)): 1040 continue 1041 1042 binop = any_expr.parent 1043 if isinstance(binop, exp.Binary): 1044 lambda_arg = exp.to_identifier("x") 1045 any_expr.replace(lambda_arg) 1046 lambda_expr = exp.Lambda(this=binop.copy(), expressions=[lambda_arg]) 1047 binop.replace(exp.Exists(this=this.unnest(), expression=lambda_expr)) 1048 1049 return expression
Transform ANY operator to Spark's EXISTS
For example, - Postgres: SELECT * FROM tbl WHERE 5 > ANY(tbl.col) - Spark: SELECT * FROM tbl WHERE EXISTS(tbl.col, x -> x < 5)
Both ANY and EXISTS accept queries but currently only array expressions are supported for this transformation
1052def eliminate_window_clause(expression: exp.Expr) -> exp.Expr: 1053 """Eliminates the `WINDOW` query clause by inling each named window.""" 1054 windows: list[exp.Expr] | None = expression.args.get("windows") 1055 if isinstance(expression, exp.Select) and windows is not None: 1056 from sqlglot.optimizer.scope import find_all_in_scope 1057 1058 expression.set("windows", None) 1059 1060 window_expression: dict[str, exp.Expr] = {} 1061 1062 def _inline_inherited_window(window: exp.Expr) -> None: 1063 inherited_window = window_expression.get(window.alias.lower()) 1064 if not inherited_window: 1065 return 1066 1067 window.set("alias", None) 1068 for key in ("partition_by", "order", "spec"): 1069 arg: exp.Expr | None = inherited_window.args.get(key) 1070 if arg is not None: 1071 window.set(key, arg.copy()) 1072 1073 for window in windows: 1074 _inline_inherited_window(window) 1075 window_expression[window.name.lower()] = window 1076 1077 for window in find_all_in_scope(expression, exp.Window): 1078 _inline_inherited_window(window) 1079 1080 return expression
Eliminates the WINDOW query clause by inling each named window.
1083def inherit_struct_field_names(expression: exp.Expr) -> exp.Expr: 1084 """ 1085 Inherit field names from the first struct in an array. 1086 1087 BigQuery supports implicitly inheriting names from the first STRUCT in an array: 1088 1089 Example: 1090 ARRAY[ 1091 STRUCT('Alice' AS name, 85 AS score), -- defines names 1092 STRUCT('Bob', 92), -- inherits names 1093 STRUCT('Diana', 95) -- inherits names 1094 ] 1095 1096 This transformation makes the field names explicit on all structs by adding 1097 PropertyEQ nodes, in order to facilitate transpilation to other dialects. 1098 1099 Args: 1100 expression: The expression tree to transform 1101 1102 Returns: 1103 The modified expression with field names inherited in all structs 1104 """ 1105 if ( 1106 isinstance(expression, exp.Array) 1107 and expression.args.get("struct_name_inheritance") 1108 and isinstance(first_item := seq_get(expression.expressions, 0), exp.Struct) 1109 and all(isinstance(fld, exp.PropertyEQ) for fld in first_item.expressions) 1110 ): 1111 field_names: list[exp.Identifier] = [fld.this for fld in first_item.expressions] 1112 1113 # Apply field names to subsequent structs that don't have them 1114 for struct in expression.expressions[1:]: 1115 if not isinstance(struct, exp.Struct) or len(struct.expressions) != len(field_names): 1116 continue 1117 1118 # Convert unnamed expressions to PropertyEQ with inherited names 1119 new_expressions: list[exp.PropertyEQ] = [] 1120 for i, expr in enumerate(struct.expressions): 1121 if not isinstance(expr, exp.PropertyEQ): 1122 # Create PropertyEQ: field_name := value, preserving the type from the inner expression 1123 property_eq = exp.PropertyEQ( 1124 this=field_names[i].copy(), 1125 expression=expr, 1126 ) 1127 property_eq.type = expr.type 1128 new_expressions.append(property_eq) 1129 else: 1130 new_expressions.append(expr) 1131 1132 struct.set("expressions", new_expressions) 1133 1134 return expression
Inherit field names from the first struct in an array.
BigQuery supports implicitly inheriting names from the first STRUCT in an array:
Example:
ARRAY[ STRUCT('Alice' AS name, 85 AS score), -- defines names STRUCT('Bob', 92), -- inherits names STRUCT('Diana', 95) -- inherits names ]
This transformation makes the field names explicit on all structs by adding PropertyEQ nodes, in order to facilitate transpilation to other dialects.
Arguments:
- expression: The expression tree to transform
Returns:
The modified expression with field names inherited in all structs