Edit on GitHub

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
class SqlHandler(typing.Protocol):
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:
        ...
SqlHandler(*args, **kwargs)
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)
def preprocess( transforms: list[typing.Callable[[sqlglot.expressions.core.Expr], sqlglot.expressions.core.Expr]], generator: Optional[Callable[[sqlglot.generator.Generator, sqlglot.expressions.core.Expr], str]] = None) -> Callable[[sqlglot.generator.Generator, sqlglot.expressions.core.Expr], str]:
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.

def unnest_generate_date_array_using_recursive_cte( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
 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
def unnest_generate_series( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def eliminate_distinct_on( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def eliminate_qualify( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def remove_precision_parameterized_types( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def unqualify_unnest( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def unnest_to_explode( expression: sqlglot.expressions.core.Expr, unnest_using_arrays_zip: bool = True) -> sqlglot.expressions.core.Expr:
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.

def explode_projection_to_unnest( index_offset: int = 0, unnest_map: bool = False) -> Callable[[sqlglot.expressions.core.Expr], sqlglot.expressions.core.Expr]:
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.

def add_within_group_for_percentiles( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def remove_within_group_for_percentiles( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def add_recursive_cte_column_names( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def epoch_cast_to_ts( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def eliminate_semi_and_anti_joins( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def eliminate_full_outer_join( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def move_ctes_to_top_level(expression: ~E) -> ~E:
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).

def ensure_bools( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def unqualify_columns( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
783def unqualify_columns(expression: exp.Expr) -> exp.Expr:
784    for column in expression.find_all(exp.Column):
785        # We only wanna pop off the table, db, catalog args
786        for part in column.parts[:-1]:
787            part.pop()
788
789    return expression
def unqualify_pivot_fields( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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'))
def remove_unique_constraints( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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
def ctas_with_tmp_tables_to_create_tmp_view( expression: sqlglot.expressions.core.Expr, tmp_storage_provider: Callable[[sqlglot.expressions.core.Expr], sqlglot.expressions.core.Expr] = <function <lambda>>) -> sqlglot.expressions.core.Expr:
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
def move_schema_columns_to_partitioned_by( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def move_partitioned_by_to_schema_columns( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def struct_kv_to_alias( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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).

def eliminate_join_marks( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
 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

  1. You cannot specify the (+) operator in a query block that also contains FROM clause join syntax.

  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.

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

def any_to_exists( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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

def eliminate_window_clause( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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.

def inherit_struct_field_names( expression: sqlglot.expressions.core.Expr) -> sqlglot.expressions.core.Expr:
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