Edit on GitHub

sqlglot.planner

  1from __future__ import annotations
  2
  3import math
  4import typing as t
  5
  6from sqlglot import alias, exp
  7from sqlglot.helper import name_sequence
  8from sqlglot.optimizer.eliminate_joins import join_condition
  9from sqlglot.optimizer.scope import find_all_in_scope, find_in_scope
 10from collections.abc import Iterator, Sequence, Iterable
 11
 12
 13class Plan:
 14    def __init__(self, expression: exp.Expr) -> None:
 15        self.expression: exp.Expr = expression.copy()
 16        with_: exp.With | None = self.expression.args.get("with_")
 17        self.ctes: exp.With | None = with_.copy() if with_ is not None else None
 18        self.root: Step = Step.from_expression(self.expression)
 19        self._dag: dict[Step, set[Step]] = {}
 20
 21    @property
 22    def dag(self) -> dict[Step, set[Step]]:
 23        if not self._dag:
 24            dag: dict[Step, set[Step]] = {}
 25            nodes = {self.root}
 26
 27            while nodes:
 28                node = nodes.pop()
 29                dag[node] = set()
 30
 31                for dep in node.dependencies:
 32                    dag[node].add(dep)
 33                    nodes.add(dep)
 34
 35            self._dag = dag
 36
 37        return self._dag
 38
 39    @property
 40    def leaves(self) -> Iterator[Step]:
 41        return (node for node, deps in self.dag.items() if not deps)
 42
 43    def __repr__(self) -> str:
 44        return f"Plan\n----\n{repr(self.root)}"
 45
 46
 47class Step:
 48    @classmethod
 49    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
 50        """
 51        Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine.
 52        Note: the expression's tables and subqueries must be aliased for this method to work. For
 53        example, given the following expression:
 54
 55        SELECT
 56          x.a,
 57          SUM(x.b)
 58        FROM x AS x
 59        JOIN y AS y
 60          ON x.a = y.a
 61        GROUP BY x.a
 62
 63        the following DAG is produced (the expression IDs might differ per execution):
 64
 65        - Aggregate: x (4347984624)
 66            Context:
 67              Aggregations:
 68                - SUM(x.b)
 69              Group:
 70                - x.a
 71            Projections:
 72              - x.a
 73              - "x".""
 74            Dependencies:
 75            - Join: x (4347985296)
 76              Context:
 77                y:
 78                On: x.a = y.a
 79              Projections:
 80              Dependencies:
 81              - Scan: x (4347983136)
 82                Context:
 83                  Source: x AS x
 84                Projections:
 85              - Scan: y (4343416624)
 86                Context:
 87                  Source: y AS y
 88                Projections:
 89
 90        Args:
 91            expression: the expression to build the DAG from.
 92            ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
 93
 94        Returns:
 95            A Step DAG corresponding to `expression`.
 96        """
 97        ctes = ctes or {}
 98        expression = expression.unnest()
 99        with_: exp.With | None = expression.args.get("with_")
100
101        # CTEs break the mold of scope and introduce themselves to all in the context.
102        if with_ is not None:
103            ctes = ctes.copy()
104            for cte in with_.expressions:
105                step = Step.from_expression(cte.this, ctes)
106                step.name = cte.alias
107                ctes[step.name] = step  # type: ignore
108
109        from_ = expression.args.get("from_")
110
111        if isinstance(expression, exp.Select) and from_:
112            step = Scan.from_expression(from_.this, ctes)
113        elif isinstance(expression, exp.SetOperation):
114            step = SetOperation.from_expression(expression, ctes)
115        else:
116            step = Scan()
117
118        joins: list[exp.Join] | None = expression.args.get("joins")
119
120        if joins is not None:
121            join = Join.from_joins(joins, ctes)
122            join.name = step.name
123            join.source_name = step.name
124            join.add_dependency(step)
125            step = join
126
127        # final selects in this chain of steps representing a select
128        projections: list[exp.Expr] = []
129        # intermediate computations of agg funcs eg x + 1 in SUM(x + 1)
130        operands: dict[exp.Expr, str] = {}
131        aggregations: dict[exp.Expr, None] = {}
132        next_operand_name = name_sequence("_a_")
133
134        def extract_agg_operands(expression: exp.Expr) -> bool:
135            agg_funcs = tuple(find_all_in_scope(expression, exp.AggFunc))
136            if agg_funcs:
137                aggregations[expression] = None
138
139            for agg in agg_funcs:
140                for operand in agg.unnest_operands():
141                    targets = (
142                        operand.expressions if isinstance(operand, exp.Distinct) else [operand]
143                    )
144
145                    for target in targets:
146                        if isinstance(target, exp.Column):
147                            continue
148                        if target not in operands:
149                            operands[target] = next_operand_name()
150
151                        target.replace(exp.column(operands[target], quoted=True))
152
153            return bool(agg_funcs)
154
155        def set_ops_and_aggs(step) -> None:
156            step.operands = tuple(alias(operand, alias_) for operand, alias_ in operands.items())
157            step.aggregations = list(aggregations)
158
159        for e in expression.expressions:
160            if find_in_scope(e, exp.AggFunc):
161                projections.append(exp.column(e.alias_or_name, step.name, quoted=True))
162                extract_agg_operands(e)
163            else:
164                projections.append(e)
165
166        where: exp.Where | None = expression.args.get("where")
167
168        if where is not None:
169            step.condition = where.this
170
171        group: exp.Group | None = expression.args.get("group")
172
173        if group is not None or aggregations:
174            aggregate = Aggregate()
175            aggregate.source = step.name
176            aggregate.name = step.name
177
178            having: exp.Having | None = expression.args.get("having")
179
180            if having is not None:
181                if extract_agg_operands(exp.alias_(having.this, "_h", quoted=True)):
182                    aggregate.condition = exp.column("_h", step.name, quoted=True)
183                else:
184                    aggregate.condition = having.this
185
186            set_ops_and_aggs(aggregate)
187
188            # give aggregates names and replace projections with references to them
189            aggregate.group = {
190                f"_g{i}": e for i, e in enumerate(group.expressions if group else [])
191            }
192
193            intermediate: dict[str | exp.Expr, str] = {}
194            for k, v in aggregate.group.items():
195                intermediate[v] = k
196                if isinstance(v, exp.Column):
197                    intermediate[v.name] = k
198
199            for projection in projections:
200                for node in projection.walk():
201                    name = intermediate.get(node)
202                    if name:
203                        node.replace(exp.column(name, step.name))
204
205            if aggregate.condition:
206                for node in aggregate.condition.walk():
207                    name = intermediate.get(node) or intermediate.get(node.name)
208                    if name:
209                        node.replace(exp.column(name, step.name))
210
211            aggregate.add_dependency(step)
212            step = aggregate
213        else:
214            aggregate = None
215
216        # Plan DISTINCT before ORDER BY, since Aggregate sorts by its own group key
217        if isinstance(expression, exp.Select) and expression.args.get("distinct"):
218            distinct = Aggregate()
219            distinct.source = step.name
220            distinct.name = step.name
221            distinct.group = {
222                e.alias_or_name: e.unalias() for e in projections or expression.expressions
223            }
224            projections = [exp.column(name, step.name, quoted=True) for name in distinct.group]
225            distinct.add_dependency(step)
226            step = distinct
227        else:
228            distinct = None
229
230        order: exp.Order | None = expression.args.get("order")
231
232        if order is not None:
233            if aggregate is not None:
234                for i, ordered in enumerate(order.expressions):
235                    if extract_agg_operands(exp.alias_(ordered.this, f"_o_{i}", quoted=True)):
236                        ordered.this.replace(exp.column(f"_o_{i}", aggregate.name, quoted=True))
237
238                set_ops_and_aggs(aggregate)
239
240            if distinct is not None:
241                for i, ordered in enumerate(order.expressions):
242                    key = ordered.this
243                    group_name = next((n for n, e in distinct.group.items() if e == key), None)
244                    if group_name:
245                        key.replace(exp.column(group_name, step.name, quoted=True))
246                        continue
247
248                    # a bare column is a reference to an output name
249                    if isinstance(key, exp.Column) and not key.table and key.name in distinct.group:
250                        continue
251
252                    key = key.copy()
253                    for node in [
254                        n
255                        for n in key.walk()
256                        if isinstance(n, exp.Column) and not n.table and n.name in distinct.group
257                    ]:
258                        node.replace(distinct.group[node.name].copy())
259
260                    # the key has no single value once DISTINCT collapses rows; follow
261                    # duckdb/sqlite and take an arbitrary one
262                    if not isinstance(key, exp.Column):
263                        distinct.operands += (alias(key, f"_a_{i}"),)
264                        key = exp.column(f"_a_{i}", quoted=True)
265
266                    distinct.aggregations.append(
267                        exp.alias_(exp.First(this=key), f"_o_{i}", quoted=True)
268                    )
269                    ordered.this.replace(exp.column(f"_o_{i}", step.name, quoted=True))
270
271            sort = Sort()
272            sort.name = step.name
273            sort.key = order.expressions
274            sort.add_dependency(step)
275            step = sort
276
277        step.projections = projections
278
279        limit: exp.Limit | None = expression.args.get("limit")
280
281        if limit is not None and not limit.is_limit_all:
282            step.limit = int(limit.text("expression"))
283
284        offset: exp.Offset | None = expression.args.get("offset")
285
286        if offset is not None:
287            step.offset = int(offset.text("expression"))
288
289        return step
290
291    def __init__(self) -> None:
292        self.name: str | None = None
293        self.dependencies: set[Step] = set()
294        self.dependents: set[Step] = set()
295        self.projections: Sequence[exp.Expr] = []
296        self.limit: float = math.inf
297        self.offset: int = 0
298        self.condition: exp.Expr | None = None
299
300    def add_dependency(self, dependency: Step) -> None:
301        self.dependencies.add(dependency)
302        dependency.dependents.add(self)
303
304    def __repr__(self) -> str:
305        return self.to_s()
306
307    def to_s(self, level: int = 0) -> str:
308        indent = "  " * level
309        nested = f"{indent}    "
310
311        context = self._to_s(f"{nested}  ")
312
313        if context:
314            context = [f"{nested}Context:"] + context
315
316        lines = [
317            f"{indent}- {self.id}",
318            *context,
319            f"{nested}Projections:",
320        ]
321
322        for expression in self.projections:
323            lines.append(f"{nested}  - {expression.sql()}")
324
325        if self.condition:
326            lines.append(f"{nested}Condition: {self.condition.sql()}")
327
328        if self.limit is not math.inf:
329            lines.append(f"{nested}Limit: {self.limit}")
330
331        if self.offset:
332            lines.append(f"{nested}Offset: {self.offset}")
333
334        if self.dependencies:
335            lines.append(f"{nested}Dependencies:")
336            for dependency in self.dependencies:
337                lines.append("  " + dependency.to_s(level + 1))
338
339        return "\n".join(lines)
340
341    @property
342    def type_name(self) -> str:
343        return self.__class__.__name__
344
345    @property
346    def id(self) -> str:
347        name = self.name
348        name = f" {name}" if name else ""
349        return f"{self.type_name}:{name} ({id(self)})"
350
351    def _to_s(self, _indent: str) -> list[str]:
352        return []
353
354
355class Scan(Step):
356    @classmethod
357    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
358        table: exp.Expr = expression
359        alias_ = expression.alias_or_name
360
361        if isinstance(expression, exp.Subquery):
362            table = expression.this
363            step = Step.from_expression(table, ctes)
364            step.name = alias_
365            return step
366
367        step = Scan()
368        step.name = alias_
369        step.source = expression
370        if ctes and table.name in ctes:
371            step.add_dependency(ctes[table.name])
372
373        return step
374
375    def __init__(self) -> None:
376        super().__init__()
377        self.source: exp.Expr | None = None
378
379    def _to_s(self, indent: str) -> list[str]:
380        return [f"{indent}Source: {self.source.sql() if self.source else '-static-'}"]  # type: ignore
381
382
383class Join(Step):
384    @classmethod
385    def from_joins(cls, joins: Iterable[exp.Join], ctes: dict[str, Step] | None = None) -> Join:
386        step = Join()
387
388        for join in joins:
389            source_key, join_key, condition = join_condition(join)
390            step.joins[join.alias_or_name] = {
391                "side": join.side,  # type: ignore
392                "join_key": join_key,
393                "source_key": source_key,
394                "condition": condition,
395            }
396
397            step.add_dependency(Scan.from_expression(join.this, ctes))
398
399        return step
400
401    def __init__(self) -> None:
402        super().__init__()
403        self.source_name: str | None = None
404        self.joins: dict[str, dict[str, list[str] | exp.Expr | list[exp.Expr]]] = {}
405
406    def _to_s(self, indent: str) -> list[str]:
407        lines = [f"{indent}Source: {self.source_name or self.name}"]
408        for name, join in self.joins.items():
409            lines.append(f"{indent}{name}: {join['side'] or 'INNER'}")
410            join_key = ", ".join(str(key) for key in t.cast(list[str], join.get("join_key") or []))
411            if join_key:
412                lines.append(f"{indent}Key: {join_key}")
413            if join.get("condition"):
414                lines.append(f"{indent}On: {join['condition'].sql()}")  # type: ignore
415        return lines
416
417
418class Aggregate(Step):
419    def __init__(self) -> None:
420        super().__init__()
421        self.aggregations: list[exp.Expr] = []
422        self.operands: tuple[exp.Expr, ...] = ()
423        self.group: dict[str, exp.Expr] = {}
424        self.source: str | None = None
425
426    def _to_s(self, indent: str) -> list[str]:
427        lines = [f"{indent}Aggregations:"]
428
429        for expression in self.aggregations:
430            lines.append(f"{indent}  - {expression.sql()}")
431
432        if self.group:
433            lines.append(f"{indent}Group:")
434            for expression in self.group.values():
435                lines.append(f"{indent}  - {expression.sql()}")
436        if self.condition:
437            lines.append(f"{indent}Having:")
438            lines.append(f"{indent}  - {self.condition.sql()}")
439        if self.operands:
440            lines.append(f"{indent}Operands:")
441            for expression in self.operands:
442                lines.append(f"{indent}  - {expression.sql()}")
443
444        return lines
445
446
447class Sort(Step):
448    def __init__(self) -> None:
449        super().__init__()
450        self.key: list[exp.Expr] | None = None
451
452    def _to_s(self, indent: str) -> list[str]:
453        lines = [f"{indent}Key:"]
454
455        for expression in self.key:  # type: ignore
456            lines.append(f"{indent}  - {expression.sql()}")
457
458        return lines
459
460
461class SetOperation(Step):
462    def __init__(self, op: type[exp.Expr], left: str, right: str, distinct: bool = False) -> None:
463        super().__init__()
464        self.op: type[exp.Expr] = op
465        self.left: str = left
466        self.right: str = right
467        self.distinct: bool = distinct
468
469    @classmethod
470    def from_expression(
471        cls, expression: exp.Expr, ctes: dict[str, Step] | None = None
472    ) -> SetOperation:
473        assert isinstance(expression, exp.SetOperation)
474
475        left = Step.from_expression(expression.left, ctes)
476        # SELECT 1 UNION SELECT 2  <-- these subqueries don't have names
477        left.name = left.name or "left"
478        right = Step.from_expression(expression.right, ctes)
479        right.name = right.name or "right"
480        step = cls(
481            op=expression.__class__,
482            left=left.name,
483            right=right.name,
484            distinct=bool(expression.args.get("distinct")),
485        )
486
487        step.add_dependency(left)
488        step.add_dependency(right)
489
490        return step
491
492    def _to_s(self, indent: str) -> list[str]:
493        lines: list[str] = []
494        if self.distinct:
495            lines.append(f"{indent}Distinct: {self.distinct}")
496        return lines
497
498    @property
499    def type_name(self) -> str:
500        return self.op.__name__
class Plan:
14class Plan:
15    def __init__(self, expression: exp.Expr) -> None:
16        self.expression: exp.Expr = expression.copy()
17        with_: exp.With | None = self.expression.args.get("with_")
18        self.ctes: exp.With | None = with_.copy() if with_ is not None else None
19        self.root: Step = Step.from_expression(self.expression)
20        self._dag: dict[Step, set[Step]] = {}
21
22    @property
23    def dag(self) -> dict[Step, set[Step]]:
24        if not self._dag:
25            dag: dict[Step, set[Step]] = {}
26            nodes = {self.root}
27
28            while nodes:
29                node = nodes.pop()
30                dag[node] = set()
31
32                for dep in node.dependencies:
33                    dag[node].add(dep)
34                    nodes.add(dep)
35
36            self._dag = dag
37
38        return self._dag
39
40    @property
41    def leaves(self) -> Iterator[Step]:
42        return (node for node, deps in self.dag.items() if not deps)
43
44    def __repr__(self) -> str:
45        return f"Plan\n----\n{repr(self.root)}"
Plan(expression: sqlglot.expressions.core.Expr)
15    def __init__(self, expression: exp.Expr) -> None:
16        self.expression: exp.Expr = expression.copy()
17        with_: exp.With | None = self.expression.args.get("with_")
18        self.ctes: exp.With | None = with_.copy() if with_ is not None else None
19        self.root: Step = Step.from_expression(self.expression)
20        self._dag: dict[Step, set[Step]] = {}
root: Step
dag: dict[Step, set[Step]]
22    @property
23    def dag(self) -> dict[Step, set[Step]]:
24        if not self._dag:
25            dag: dict[Step, set[Step]] = {}
26            nodes = {self.root}
27
28            while nodes:
29                node = nodes.pop()
30                dag[node] = set()
31
32                for dep in node.dependencies:
33                    dag[node].add(dep)
34                    nodes.add(dep)
35
36            self._dag = dag
37
38        return self._dag
leaves: Iterator[Step]
40    @property
41    def leaves(self) -> Iterator[Step]:
42        return (node for node, deps in self.dag.items() if not deps)
class Step:
 48class Step:
 49    @classmethod
 50    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
 51        """
 52        Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine.
 53        Note: the expression's tables and subqueries must be aliased for this method to work. For
 54        example, given the following expression:
 55
 56        SELECT
 57          x.a,
 58          SUM(x.b)
 59        FROM x AS x
 60        JOIN y AS y
 61          ON x.a = y.a
 62        GROUP BY x.a
 63
 64        the following DAG is produced (the expression IDs might differ per execution):
 65
 66        - Aggregate: x (4347984624)
 67            Context:
 68              Aggregations:
 69                - SUM(x.b)
 70              Group:
 71                - x.a
 72            Projections:
 73              - x.a
 74              - "x".""
 75            Dependencies:
 76            - Join: x (4347985296)
 77              Context:
 78                y:
 79                On: x.a = y.a
 80              Projections:
 81              Dependencies:
 82              - Scan: x (4347983136)
 83                Context:
 84                  Source: x AS x
 85                Projections:
 86              - Scan: y (4343416624)
 87                Context:
 88                  Source: y AS y
 89                Projections:
 90
 91        Args:
 92            expression: the expression to build the DAG from.
 93            ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
 94
 95        Returns:
 96            A Step DAG corresponding to `expression`.
 97        """
 98        ctes = ctes or {}
 99        expression = expression.unnest()
100        with_: exp.With | None = expression.args.get("with_")
101
102        # CTEs break the mold of scope and introduce themselves to all in the context.
103        if with_ is not None:
104            ctes = ctes.copy()
105            for cte in with_.expressions:
106                step = Step.from_expression(cte.this, ctes)
107                step.name = cte.alias
108                ctes[step.name] = step  # type: ignore
109
110        from_ = expression.args.get("from_")
111
112        if isinstance(expression, exp.Select) and from_:
113            step = Scan.from_expression(from_.this, ctes)
114        elif isinstance(expression, exp.SetOperation):
115            step = SetOperation.from_expression(expression, ctes)
116        else:
117            step = Scan()
118
119        joins: list[exp.Join] | None = expression.args.get("joins")
120
121        if joins is not None:
122            join = Join.from_joins(joins, ctes)
123            join.name = step.name
124            join.source_name = step.name
125            join.add_dependency(step)
126            step = join
127
128        # final selects in this chain of steps representing a select
129        projections: list[exp.Expr] = []
130        # intermediate computations of agg funcs eg x + 1 in SUM(x + 1)
131        operands: dict[exp.Expr, str] = {}
132        aggregations: dict[exp.Expr, None] = {}
133        next_operand_name = name_sequence("_a_")
134
135        def extract_agg_operands(expression: exp.Expr) -> bool:
136            agg_funcs = tuple(find_all_in_scope(expression, exp.AggFunc))
137            if agg_funcs:
138                aggregations[expression] = None
139
140            for agg in agg_funcs:
141                for operand in agg.unnest_operands():
142                    targets = (
143                        operand.expressions if isinstance(operand, exp.Distinct) else [operand]
144                    )
145
146                    for target in targets:
147                        if isinstance(target, exp.Column):
148                            continue
149                        if target not in operands:
150                            operands[target] = next_operand_name()
151
152                        target.replace(exp.column(operands[target], quoted=True))
153
154            return bool(agg_funcs)
155
156        def set_ops_and_aggs(step) -> None:
157            step.operands = tuple(alias(operand, alias_) for operand, alias_ in operands.items())
158            step.aggregations = list(aggregations)
159
160        for e in expression.expressions:
161            if find_in_scope(e, exp.AggFunc):
162                projections.append(exp.column(e.alias_or_name, step.name, quoted=True))
163                extract_agg_operands(e)
164            else:
165                projections.append(e)
166
167        where: exp.Where | None = expression.args.get("where")
168
169        if where is not None:
170            step.condition = where.this
171
172        group: exp.Group | None = expression.args.get("group")
173
174        if group is not None or aggregations:
175            aggregate = Aggregate()
176            aggregate.source = step.name
177            aggregate.name = step.name
178
179            having: exp.Having | None = expression.args.get("having")
180
181            if having is not None:
182                if extract_agg_operands(exp.alias_(having.this, "_h", quoted=True)):
183                    aggregate.condition = exp.column("_h", step.name, quoted=True)
184                else:
185                    aggregate.condition = having.this
186
187            set_ops_and_aggs(aggregate)
188
189            # give aggregates names and replace projections with references to them
190            aggregate.group = {
191                f"_g{i}": e for i, e in enumerate(group.expressions if group else [])
192            }
193
194            intermediate: dict[str | exp.Expr, str] = {}
195            for k, v in aggregate.group.items():
196                intermediate[v] = k
197                if isinstance(v, exp.Column):
198                    intermediate[v.name] = k
199
200            for projection in projections:
201                for node in projection.walk():
202                    name = intermediate.get(node)
203                    if name:
204                        node.replace(exp.column(name, step.name))
205
206            if aggregate.condition:
207                for node in aggregate.condition.walk():
208                    name = intermediate.get(node) or intermediate.get(node.name)
209                    if name:
210                        node.replace(exp.column(name, step.name))
211
212            aggregate.add_dependency(step)
213            step = aggregate
214        else:
215            aggregate = None
216
217        # Plan DISTINCT before ORDER BY, since Aggregate sorts by its own group key
218        if isinstance(expression, exp.Select) and expression.args.get("distinct"):
219            distinct = Aggregate()
220            distinct.source = step.name
221            distinct.name = step.name
222            distinct.group = {
223                e.alias_or_name: e.unalias() for e in projections or expression.expressions
224            }
225            projections = [exp.column(name, step.name, quoted=True) for name in distinct.group]
226            distinct.add_dependency(step)
227            step = distinct
228        else:
229            distinct = None
230
231        order: exp.Order | None = expression.args.get("order")
232
233        if order is not None:
234            if aggregate is not None:
235                for i, ordered in enumerate(order.expressions):
236                    if extract_agg_operands(exp.alias_(ordered.this, f"_o_{i}", quoted=True)):
237                        ordered.this.replace(exp.column(f"_o_{i}", aggregate.name, quoted=True))
238
239                set_ops_and_aggs(aggregate)
240
241            if distinct is not None:
242                for i, ordered in enumerate(order.expressions):
243                    key = ordered.this
244                    group_name = next((n for n, e in distinct.group.items() if e == key), None)
245                    if group_name:
246                        key.replace(exp.column(group_name, step.name, quoted=True))
247                        continue
248
249                    # a bare column is a reference to an output name
250                    if isinstance(key, exp.Column) and not key.table and key.name in distinct.group:
251                        continue
252
253                    key = key.copy()
254                    for node in [
255                        n
256                        for n in key.walk()
257                        if isinstance(n, exp.Column) and not n.table and n.name in distinct.group
258                    ]:
259                        node.replace(distinct.group[node.name].copy())
260
261                    # the key has no single value once DISTINCT collapses rows; follow
262                    # duckdb/sqlite and take an arbitrary one
263                    if not isinstance(key, exp.Column):
264                        distinct.operands += (alias(key, f"_a_{i}"),)
265                        key = exp.column(f"_a_{i}", quoted=True)
266
267                    distinct.aggregations.append(
268                        exp.alias_(exp.First(this=key), f"_o_{i}", quoted=True)
269                    )
270                    ordered.this.replace(exp.column(f"_o_{i}", step.name, quoted=True))
271
272            sort = Sort()
273            sort.name = step.name
274            sort.key = order.expressions
275            sort.add_dependency(step)
276            step = sort
277
278        step.projections = projections
279
280        limit: exp.Limit | None = expression.args.get("limit")
281
282        if limit is not None and not limit.is_limit_all:
283            step.limit = int(limit.text("expression"))
284
285        offset: exp.Offset | None = expression.args.get("offset")
286
287        if offset is not None:
288            step.offset = int(offset.text("expression"))
289
290        return step
291
292    def __init__(self) -> None:
293        self.name: str | None = None
294        self.dependencies: set[Step] = set()
295        self.dependents: set[Step] = set()
296        self.projections: Sequence[exp.Expr] = []
297        self.limit: float = math.inf
298        self.offset: int = 0
299        self.condition: exp.Expr | None = None
300
301    def add_dependency(self, dependency: Step) -> None:
302        self.dependencies.add(dependency)
303        dependency.dependents.add(self)
304
305    def __repr__(self) -> str:
306        return self.to_s()
307
308    def to_s(self, level: int = 0) -> str:
309        indent = "  " * level
310        nested = f"{indent}    "
311
312        context = self._to_s(f"{nested}  ")
313
314        if context:
315            context = [f"{nested}Context:"] + context
316
317        lines = [
318            f"{indent}- {self.id}",
319            *context,
320            f"{nested}Projections:",
321        ]
322
323        for expression in self.projections:
324            lines.append(f"{nested}  - {expression.sql()}")
325
326        if self.condition:
327            lines.append(f"{nested}Condition: {self.condition.sql()}")
328
329        if self.limit is not math.inf:
330            lines.append(f"{nested}Limit: {self.limit}")
331
332        if self.offset:
333            lines.append(f"{nested}Offset: {self.offset}")
334
335        if self.dependencies:
336            lines.append(f"{nested}Dependencies:")
337            for dependency in self.dependencies:
338                lines.append("  " + dependency.to_s(level + 1))
339
340        return "\n".join(lines)
341
342    @property
343    def type_name(self) -> str:
344        return self.__class__.__name__
345
346    @property
347    def id(self) -> str:
348        name = self.name
349        name = f" {name}" if name else ""
350        return f"{self.type_name}:{name} ({id(self)})"
351
352    def _to_s(self, _indent: str) -> list[str]:
353        return []
@classmethod
def from_expression( cls, expression: sqlglot.expressions.core.Expr, ctes: dict[str, Step] | None = None) -> Step:
 49    @classmethod
 50    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
 51        """
 52        Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine.
 53        Note: the expression's tables and subqueries must be aliased for this method to work. For
 54        example, given the following expression:
 55
 56        SELECT
 57          x.a,
 58          SUM(x.b)
 59        FROM x AS x
 60        JOIN y AS y
 61          ON x.a = y.a
 62        GROUP BY x.a
 63
 64        the following DAG is produced (the expression IDs might differ per execution):
 65
 66        - Aggregate: x (4347984624)
 67            Context:
 68              Aggregations:
 69                - SUM(x.b)
 70              Group:
 71                - x.a
 72            Projections:
 73              - x.a
 74              - "x".""
 75            Dependencies:
 76            - Join: x (4347985296)
 77              Context:
 78                y:
 79                On: x.a = y.a
 80              Projections:
 81              Dependencies:
 82              - Scan: x (4347983136)
 83                Context:
 84                  Source: x AS x
 85                Projections:
 86              - Scan: y (4343416624)
 87                Context:
 88                  Source: y AS y
 89                Projections:
 90
 91        Args:
 92            expression: the expression to build the DAG from.
 93            ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
 94
 95        Returns:
 96            A Step DAG corresponding to `expression`.
 97        """
 98        ctes = ctes or {}
 99        expression = expression.unnest()
100        with_: exp.With | None = expression.args.get("with_")
101
102        # CTEs break the mold of scope and introduce themselves to all in the context.
103        if with_ is not None:
104            ctes = ctes.copy()
105            for cte in with_.expressions:
106                step = Step.from_expression(cte.this, ctes)
107                step.name = cte.alias
108                ctes[step.name] = step  # type: ignore
109
110        from_ = expression.args.get("from_")
111
112        if isinstance(expression, exp.Select) and from_:
113            step = Scan.from_expression(from_.this, ctes)
114        elif isinstance(expression, exp.SetOperation):
115            step = SetOperation.from_expression(expression, ctes)
116        else:
117            step = Scan()
118
119        joins: list[exp.Join] | None = expression.args.get("joins")
120
121        if joins is not None:
122            join = Join.from_joins(joins, ctes)
123            join.name = step.name
124            join.source_name = step.name
125            join.add_dependency(step)
126            step = join
127
128        # final selects in this chain of steps representing a select
129        projections: list[exp.Expr] = []
130        # intermediate computations of agg funcs eg x + 1 in SUM(x + 1)
131        operands: dict[exp.Expr, str] = {}
132        aggregations: dict[exp.Expr, None] = {}
133        next_operand_name = name_sequence("_a_")
134
135        def extract_agg_operands(expression: exp.Expr) -> bool:
136            agg_funcs = tuple(find_all_in_scope(expression, exp.AggFunc))
137            if agg_funcs:
138                aggregations[expression] = None
139
140            for agg in agg_funcs:
141                for operand in agg.unnest_operands():
142                    targets = (
143                        operand.expressions if isinstance(operand, exp.Distinct) else [operand]
144                    )
145
146                    for target in targets:
147                        if isinstance(target, exp.Column):
148                            continue
149                        if target not in operands:
150                            operands[target] = next_operand_name()
151
152                        target.replace(exp.column(operands[target], quoted=True))
153
154            return bool(agg_funcs)
155
156        def set_ops_and_aggs(step) -> None:
157            step.operands = tuple(alias(operand, alias_) for operand, alias_ in operands.items())
158            step.aggregations = list(aggregations)
159
160        for e in expression.expressions:
161            if find_in_scope(e, exp.AggFunc):
162                projections.append(exp.column(e.alias_or_name, step.name, quoted=True))
163                extract_agg_operands(e)
164            else:
165                projections.append(e)
166
167        where: exp.Where | None = expression.args.get("where")
168
169        if where is not None:
170            step.condition = where.this
171
172        group: exp.Group | None = expression.args.get("group")
173
174        if group is not None or aggregations:
175            aggregate = Aggregate()
176            aggregate.source = step.name
177            aggregate.name = step.name
178
179            having: exp.Having | None = expression.args.get("having")
180
181            if having is not None:
182                if extract_agg_operands(exp.alias_(having.this, "_h", quoted=True)):
183                    aggregate.condition = exp.column("_h", step.name, quoted=True)
184                else:
185                    aggregate.condition = having.this
186
187            set_ops_and_aggs(aggregate)
188
189            # give aggregates names and replace projections with references to them
190            aggregate.group = {
191                f"_g{i}": e for i, e in enumerate(group.expressions if group else [])
192            }
193
194            intermediate: dict[str | exp.Expr, str] = {}
195            for k, v in aggregate.group.items():
196                intermediate[v] = k
197                if isinstance(v, exp.Column):
198                    intermediate[v.name] = k
199
200            for projection in projections:
201                for node in projection.walk():
202                    name = intermediate.get(node)
203                    if name:
204                        node.replace(exp.column(name, step.name))
205
206            if aggregate.condition:
207                for node in aggregate.condition.walk():
208                    name = intermediate.get(node) or intermediate.get(node.name)
209                    if name:
210                        node.replace(exp.column(name, step.name))
211
212            aggregate.add_dependency(step)
213            step = aggregate
214        else:
215            aggregate = None
216
217        # Plan DISTINCT before ORDER BY, since Aggregate sorts by its own group key
218        if isinstance(expression, exp.Select) and expression.args.get("distinct"):
219            distinct = Aggregate()
220            distinct.source = step.name
221            distinct.name = step.name
222            distinct.group = {
223                e.alias_or_name: e.unalias() for e in projections or expression.expressions
224            }
225            projections = [exp.column(name, step.name, quoted=True) for name in distinct.group]
226            distinct.add_dependency(step)
227            step = distinct
228        else:
229            distinct = None
230
231        order: exp.Order | None = expression.args.get("order")
232
233        if order is not None:
234            if aggregate is not None:
235                for i, ordered in enumerate(order.expressions):
236                    if extract_agg_operands(exp.alias_(ordered.this, f"_o_{i}", quoted=True)):
237                        ordered.this.replace(exp.column(f"_o_{i}", aggregate.name, quoted=True))
238
239                set_ops_and_aggs(aggregate)
240
241            if distinct is not None:
242                for i, ordered in enumerate(order.expressions):
243                    key = ordered.this
244                    group_name = next((n for n, e in distinct.group.items() if e == key), None)
245                    if group_name:
246                        key.replace(exp.column(group_name, step.name, quoted=True))
247                        continue
248
249                    # a bare column is a reference to an output name
250                    if isinstance(key, exp.Column) and not key.table and key.name in distinct.group:
251                        continue
252
253                    key = key.copy()
254                    for node in [
255                        n
256                        for n in key.walk()
257                        if isinstance(n, exp.Column) and not n.table and n.name in distinct.group
258                    ]:
259                        node.replace(distinct.group[node.name].copy())
260
261                    # the key has no single value once DISTINCT collapses rows; follow
262                    # duckdb/sqlite and take an arbitrary one
263                    if not isinstance(key, exp.Column):
264                        distinct.operands += (alias(key, f"_a_{i}"),)
265                        key = exp.column(f"_a_{i}", quoted=True)
266
267                    distinct.aggregations.append(
268                        exp.alias_(exp.First(this=key), f"_o_{i}", quoted=True)
269                    )
270                    ordered.this.replace(exp.column(f"_o_{i}", step.name, quoted=True))
271
272            sort = Sort()
273            sort.name = step.name
274            sort.key = order.expressions
275            sort.add_dependency(step)
276            step = sort
277
278        step.projections = projections
279
280        limit: exp.Limit | None = expression.args.get("limit")
281
282        if limit is not None and not limit.is_limit_all:
283            step.limit = int(limit.text("expression"))
284
285        offset: exp.Offset | None = expression.args.get("offset")
286
287        if offset is not None:
288            step.offset = int(offset.text("expression"))
289
290        return step

Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine. Note: the expression's tables and subqueries must be aliased for this method to work. For example, given the following expression:

SELECT x.a, SUM(x.b) FROM x AS x JOIN y AS y ON x.a = y.a GROUP BY x.a

the following DAG is produced (the expression IDs might differ per execution):

  • Aggregate: x (4347984624) Context: Aggregations: - SUM(x.b) Group: - x.a Projections:
    • x.a
    • "x"."" Dependencies:
      • Join: x (4347985296) Context: y: On: x.a = y.a Projections: Dependencies:
    • Scan: x (4347983136) Context: Source: x AS x Projections:
    • Scan: y (4343416624) Context: Source: y AS y Projections:
Arguments:
  • expression: the expression to build the DAG from.
  • ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
Returns:

A Step DAG corresponding to expression.

name: str | None
dependencies: set[Step]
dependents: set[Step]
projections: Sequence[sqlglot.expressions.core.Expr]
limit: float
offset: int
condition: sqlglot.expressions.core.Expr | None
def add_dependency(self, dependency: Step) -> None:
301    def add_dependency(self, dependency: Step) -> None:
302        self.dependencies.add(dependency)
303        dependency.dependents.add(self)
def to_s(self, level: int = 0) -> str:
308    def to_s(self, level: int = 0) -> str:
309        indent = "  " * level
310        nested = f"{indent}    "
311
312        context = self._to_s(f"{nested}  ")
313
314        if context:
315            context = [f"{nested}Context:"] + context
316
317        lines = [
318            f"{indent}- {self.id}",
319            *context,
320            f"{nested}Projections:",
321        ]
322
323        for expression in self.projections:
324            lines.append(f"{nested}  - {expression.sql()}")
325
326        if self.condition:
327            lines.append(f"{nested}Condition: {self.condition.sql()}")
328
329        if self.limit is not math.inf:
330            lines.append(f"{nested}Limit: {self.limit}")
331
332        if self.offset:
333            lines.append(f"{nested}Offset: {self.offset}")
334
335        if self.dependencies:
336            lines.append(f"{nested}Dependencies:")
337            for dependency in self.dependencies:
338                lines.append("  " + dependency.to_s(level + 1))
339
340        return "\n".join(lines)
type_name: str
342    @property
343    def type_name(self) -> str:
344        return self.__class__.__name__
id: str
346    @property
347    def id(self) -> str:
348        name = self.name
349        name = f" {name}" if name else ""
350        return f"{self.type_name}:{name} ({id(self)})"
class Scan(Step):
356class Scan(Step):
357    @classmethod
358    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
359        table: exp.Expr = expression
360        alias_ = expression.alias_or_name
361
362        if isinstance(expression, exp.Subquery):
363            table = expression.this
364            step = Step.from_expression(table, ctes)
365            step.name = alias_
366            return step
367
368        step = Scan()
369        step.name = alias_
370        step.source = expression
371        if ctes and table.name in ctes:
372            step.add_dependency(ctes[table.name])
373
374        return step
375
376    def __init__(self) -> None:
377        super().__init__()
378        self.source: exp.Expr | None = None
379
380    def _to_s(self, indent: str) -> list[str]:
381        return [f"{indent}Source: {self.source.sql() if self.source else '-static-'}"]  # type: ignore
@classmethod
def from_expression( cls, expression: sqlglot.expressions.core.Expr, ctes: dict[str, Step] | None = None) -> Step:
357    @classmethod
358    def from_expression(cls, expression: exp.Expr, ctes: dict[str, Step] | None = None) -> Step:
359        table: exp.Expr = expression
360        alias_ = expression.alias_or_name
361
362        if isinstance(expression, exp.Subquery):
363            table = expression.this
364            step = Step.from_expression(table, ctes)
365            step.name = alias_
366            return step
367
368        step = Scan()
369        step.name = alias_
370        step.source = expression
371        if ctes and table.name in ctes:
372            step.add_dependency(ctes[table.name])
373
374        return step

Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine. Note: the expression's tables and subqueries must be aliased for this method to work. For example, given the following expression:

SELECT x.a, SUM(x.b) FROM x AS x JOIN y AS y ON x.a = y.a GROUP BY x.a

the following DAG is produced (the expression IDs might differ per execution):

  • Aggregate: x (4347984624) Context: Aggregations: - SUM(x.b) Group: - x.a Projections:
    • x.a
    • "x"."" Dependencies:
      • Join: x (4347985296) Context: y: On: x.a = y.a Projections: Dependencies:
    • Scan: x (4347983136) Context: Source: x AS x Projections:
    • Scan: y (4343416624) Context: Source: y AS y Projections:
Arguments:
  • expression: the expression to build the DAG from.
  • ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
Returns:

A Step DAG corresponding to expression.

class Join(Step):
384class Join(Step):
385    @classmethod
386    def from_joins(cls, joins: Iterable[exp.Join], ctes: dict[str, Step] | None = None) -> Join:
387        step = Join()
388
389        for join in joins:
390            source_key, join_key, condition = join_condition(join)
391            step.joins[join.alias_or_name] = {
392                "side": join.side,  # type: ignore
393                "join_key": join_key,
394                "source_key": source_key,
395                "condition": condition,
396            }
397
398            step.add_dependency(Scan.from_expression(join.this, ctes))
399
400        return step
401
402    def __init__(self) -> None:
403        super().__init__()
404        self.source_name: str | None = None
405        self.joins: dict[str, dict[str, list[str] | exp.Expr | list[exp.Expr]]] = {}
406
407    def _to_s(self, indent: str) -> list[str]:
408        lines = [f"{indent}Source: {self.source_name or self.name}"]
409        for name, join in self.joins.items():
410            lines.append(f"{indent}{name}: {join['side'] or 'INNER'}")
411            join_key = ", ".join(str(key) for key in t.cast(list[str], join.get("join_key") or []))
412            if join_key:
413                lines.append(f"{indent}Key: {join_key}")
414            if join.get("condition"):
415                lines.append(f"{indent}On: {join['condition'].sql()}")  # type: ignore
416        return lines
@classmethod
def from_joins( cls, joins: Iterable[sqlglot.expressions.query.Join], ctes: dict[str, Step] | None = None) -> Join:
385    @classmethod
386    def from_joins(cls, joins: Iterable[exp.Join], ctes: dict[str, Step] | None = None) -> Join:
387        step = Join()
388
389        for join in joins:
390            source_key, join_key, condition = join_condition(join)
391            step.joins[join.alias_or_name] = {
392                "side": join.side,  # type: ignore
393                "join_key": join_key,
394                "source_key": source_key,
395                "condition": condition,
396            }
397
398            step.add_dependency(Scan.from_expression(join.this, ctes))
399
400        return step
source_name: str | None
joins: dict[str, dict[str, list[str] | sqlglot.expressions.core.Expr | list[sqlglot.expressions.core.Expr]]]
class Aggregate(Step):
419class Aggregate(Step):
420    def __init__(self) -> None:
421        super().__init__()
422        self.aggregations: list[exp.Expr] = []
423        self.operands: tuple[exp.Expr, ...] = ()
424        self.group: dict[str, exp.Expr] = {}
425        self.source: str | None = None
426
427    def _to_s(self, indent: str) -> list[str]:
428        lines = [f"{indent}Aggregations:"]
429
430        for expression in self.aggregations:
431            lines.append(f"{indent}  - {expression.sql()}")
432
433        if self.group:
434            lines.append(f"{indent}Group:")
435            for expression in self.group.values():
436                lines.append(f"{indent}  - {expression.sql()}")
437        if self.condition:
438            lines.append(f"{indent}Having:")
439            lines.append(f"{indent}  - {self.condition.sql()}")
440        if self.operands:
441            lines.append(f"{indent}Operands:")
442            for expression in self.operands:
443                lines.append(f"{indent}  - {expression.sql()}")
444
445        return lines
aggregations: list[sqlglot.expressions.core.Expr]
operands: tuple[sqlglot.expressions.core.Expr, ...]
group: dict[str, sqlglot.expressions.core.Expr]
source: str | None
class Sort(Step):
448class Sort(Step):
449    def __init__(self) -> None:
450        super().__init__()
451        self.key: list[exp.Expr] | None = None
452
453    def _to_s(self, indent: str) -> list[str]:
454        lines = [f"{indent}Key:"]
455
456        for expression in self.key:  # type: ignore
457            lines.append(f"{indent}  - {expression.sql()}")
458
459        return lines
key: list[sqlglot.expressions.core.Expr] | None
class SetOperation(Step):
462class SetOperation(Step):
463    def __init__(self, op: type[exp.Expr], left: str, right: str, distinct: bool = False) -> None:
464        super().__init__()
465        self.op: type[exp.Expr] = op
466        self.left: str = left
467        self.right: str = right
468        self.distinct: bool = distinct
469
470    @classmethod
471    def from_expression(
472        cls, expression: exp.Expr, ctes: dict[str, Step] | None = None
473    ) -> SetOperation:
474        assert isinstance(expression, exp.SetOperation)
475
476        left = Step.from_expression(expression.left, ctes)
477        # SELECT 1 UNION SELECT 2  <-- these subqueries don't have names
478        left.name = left.name or "left"
479        right = Step.from_expression(expression.right, ctes)
480        right.name = right.name or "right"
481        step = cls(
482            op=expression.__class__,
483            left=left.name,
484            right=right.name,
485            distinct=bool(expression.args.get("distinct")),
486        )
487
488        step.add_dependency(left)
489        step.add_dependency(right)
490
491        return step
492
493    def _to_s(self, indent: str) -> list[str]:
494        lines: list[str] = []
495        if self.distinct:
496            lines.append(f"{indent}Distinct: {self.distinct}")
497        return lines
498
499    @property
500    def type_name(self) -> str:
501        return self.op.__name__
SetOperation( op: type[sqlglot.expressions.core.Expr], left: str, right: str, distinct: bool = False)
463    def __init__(self, op: type[exp.Expr], left: str, right: str, distinct: bool = False) -> None:
464        super().__init__()
465        self.op: type[exp.Expr] = op
466        self.left: str = left
467        self.right: str = right
468        self.distinct: bool = distinct
left: str
right: str
distinct: bool
@classmethod
def from_expression( cls, expression: sqlglot.expressions.core.Expr, ctes: dict[str, Step] | None = None) -> SetOperation:
470    @classmethod
471    def from_expression(
472        cls, expression: exp.Expr, ctes: dict[str, Step] | None = None
473    ) -> SetOperation:
474        assert isinstance(expression, exp.SetOperation)
475
476        left = Step.from_expression(expression.left, ctes)
477        # SELECT 1 UNION SELECT 2  <-- these subqueries don't have names
478        left.name = left.name or "left"
479        right = Step.from_expression(expression.right, ctes)
480        right.name = right.name or "right"
481        step = cls(
482            op=expression.__class__,
483            left=left.name,
484            right=right.name,
485            distinct=bool(expression.args.get("distinct")),
486        )
487
488        step.add_dependency(left)
489        step.add_dependency(right)
490
491        return step

Builds a DAG of Steps from a SQL expression so that it's easier to execute in an engine. Note: the expression's tables and subqueries must be aliased for this method to work. For example, given the following expression:

SELECT x.a, SUM(x.b) FROM x AS x JOIN y AS y ON x.a = y.a GROUP BY x.a

the following DAG is produced (the expression IDs might differ per execution):

  • Aggregate: x (4347984624) Context: Aggregations: - SUM(x.b) Group: - x.a Projections:
    • x.a
    • "x"."" Dependencies:
      • Join: x (4347985296) Context: y: On: x.a = y.a Projections: Dependencies:
    • Scan: x (4347983136) Context: Source: x AS x Projections:
    • Scan: y (4343416624) Context: Source: y AS y Projections:
Arguments:
  • expression: the expression to build the DAG from.
  • ctes: a dictionary that maps CTEs to their corresponding Step DAG by name.
Returns:

A Step DAG corresponding to expression.

type_name: str
499    @property
500    def type_name(self) -> str:
501        return self.op.__name__