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__
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)}"
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]] = {}
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
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 []
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.
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)
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
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.
Inherited Members
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
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
Inherited Members
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
Inherited Members
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
Inherited Members
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__
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.