mz_expr/relation.rs
1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10#![warn(missing_docs)]
11
12use std::cell::RefCell;
13use std::cmp::{Ordering, max};
14use std::collections::{BTreeMap, BTreeSet};
15use std::fmt;
16use std::fmt::{Display, Formatter};
17use std::hash::{DefaultHasher, Hash, Hasher};
18use std::num::NonZeroU64;
19use std::time::Instant;
20
21use bytesize::ByteSize;
22use columnation::{Columnation, CopyRegion};
23use itertools::Itertools;
24use mz_ore::cast::{CastFrom, CastInto};
25use mz_ore::collections::CollectionExt;
26use mz_ore::id_gen::IdGen;
27use mz_ore::metrics::Histogram;
28use mz_ore::num::NonNeg;
29use mz_ore::soft_assert_no_log;
30use mz_ore::stack::RecursionLimitError;
31use mz_ore::str::Indent;
32use mz_repr::adt::numeric::NumericMaxScale;
33use mz_repr::explain::text::text_string_at;
34use mz_repr::explain::{
35 DummyHumanizer, ExplainConfig, ExprHumanizer, IndexUsageType, PlanRenderingContext,
36};
37use mz_repr::{
38 ColumnName, Datum, DatumVec, Diff, GlobalId, IntoRowIterator, ReprColumnType, ReprRelationType,
39 ReprScalarType, Row, RowIterator, RowRef, SqlColumnType, SqlRelationType, SqlScalarType,
40};
41use serde::{Deserialize, Serialize};
42
43use crate::Id::Local;
44use crate::explain::{HumanizedExpr, HumanizerMode};
45use crate::relation::func::{AggregateFunc, LagLeadType, TableFunc};
46use crate::row::{RowCollection, RowCollectionIter};
47use crate::scalar::columns::Columns;
48use crate::scalar::func::variadic::{
49 JsonbBuildArray, JsonbBuildObject, ListCreate, ListIndex, MapBuild, RecordCreate,
50};
51use crate::visit::{Visit, VisitChildren};
52use crate::{
53 EvalError, FilterCharacteristics, Id, LocalId, MirScalarExpr, UnaryFunc, func as scalar_func,
54};
55
56pub mod canonicalize;
57pub mod func;
58pub mod join_input_mapper;
59
60/// A recursion limit to be used for stack-safe traversals of [`MirRelationExpr`] trees.
61///
62/// The recursion limit must be large enough to accommodate for the linear representation
63/// of some pathological but frequently occurring query fragments.
64///
65/// For example, in MIR we could have long chains of
66/// - (1) `Let` bindings,
67/// - (2) `CallBinary` calls with associative functions such as `+`
68///
69/// Until we fix those, we need to stick with the larger recursion limit.
70pub const RECURSION_LIMIT: usize = 2048;
71
72/// A trait for types that describe how to build a collection.
73pub trait CollectionPlan {
74 /// Collects the set of global identifiers from dataflows referenced in Get.
75 fn depends_on_into(&self, out: &mut BTreeSet<GlobalId>);
76
77 /// Returns the set of global identifiers from dataflows referenced in Get.
78 ///
79 /// See [`CollectionPlan::depends_on_into`] to reuse an existing `BTreeSet`.
80 fn depends_on(&self) -> BTreeSet<GlobalId> {
81 let mut out = BTreeSet::new();
82 self.depends_on_into(&mut out);
83 out
84 }
85}
86
87/// An abstract syntax tree which defines a collection.
88///
89/// The AST is meant to reflect the capabilities of the `differential_dataflow::Collection` type,
90/// written generically enough to avoid run-time compilation work.
91///
92/// `derived_hash_with_manual_eq` was complaining for the wrong reason: This lint exists because
93/// it's bad when `Eq` doesn't agree with `Hash`, which is often quite likely if one of them is
94/// implemented manually. However, our manual implementation of `Eq` _will_ agree with the derived
95/// one. This is because the reason for the manual implementation is not to change the semantics
96/// from the derived one, but to avoid stack overflows.
97#[allow(clippy::derived_hash_with_manual_eq)]
98#[derive(Clone, Debug, Ord, PartialOrd, Serialize, Deserialize, Hash)]
99pub enum MirRelationExpr {
100 /// A constant relation containing specified rows.
101 ///
102 /// The runtime memory footprint of this operator is zero.
103 ///
104 /// When you would like to pattern match on this, consider using `MirRelationExpr::as_const`
105 /// instead, which looks behind `ArrangeBy`s. You might want this matching behavior because
106 /// constant folding doesn't remove `ArrangeBy`s.
107 Constant {
108 /// Rows of the constant collection and their multiplicities.
109 rows: Result<Vec<(Row, Diff)>, EvalError>,
110 /// Schema of the collection.
111 typ: ReprRelationType,
112 },
113 /// Get an existing dataflow.
114 ///
115 /// The runtime memory footprint of this operator is zero.
116 Get {
117 /// The identifier for the collection to load.
118 id: Id,
119 /// Schema of the collection.
120 typ: ReprRelationType,
121 /// If this is a global Get, this will indicate whether we are going to read from Persist or
122 /// from an index, or from a different object in `objects_to_build`. If it's an index, then
123 /// how downstream dataflow operations will use this index is also recorded. This is filled
124 /// by `prune_and_annotate_dataflow_index_imports`. Note that this is not used by the
125 /// lowering to LIR, but is used only by EXPLAIN.
126 access_strategy: AccessStrategy,
127 },
128 /// Introduce a temporary dataflow.
129 ///
130 /// The runtime memory footprint of this operator is zero.
131 Let {
132 /// The identifier to be used in `Get` variants to retrieve `value`.
133 id: LocalId,
134 /// The collection to be bound to `id`.
135 value: Box<MirRelationExpr>,
136 /// The result of the `Let`, evaluated with `id` bound to `value`.
137 body: Box<MirRelationExpr>,
138 },
139 /// Introduce mutually recursive bindings.
140 ///
141 /// Each `LocalId` is immediately bound to an initially empty collection
142 /// with the type of its corresponding `MirRelationExpr`. Repeatedly, each
143 /// binding is evaluated using the current contents of each other binding,
144 /// and is refreshed to contain the new evaluation. This process continues
145 /// through all bindings, and repeats as long as changes continue to occur.
146 ///
147 /// The resulting value of the expression is `body` evaluated once in the
148 /// context of the final iterates.
149 ///
150 /// A zero-binding instance can be replaced by `body`.
151 /// A single-binding instance is equivalent to `MirRelationExpr::Let`.
152 ///
153 /// The runtime memory footprint of this operator is zero.
154 LetRec {
155 /// The identifiers to be used in `Get` variants to retrieve each `value`.
156 ids: Vec<LocalId>,
157 /// The collections to be bound to each `id`.
158 values: Vec<MirRelationExpr>,
159 /// Maximum number of iterations, after which we should artificially force a fixpoint.
160 /// (Whether we error or just stop is configured by `LetRecLimit::return_at_limit`.)
161 /// The per-`LetRec` limit that the user specified is initially copied to each binding to
162 /// accommodate slicing and merging of `LetRec`s in MIR transforms (e.g., `NormalizeLets`).
163 limits: Vec<Option<LetRecLimit>>,
164 /// The result of the `Let`, evaluated with `id` bound to `value`.
165 body: Box<MirRelationExpr>,
166 },
167 /// Project out some columns from a dataflow
168 ///
169 /// The runtime memory footprint of this operator is zero.
170 Project {
171 /// The source collection.
172 input: Box<MirRelationExpr>,
173 /// Indices of columns to retain.
174 outputs: Vec<usize>,
175 },
176 /// Append new columns to a dataflow
177 ///
178 /// The runtime memory footprint of this operator is zero.
179 Map {
180 /// The source collection.
181 input: Box<MirRelationExpr>,
182 /// Expressions which determine values to append to each row.
183 /// An expression may refer to columns in `input` or
184 /// expressions defined earlier in the vector
185 scalars: Vec<MirScalarExpr>,
186 },
187 /// Like Map, but yields zero-or-more output rows per input row
188 ///
189 /// The runtime memory footprint of this operator is zero.
190 FlatMap {
191 /// The source collection
192 input: Box<MirRelationExpr>,
193 /// The table func to apply
194 func: TableFunc,
195 /// The argument to the table func
196 exprs: Vec<MirScalarExpr>,
197 },
198 /// Keep rows from a dataflow where all the predicates are true
199 ///
200 /// The runtime memory footprint of this operator is zero.
201 Filter {
202 /// The source collection.
203 input: Box<MirRelationExpr>,
204 /// Predicates, each of which must be true.
205 predicates: Vec<MirScalarExpr>,
206 },
207 /// Join several collections, where some columns must be equal.
208 ///
209 /// For further details consult the documentation for [`MirRelationExpr::join`].
210 ///
211 /// The runtime memory footprint of this operator can be proportional to
212 /// the sizes of all inputs and the size of all joins of prefixes.
213 /// This may be reduced due to arrangements available at rendering time.
214 Join {
215 /// A sequence of input relations.
216 inputs: Vec<MirRelationExpr>,
217 /// A sequence of equivalence classes of expressions on the cross product of inputs.
218 ///
219 /// Each equivalence class is a list of scalar expressions, where for each class the
220 /// intended interpretation is that all evaluated expressions should be equal.
221 ///
222 /// Each scalar expression is to be evaluated over the cross-product of all records
223 /// from all inputs. In many cases this may just be column selection from specific
224 /// inputs, but more general cases exist (e.g. complex functions of multiple columns
225 /// from multiple inputs, or just constant literals).
226 equivalences: Vec<Vec<MirScalarExpr>>,
227 /// Join implementation information.
228 #[serde(default)]
229 implementation: JoinImplementation,
230 },
231 /// Group a dataflow by some columns and aggregate over each group
232 ///
233 /// The runtime memory footprint of this operator is at most proportional to the
234 /// number of distinct records in the input and output. The actual requirements
235 /// can be less: the number of distinct inputs to each aggregate, summed across
236 /// each aggregate, plus the output size. For more details consult the code that
237 /// builds the associated dataflow.
238 Reduce {
239 /// The source collection.
240 input: Box<MirRelationExpr>,
241 /// Column indices used to form groups.
242 group_key: Vec<MirScalarExpr>,
243 /// Expressions which determine values to append to each row, after the group keys.
244 aggregates: Vec<AggregateExpr>,
245 /// True iff the input is known to monotonically increase (only addition of records).
246 #[serde(default)]
247 monotonic: bool,
248 /// User hint: expected number of values per group key. Used to optimize physical rendering.
249 #[serde(default)]
250 expected_group_size: Option<u64>,
251 },
252 /// Groups and orders within each group, limiting output.
253 ///
254 /// The runtime memory footprint of this operator is proportional to its input and output.
255 TopK {
256 /// The source collection.
257 input: Box<MirRelationExpr>,
258 /// Column indices used to form groups.
259 group_key: Vec<usize>,
260 /// Column indices used to order rows within groups.
261 order_key: Vec<ColumnOrder>,
262 /// Number of records to retain
263 #[serde(default)]
264 limit: Option<MirScalarExpr>,
265 /// Number of records to skip
266 #[serde(default)]
267 offset: usize,
268 /// True iff the input is known to monotonically increase (only addition of records).
269 #[serde(default)]
270 monotonic: bool,
271 /// User-supplied hint: how many rows will have the same group key.
272 #[serde(default)]
273 expected_group_size: Option<u64>,
274 },
275 /// Return a dataflow where the row counts are negated
276 ///
277 /// The runtime memory footprint of this operator is zero.
278 Negate {
279 /// The source collection.
280 input: Box<MirRelationExpr>,
281 },
282 /// Keep rows from a dataflow where the row counts are positive
283 ///
284 /// The runtime memory footprint of this operator is proportional to its input and output.
285 Threshold {
286 /// The source collection.
287 input: Box<MirRelationExpr>,
288 },
289 /// Adds the frequencies of elements in contained sets.
290 ///
291 /// The runtime memory footprint of this operator is zero.
292 Union {
293 /// A source collection.
294 base: Box<MirRelationExpr>,
295 /// Source collections to union.
296 inputs: Vec<MirRelationExpr>,
297 },
298 /// Technically a no-op. Used to render an index. Will be used to optimize queries
299 /// on finer grain. Each `keys` item represents a different index that should be
300 /// produced from the `keys`.
301 ///
302 /// The runtime memory footprint of this operator is proportional to its input.
303 ArrangeBy {
304 /// The source collection
305 input: Box<MirRelationExpr>,
306 /// Columns to arrange `input` by, in order of decreasing primacy
307 keys: Vec<Vec<MirScalarExpr>>,
308 },
309}
310
311impl PartialEq for MirRelationExpr {
312 fn eq(&self, other: &Self) -> bool {
313 // Capture the result and test it wrt `Ord` implementation in test environments.
314 let result = structured_diff::MreDiff::new(self, other).next().is_none();
315 mz_ore::soft_assert_eq_no_log!(result, self.cmp(other) == Ordering::Equal);
316 result
317 }
318}
319impl Eq for MirRelationExpr {}
320
321impl MirRelationExpr {
322 /// Reports the schema of the relation.
323 ///
324 /// This is the SQL-type parallel of [`Self::typ`]; it is merely
325 /// a wrapper around it, returning a [`SqlRelationType`] instead of
326 /// a [`ReprRelationType`].
327 pub fn sql_typ(&self) -> SqlRelationType {
328 let repr_typ = self.typ();
329 SqlRelationType::from_repr(&repr_typ)
330 }
331
332 /// Reports the repr schema of the relation.
333 ///
334 /// This method determines the type through recursive traversal of the
335 /// relation expression, drawing from the types of base collections.
336 /// As such, this is not an especially cheap method, and should be used
337 /// judiciously.
338 ///
339 /// The relation type is computed incrementally with a recursive post-order
340 /// traversal, that accumulates the input types for the relations yet to be
341 /// visited in `type_stack`.
342 pub fn typ(&self) -> ReprRelationType {
343 let mut type_stack = Vec::new();
344 self.visit_pre_post(
345 &mut |e: &MirRelationExpr| -> Option<Vec<&MirRelationExpr>> {
346 match &e {
347 MirRelationExpr::Let { body, .. } => Some(vec![&*body]),
348 MirRelationExpr::LetRec { body, .. } => Some(vec![&*body]),
349 _ => None,
350 }
351 },
352 &mut |e: &MirRelationExpr| {
353 match e {
354 MirRelationExpr::Let { .. } => {
355 let body_typ = type_stack.pop().unwrap();
356 // Insert a dummy relation type for the value, since `typ_with_input_types`
357 // won't look at it, but expects the relation type of the body to be second.
358 type_stack.push(ReprRelationType::empty());
359 type_stack.push(body_typ);
360 }
361 MirRelationExpr::LetRec { values, .. } => {
362 let body_typ = type_stack.pop().unwrap();
363 type_stack.extend(
364 std::iter::repeat(ReprRelationType::empty()).take(values.len()),
365 );
366 // Insert dummy relation types for the values, since `typ_with_input_types`
367 // won't look at them, but expects the relation type of the body to be last.
368 type_stack.push(body_typ);
369 }
370 _ => {}
371 }
372 let num_inputs = e.num_inputs();
373 let relation_type =
374 e.typ_with_input_types(&type_stack[type_stack.len() - num_inputs..]);
375 type_stack.truncate(type_stack.len() - num_inputs);
376 type_stack.push(relation_type);
377 },
378 );
379 assert_eq!(type_stack.len(), 1);
380 type_stack.pop().unwrap()
381 }
382
383 /// Reports the repr schema of the relation given the repr schema of the input relations.
384 pub fn typ_with_input_types(&self, input_types: &[ReprRelationType]) -> ReprRelationType {
385 let column_types = self.col_with_input_cols(input_types.iter().map(|i| &i.column_types));
386 let unique_keys = self.keys_with_input_keys(
387 input_types.iter().map(|i| i.arity()),
388 input_types.iter().map(|i| &i.keys),
389 );
390 ReprRelationType::new(column_types).with_keys(unique_keys)
391 }
392
393 /// Reports the column types of the relation given the column types of the
394 /// input relations.
395 ///
396 /// This method delegates to `try_col_with_input_cols`, panicking if an `Err`
397 /// variant is returned.
398 pub fn col_with_input_cols<'a, I>(&self, input_types: I) -> Vec<ReprColumnType>
399 where
400 I: Iterator<Item = &'a Vec<ReprColumnType>>,
401 {
402 match self.try_col_with_input_cols(input_types) {
403 Ok(col_types) => col_types,
404 Err(err) => panic!("{err}"),
405 }
406 }
407
408 /// Reports the column types of the relation given the column types of the input relations.
409 ///
410 /// `input_types` is required to contain the column types for the input relations of
411 /// the current relation in the same order as they are visited by `try_visit_children`
412 /// method, even though not all may be used for computing the schema of the
413 /// current relation. For example, `Let` expects two input types, one for the
414 /// value relation and one for the body, in that order, but only the one for the
415 /// body is used to determine the type of the `Let` relation.
416 ///
417 /// It is meant to be used during post-order traversals to compute column types
418 /// incrementally.
419 pub fn try_col_with_input_cols<'a, I>(
420 &self,
421 mut input_types: I,
422 ) -> Result<Vec<ReprColumnType>, String>
423 where
424 I: Iterator<Item = &'a Vec<ReprColumnType>>,
425 {
426 use MirRelationExpr::*;
427
428 let col_types = match self {
429 Constant { rows, typ } => {
430 let mut col_types = typ.column_types.clone();
431 let mut seen_null = vec![false; typ.arity()];
432 if let Ok(rows) = rows {
433 for (row, _diff) in rows {
434 for (datum, i) in row.iter().zip_eq(0..typ.arity()) {
435 if datum.is_null() {
436 seen_null[i] = true;
437 }
438 }
439 }
440 }
441 for (&seen_null, i) in seen_null.iter().zip_eq(0..typ.arity()) {
442 if !seen_null {
443 col_types[i].nullable = false;
444 } else {
445 assert!(col_types[i].nullable);
446 }
447 }
448 col_types
449 }
450 Get { typ, .. } => typ.column_types.clone(),
451 Project { outputs, .. } => {
452 let input = input_types.next().unwrap();
453 outputs.iter().map(|&i| input[i].clone()).collect()
454 }
455 Map { scalars, .. } => {
456 let mut result = input_types.next().unwrap().clone();
457 for scalar in scalars.iter() {
458 result.push(scalar.typ(&result))
459 }
460 result
461 }
462 FlatMap { func, .. } => {
463 let mut result = input_types.next().unwrap().clone();
464 result.extend(
465 func.output_sql_type()
466 .column_types
467 .iter()
468 .map(ReprColumnType::from),
469 );
470 result
471 }
472 Filter { predicates, .. } => {
473 let mut result = input_types.next().unwrap().clone();
474
475 // Set as nonnull any columns where null values would cause
476 // any predicate to evaluate to null.
477 for column in non_nullable_columns(predicates) {
478 result[column].nullable = false;
479 }
480 result
481 }
482 Join { equivalences, .. } => {
483 // Concatenate input column types
484 let mut types = input_types.flat_map(|cols| cols.to_owned()).collect_vec();
485 // In an equivalence class, if any column is non-null, then make all non-null
486 for equivalence in equivalences {
487 let col_inds = equivalence
488 .iter()
489 .filter_map(|expr| match expr {
490 MirScalarExpr::Column(col, _name) => Some(*col),
491 _ => None,
492 })
493 .collect_vec();
494 if col_inds.iter().any(|i| !types.get(*i).unwrap().nullable) {
495 for i in col_inds {
496 types.get_mut(i).unwrap().nullable = false;
497 }
498 }
499 }
500 types
501 }
502 Reduce {
503 group_key,
504 aggregates,
505 ..
506 } => {
507 let input = input_types.next().unwrap();
508 group_key
509 .iter()
510 .map(|e| e.typ(input))
511 .chain(aggregates.iter().map(|agg| agg.typ(input)))
512 .collect()
513 }
514 TopK { .. } | Negate { .. } | Threshold { .. } | ArrangeBy { .. } => {
515 input_types.next().unwrap().clone()
516 }
517 Let { .. } => {
518 // skip over the input types for `value`.
519 input_types.nth(1).unwrap().clone()
520 }
521 LetRec { values, .. } => {
522 // skip over the input types for `values`.
523 input_types.nth(values.len()).unwrap().clone()
524 }
525 Union { .. } => {
526 let mut result = input_types.next().unwrap().clone();
527 for input_col_types in input_types {
528 for (base_col, col) in result.iter_mut().zip_eq(input_col_types) {
529 *base_col = base_col
530 .union(col)
531 .map_err(|e| format!("{}\nin plan:\n{}", e, self.pretty()))?;
532 }
533 }
534 result
535 }
536 };
537
538 Ok(col_types)
539 }
540
541 /// Reports the unique keys of the relation given the arities and the unique
542 /// keys of the input relations.
543 ///
544 /// `input_arities` and `input_keys` are required to contain the
545 /// corresponding info for the input relations of
546 /// the current relation in the same order as they are visited by `try_visit_children`
547 /// method, even though not all may be used for computing the schema of the
548 /// current relation. For example, `Let` expects two input types, one for the
549 /// value relation and one for the body, in that order, but only the one for the
550 /// body is used to determine the type of the `Let` relation.
551 ///
552 /// It is meant to be used during post-order traversals to compute unique keys
553 /// incrementally.
554 pub fn keys_with_input_keys<'a, I, J>(
555 &self,
556 mut input_arities: I,
557 mut input_keys: J,
558 ) -> Vec<Vec<usize>>
559 where
560 I: Iterator<Item = usize>,
561 J: Iterator<Item = &'a Vec<Vec<usize>>>,
562 {
563 use MirRelationExpr::*;
564
565 let mut keys = match self {
566 Constant {
567 rows: Ok(rows),
568 typ,
569 } => {
570 let n_cols = typ.arity();
571 // If the `i`th entry is `Some`, then we have not yet observed non-uniqueness in the `i`th column.
572 let mut unique_values_per_col = vec![Some(BTreeSet::<Datum>::default()); n_cols];
573 for (row, diff) in rows {
574 for (i, datum) in row.iter().enumerate() {
575 if datum != Datum::Dummy {
576 if let Some(unique_vals) = &mut unique_values_per_col[i] {
577 let is_dupe = *diff != Diff::ONE || !unique_vals.insert(datum);
578 if is_dupe {
579 unique_values_per_col[i] = None;
580 }
581 }
582 }
583 }
584 }
585 if rows.len() == 0 || (rows.len() == 1 && rows[0].1 == Diff::ONE) {
586 vec![vec![]]
587 } else {
588 // XXX - Multi-column keys are not detected.
589 typ.keys
590 .iter()
591 .cloned()
592 .chain(
593 unique_values_per_col
594 .into_iter()
595 .enumerate()
596 .filter(|(_idx, unique_vals)| unique_vals.is_some())
597 .map(|(idx, _)| vec![idx]),
598 )
599 .collect()
600 }
601 }
602 Constant { rows: Err(_), typ } | Get { typ, .. } => typ.keys.clone(),
603 Threshold { .. } | ArrangeBy { .. } => input_keys.next().unwrap().clone(),
604 Let { .. } => {
605 // skip over the unique keys for value
606 input_keys.nth(1).unwrap().clone()
607 }
608 LetRec { values, .. } => {
609 // skip over the unique keys for value
610 input_keys.nth(values.len()).unwrap().clone()
611 }
612 Project { outputs, .. } => {
613 let input = input_keys.next().unwrap();
614 input
615 .iter()
616 .filter_map(|key_set| {
617 if key_set.iter().all(|k| outputs.contains(k)) {
618 Some(
619 key_set
620 .iter()
621 .map(|c| outputs.iter().position(|o| o == c).unwrap())
622 .collect(),
623 )
624 } else {
625 None
626 }
627 })
628 .collect()
629 }
630 Map { scalars, .. } => {
631 let mut remappings = Vec::new();
632 let arity = input_arities.next().unwrap();
633 for (column, scalar) in scalars.iter().enumerate() {
634 // assess whether the scalar preserves uniqueness,
635 // and could participate in a key!
636
637 fn uniqueness(expr: &MirScalarExpr) -> Option<usize> {
638 match expr {
639 MirScalarExpr::CallUnary { func, expr } => {
640 if func.preserves_uniqueness() {
641 uniqueness(expr)
642 } else {
643 None
644 }
645 }
646 MirScalarExpr::Column(c, _name) => Some(*c),
647 _ => None,
648 }
649 }
650
651 if let Some(c) = uniqueness(scalar) {
652 remappings.push((c, column + arity));
653 }
654 }
655
656 let mut result = input_keys.next().unwrap().clone();
657 let mut new_keys = Vec::new();
658 // Any column in `remappings` could be replaced in a key
659 // by the corresponding c. This could lead to combinatorial
660 // explosion using our current representation, so we wont
661 // do that. Instead, we'll handle the case of one remapping.
662 if remappings.len() == 1 {
663 let (old, new) = remappings.pop().unwrap();
664 for key in &result {
665 if key.contains(&old) {
666 let mut new_key: Vec<usize> =
667 key.iter().cloned().filter(|k| k != &old).collect();
668 new_key.push(new);
669 new_key.sort_unstable();
670 new_keys.push(new_key);
671 }
672 }
673 result.append(&mut new_keys);
674 }
675 result
676 }
677 FlatMap { .. } => {
678 // FlatMap can add duplicate rows, so input keys are no longer
679 // valid
680 vec![]
681 }
682 Negate { .. } => {
683 // Although negate may have distinct records for each key,
684 // the multiplicity is -1 rather than 1. This breaks many
685 // of the optimization uses of "keys".
686 vec![]
687 }
688 Filter { predicates, .. } => {
689 // A filter inherits the keys of its input unless the filters
690 // have reduced the input to a single row, in which case the
691 // keys of the input are `()`.
692 let mut input = input_keys.next().unwrap().clone();
693
694 if !input.is_empty() {
695 // Track columns equated to literals, which we can prune.
696 let mut cols_equal_to_literal = BTreeSet::new();
697
698 // Perform union find on `col1 = col2` to establish
699 // connected components of equated columns. Absent any
700 // equalities, this will be `0 .. #c` (where #c is the
701 // greatest column referenced by a predicate), but each
702 // equality will orient the root of the greater to the root
703 // of the lesser.
704 let mut union_find = Vec::new();
705
706 for expr in predicates.iter() {
707 if let MirScalarExpr::CallBinary {
708 func: crate::BinaryFunc::Eq(_),
709 expr1,
710 expr2,
711 } = expr
712 {
713 if let MirScalarExpr::Column(c, _name) = &**expr1 {
714 if expr2.is_literal_ok() {
715 cols_equal_to_literal.insert(c);
716 }
717 }
718 if let MirScalarExpr::Column(c, _name) = &**expr2 {
719 if expr1.is_literal_ok() {
720 cols_equal_to_literal.insert(c);
721 }
722 }
723 // Perform union-find to equate columns.
724 if let (Some(c1), Some(c2)) = (expr1.as_column(), expr2.as_column()) {
725 if c1 != c2 {
726 // Ensure union_find has entries up to
727 // max(c1, c2) by filling up missing
728 // positions with identity mappings.
729 while union_find.len() <= std::cmp::max(c1, c2) {
730 union_find.push(union_find.len());
731 }
732 let mut r1 = c1; // Find the representative column of [c1].
733 while r1 != union_find[r1] {
734 assert!(union_find[r1] < r1);
735 r1 = union_find[r1];
736 }
737 let mut r2 = c2; // Find the representative column of [c2].
738 while r2 != union_find[r2] {
739 assert!(union_find[r2] < r2);
740 r2 = union_find[r2];
741 }
742 // Union [c1] and [c2] by pointing the
743 // larger to the smaller representative (we
744 // update the remaining equivalence class
745 // members only once after this for-loop).
746 union_find[std::cmp::max(r1, r2)] = std::cmp::min(r1, r2);
747 }
748 }
749 }
750 }
751
752 // Complete union-find by pointing each element at its representative column.
753 for i in 0..union_find.len() {
754 // Iteration not required, as each prior already references the right column.
755 union_find[i] = union_find[union_find[i]];
756 }
757
758 // Remove columns bound to literals, and remap columns equated to earlier columns.
759 // We will re-expand remapped columns in a moment, but this avoids exponential work.
760 for key_set in &mut input {
761 key_set.retain(|k| !cols_equal_to_literal.contains(&k));
762 for col in key_set.iter_mut() {
763 if let Some(equiv) = union_find.get(*col) {
764 *col = *equiv;
765 }
766 }
767 key_set.sort();
768 key_set.dedup();
769 }
770 input.sort();
771 input.dedup();
772
773 // Expand out each key to each of its equivalent forms.
774 // Each instance of `col` can be replaced by any equivalent column.
775 // This has the potential to result in exponentially sized number of unique keys,
776 // and in the future we should probably maintain unique keys modulo equivalence.
777
778 // First, compute an inverse map from each representative
779 // column `sub` to all other equivalent columns `col`.
780 let mut subs = Vec::new();
781 for (col, sub) in union_find.iter().enumerate() {
782 if *sub != col {
783 assert!(*sub < col);
784 while subs.len() <= *sub {
785 subs.push(Vec::new());
786 }
787 subs[*sub].push(col);
788 }
789 }
790 // For each column, substitute for it in each occurrence.
791 let mut to_add = Vec::new();
792 for (col, subs) in subs.iter().enumerate() {
793 if !subs.is_empty() {
794 for key_set in input.iter() {
795 if key_set.contains(&col) {
796 let mut to_extend = key_set.clone();
797 to_extend.retain(|c| c != &col);
798 for sub in subs {
799 to_extend.push(*sub);
800 to_add.push(to_extend.clone());
801 to_extend.pop();
802 }
803 }
804 }
805 }
806 // No deduplication, as we cannot introduce duplicates.
807 input.append(&mut to_add);
808 }
809 for key_set in input.iter_mut() {
810 key_set.sort();
811 key_set.dedup();
812 }
813 }
814 input
815 }
816 Join { equivalences, .. } => {
817 // It is important the `new_from_input_arities` constructor is
818 // used. Otherwise, Materialize may potentially end up in an
819 // infinite loop.
820 let input_mapper = crate::JoinInputMapper::new_from_input_arities(input_arities);
821
822 input_mapper.global_keys(input_keys, equivalences)
823 }
824 Reduce { group_key, .. } => {
825 // The group key should form a key, but we might already have
826 // keys that are subsets of the group key, and should retain
827 // those instead, if so.
828 let mut result = Vec::new();
829 for key_set in input_keys.next().unwrap() {
830 if key_set
831 .iter()
832 .all(|k| group_key.contains(&MirScalarExpr::column(*k)))
833 {
834 result.push(
835 key_set
836 .iter()
837 .map(|i| {
838 group_key
839 .iter()
840 .position(|k| k == &MirScalarExpr::column(*i))
841 .unwrap()
842 })
843 .collect::<Vec<_>>(),
844 );
845 }
846 }
847 if result.is_empty() {
848 result.push((0..group_key.len()).collect());
849 }
850 result
851 }
852 TopK {
853 group_key, limit, ..
854 } => {
855 // If `limit` is `Some(1)` then the group key will become
856 // a unique key, as there will be only one record with that key.
857 let mut result = input_keys.next().unwrap().clone();
858 if limit.as_ref().and_then(|x| x.as_literal_int64()) == Some(1) {
859 result.push(group_key.clone())
860 }
861 result
862 }
863 Union { base, inputs } => {
864 // Generally, unions do not have any unique keys, because
865 // each input might duplicate some. However, there is at
866 // least one idiomatic structure that does preserve keys,
867 // which results from SQL aggregations that must populate
868 // absent records with default values. In that pattern,
869 // the union of one GET with its negation, which has first
870 // been subjected to a projection and map, we can remove
871 // their influence on the key structure.
872 //
873 // If there are A, B, each with a unique `key` such that
874 // we are looking at
875 //
876 // A.proj(set_containing_key) + (B - A.proj(key)).map(stuff)
877 //
878 // Then we can report `key` as a unique key.
879 //
880 // TODO: make unique key structure an optimization analysis
881 // rather than part of the type information.
882 // TODO: perhaps ensure that (above) A.proj(key) is a
883 // subset of B, as otherwise there are negative records
884 // and who knows what is true (not expected, but again
885 // who knows what the query plan might look like).
886
887 let arity = input_arities.next().unwrap();
888 let (base_projection, base_with_project_stripped) =
889 if let MirRelationExpr::Project { input, outputs } = &**base {
890 (outputs.clone(), &**input)
891 } else {
892 // A input without a project is equivalent to an input
893 // with the project being all columns in the input in order.
894 ((0..arity).collect::<Vec<_>>(), &**base)
895 };
896 let mut result = Vec::new();
897 if let MirRelationExpr::Get {
898 id: first_id,
899 typ: _,
900 ..
901 } = base_with_project_stripped
902 {
903 if inputs.len() == 1 {
904 if let MirRelationExpr::Map { input, .. } = &inputs[0] {
905 if let MirRelationExpr::Union { base, inputs } = &**input {
906 if inputs.len() == 1 {
907 if let Some((input, outputs)) = base.is_negated_project() {
908 if let MirRelationExpr::Get {
909 id: second_id,
910 typ: _,
911 ..
912 } = input
913 {
914 if first_id == second_id {
915 result.extend(
916 input_keys
917 .next()
918 .unwrap()
919 .into_iter()
920 .filter(|key| {
921 key.iter().all(|c| {
922 outputs.get(*c) == Some(c)
923 && base_projection.get(*c)
924 == Some(c)
925 })
926 })
927 .cloned(),
928 );
929 }
930 }
931 }
932 }
933 }
934 }
935 }
936 }
937 // Important: do not inherit keys of either input, as not unique.
938 result
939 }
940 };
941 keys.sort();
942 keys.dedup();
943 keys
944 }
945
946 /// The number of columns in the relation.
947 ///
948 /// This number is determined from the type, which is determined recursively
949 /// at non-trivial cost.
950 ///
951 /// The arity is computed incrementally with a recursive post-order
952 /// traversal, that accumulates the arities for the relations yet to be
953 /// visited in `arity_stack`.
954 pub fn arity(&self) -> usize {
955 let mut arity_stack = Vec::new();
956 self.visit_pre_post(
957 &mut |e: &MirRelationExpr| -> Option<Vec<&MirRelationExpr>> {
958 match &e {
959 MirRelationExpr::Let { body, .. } => {
960 // Do not traverse the value sub-graph, since it's not relevant for
961 // determining the arity of Let operators.
962 Some(vec![&*body])
963 }
964 MirRelationExpr::LetRec { body, .. } => {
965 // Do not traverse the value sub-graph, since it's not relevant for
966 // determining the arity of Let operators.
967 Some(vec![&*body])
968 }
969 MirRelationExpr::Project { .. } | MirRelationExpr::Reduce { .. } => {
970 // No further traversal is required; these operators know their arity.
971 Some(Vec::new())
972 }
973 _ => None,
974 }
975 },
976 &mut |e: &MirRelationExpr| {
977 match &e {
978 MirRelationExpr::Let { .. } => {
979 let body_arity = arity_stack.pop().unwrap();
980 arity_stack.push(0);
981 arity_stack.push(body_arity);
982 }
983 MirRelationExpr::LetRec { values, .. } => {
984 let body_arity = arity_stack.pop().unwrap();
985 arity_stack.extend(std::iter::repeat(0).take(values.len()));
986 arity_stack.push(body_arity);
987 }
988 MirRelationExpr::Project { .. } | MirRelationExpr::Reduce { .. } => {
989 arity_stack.push(0);
990 }
991 _ => {}
992 }
993 let num_inputs = e.num_inputs();
994 let input_arities = arity_stack.drain(arity_stack.len() - num_inputs..);
995 let arity = e.arity_with_input_arities(input_arities);
996 arity_stack.push(arity);
997 },
998 );
999 assert_eq!(arity_stack.len(), 1);
1000 arity_stack.pop().unwrap()
1001 }
1002
1003 /// Reports the arity of the relation given the schema of the input relations.
1004 ///
1005 /// `input_arities` is required to contain the arities for the input relations of
1006 /// the current relation in the same order as they are visited by `try_visit_children`
1007 /// method, even though not all may be used for computing the schema of the
1008 /// current relation. For example, `Let` expects two input types, one for the
1009 /// value relation and one for the body, in that order, but only the one for the
1010 /// body is used to determine the type of the `Let` relation.
1011 ///
1012 /// It is meant to be used during post-order traversals to compute arities
1013 /// incrementally.
1014 pub fn arity_with_input_arities<I>(&self, mut input_arities: I) -> usize
1015 where
1016 I: Iterator<Item = usize>,
1017 {
1018 use MirRelationExpr::*;
1019
1020 match self {
1021 Constant { rows: _, typ } => typ.arity(),
1022 Get { typ, .. } => typ.arity(),
1023 Let { .. } => {
1024 input_arities.next();
1025 input_arities.next().unwrap()
1026 }
1027 LetRec { values, .. } => {
1028 for _ in 0..values.len() {
1029 input_arities.next();
1030 }
1031 input_arities.next().unwrap()
1032 }
1033 Project { outputs, .. } => outputs.len(),
1034 Map { scalars, .. } => input_arities.next().unwrap() + scalars.len(),
1035 FlatMap { func, .. } => input_arities.next().unwrap() + func.output_arity(),
1036 Join { .. } => input_arities.sum(),
1037 Reduce {
1038 input: _,
1039 group_key,
1040 aggregates,
1041 ..
1042 } => group_key.len() + aggregates.len(),
1043 Filter { .. }
1044 | TopK { .. }
1045 | Negate { .. }
1046 | Threshold { .. }
1047 | Union { .. }
1048 | ArrangeBy { .. } => input_arities.next().unwrap(),
1049 }
1050 }
1051
1052 /// The number of child relations this relation has.
1053 pub fn num_inputs(&self) -> usize {
1054 let mut count = 0;
1055
1056 self.visit_children(|_| count += 1);
1057
1058 count
1059 }
1060
1061 /// Constructs a constant collection from specific rows and schema, where
1062 /// each row will have a multiplicity of one.
1063 pub fn constant(rows: Vec<Vec<Datum>>, typ: ReprRelationType) -> Self {
1064 let rows = rows.into_iter().map(|row| (row, Diff::ONE)).collect();
1065 MirRelationExpr::constant_diff(rows, typ)
1066 }
1067
1068 /// Constructs a constant collection from specific rows and schema, where
1069 /// each row can have an arbitrary multiplicity.
1070 pub fn constant_diff(rows: Vec<(Vec<Datum>, Diff)>, typ: ReprRelationType) -> Self {
1071 for (row, _diff) in &rows {
1072 for (datum, column_typ) in row.iter().zip_eq(typ.column_types.iter()) {
1073 assert!(
1074 datum.is_instance_of(column_typ),
1075 "Expected datum of type {:?}, got value {:?}",
1076 column_typ,
1077 datum
1078 );
1079 }
1080 }
1081 let rows = Ok(rows
1082 .into_iter()
1083 .map(move |(row, diff)| (Row::pack_slice(&row), diff))
1084 .collect());
1085 MirRelationExpr::Constant { rows, typ }
1086 }
1087
1088 /// If self is a constant, return the value and the type, otherwise `None`.
1089 /// Looks behind `ArrangeBy`s.
1090 pub fn as_const(&self) -> Option<(&Result<Vec<(Row, Diff)>, EvalError>, &ReprRelationType)> {
1091 match self {
1092 MirRelationExpr::Constant { rows, typ } => Some((rows, typ)),
1093 MirRelationExpr::ArrangeBy { input, .. } => input.as_const(),
1094 _ => None,
1095 }
1096 }
1097
1098 /// If self is a constant, mutably return the value and the type, otherwise `None`.
1099 /// Looks behind `ArrangeBy`s.
1100 pub fn as_const_mut(
1101 &mut self,
1102 ) -> Option<(
1103 &mut Result<Vec<(Row, Diff)>, EvalError>,
1104 &mut ReprRelationType,
1105 )> {
1106 match self {
1107 MirRelationExpr::Constant { rows, typ } => Some((rows, typ)),
1108 MirRelationExpr::ArrangeBy { input, .. } => input.as_const_mut(),
1109 _ => None,
1110 }
1111 }
1112
1113 /// If self is a constant error, return the error, otherwise `None`.
1114 /// Looks behind `ArrangeBy`s.
1115 pub fn as_const_err(&self) -> Option<&EvalError> {
1116 match self {
1117 MirRelationExpr::Constant { rows: Err(e), .. } => Some(e),
1118 MirRelationExpr::ArrangeBy { input, .. } => input.as_const_err(),
1119 _ => None,
1120 }
1121 }
1122
1123 /// Checks if `self` is the single element collection with no columns.
1124 pub fn is_constant_singleton(&self) -> bool {
1125 if let Some((Ok(rows), typ)) = self.as_const() {
1126 rows.len() == 1 && typ.column_types.len() == 0 && rows[0].1 == Diff::ONE
1127 } else {
1128 false
1129 }
1130 }
1131
1132 /// Constructs the expression for getting a local collection.
1133 pub fn local_get(id: LocalId, typ: ReprRelationType) -> Self {
1134 MirRelationExpr::Get {
1135 id: Id::Local(id),
1136 typ,
1137 access_strategy: AccessStrategy::UnknownOrLocal,
1138 }
1139 }
1140
1141 /// Constructs the expression for getting a global collection
1142 pub fn global_get(id: GlobalId, typ: ReprRelationType) -> Self {
1143 MirRelationExpr::Get {
1144 id: Id::Global(id),
1145 typ,
1146 access_strategy: AccessStrategy::UnknownOrLocal,
1147 }
1148 }
1149
1150 /// Retains only the columns specified by `output`.
1151 pub fn project(mut self, mut outputs: Vec<usize>) -> Self {
1152 if let MirRelationExpr::Project {
1153 outputs: columns, ..
1154 } = &mut self
1155 {
1156 // Update `outputs` to reference base columns of `input`.
1157 for column in outputs.iter_mut() {
1158 *column = columns[*column];
1159 }
1160 *columns = outputs;
1161 self
1162 } else {
1163 MirRelationExpr::Project {
1164 input: Box::new(self),
1165 outputs,
1166 }
1167 }
1168 }
1169
1170 /// Append to each row the results of applying elements of `scalar`.
1171 pub fn map(mut self, scalars: Vec<MirScalarExpr>) -> Self {
1172 if let MirRelationExpr::Map { scalars: s, .. } = &mut self {
1173 s.extend(scalars);
1174 self
1175 } else if !scalars.is_empty() {
1176 MirRelationExpr::Map {
1177 input: Box::new(self),
1178 scalars,
1179 }
1180 } else {
1181 self
1182 }
1183 }
1184
1185 /// Append to each row a single `scalar`.
1186 pub fn map_one(self, scalar: MirScalarExpr) -> Self {
1187 self.map(vec![scalar])
1188 }
1189
1190 /// Like `map`, but yields zero-or-more output rows per input row
1191 pub fn flat_map(self, func: TableFunc, exprs: Vec<MirScalarExpr>) -> Self {
1192 MirRelationExpr::FlatMap {
1193 input: Box::new(self),
1194 func,
1195 exprs,
1196 }
1197 }
1198
1199 /// Retain only the rows satisfying each of several predicates.
1200 pub fn filter<I>(mut self, predicates: I) -> Self
1201 where
1202 I: IntoIterator<Item = MirScalarExpr>,
1203 {
1204 // Extract existing predicates
1205 let mut new_predicates = if let MirRelationExpr::Filter { input, predicates } = self {
1206 self = *input;
1207 predicates
1208 } else {
1209 Vec::new()
1210 };
1211 // Normalize collection of predicates.
1212 new_predicates.extend(predicates);
1213 new_predicates.retain(|p| !p.is_literal_true());
1214 new_predicates.sort();
1215 new_predicates.dedup();
1216 // Introduce a `Filter` only if we have predicates.
1217 if !new_predicates.is_empty() {
1218 self = MirRelationExpr::Filter {
1219 input: Box::new(self),
1220 predicates: new_predicates,
1221 };
1222 }
1223
1224 self
1225 }
1226
1227 /// Form the Cartesian outer-product of rows in both inputs.
1228 pub fn product(mut self, right: Self) -> Self {
1229 if right.is_constant_singleton() {
1230 self
1231 } else if self.is_constant_singleton() {
1232 right
1233 } else if let MirRelationExpr::Join { inputs, .. } = &mut self {
1234 inputs.push(right);
1235 self
1236 } else {
1237 MirRelationExpr::join(vec![self, right], vec![])
1238 }
1239 }
1240
1241 /// Performs a relational equijoin among the input collections.
1242 ///
1243 /// The sequence `inputs` each describe different input collections, and the sequence `variables` describes
1244 /// equality constraints that some of their columns must satisfy. Each element in `variable` describes a set
1245 /// of pairs `(input_index, column_index)` where every value described by that set must be equal.
1246 ///
1247 /// For example, the pair `(input, column)` indexes into `inputs[input][column]`, extracting the `input`th
1248 /// input collection and for each row examining its `column`th column.
1249 ///
1250 /// # Example
1251 ///
1252 /// ```rust
1253 /// use mz_repr::{Datum, SqlColumnType, ReprRelationType, ReprScalarType};
1254 /// use mz_expr::MirRelationExpr;
1255 ///
1256 /// // A common schema for each input.
1257 /// let schema = ReprRelationType::new(vec![
1258 /// ReprScalarType::Int32.nullable(false),
1259 /// ReprScalarType::Int32.nullable(false),
1260 /// ]);
1261 ///
1262 /// // the specific data are not important here.
1263 /// let data = vec![Datum::Int32(0), Datum::Int32(1)];
1264 ///
1265 /// // Three collections that could have been different.
1266 /// let input0 = MirRelationExpr::constant(vec![data.clone()], schema.clone());
1267 /// let input1 = MirRelationExpr::constant(vec![data.clone()], schema.clone());
1268 /// let input2 = MirRelationExpr::constant(vec![data.clone()], schema.clone());
1269 ///
1270 /// // Join the three relations looking for triangles, like so.
1271 /// //
1272 /// // Output(A,B,C) := Input0(A,B), Input1(B,C), Input2(A,C)
1273 /// let joined = MirRelationExpr::join(
1274 /// vec![input0, input1, input2],
1275 /// vec![
1276 /// vec![(0,0), (2,0)], // fields A of inputs 0 and 2.
1277 /// vec![(0,1), (1,0)], // fields B of inputs 0 and 1.
1278 /// vec![(1,1), (2,1)], // fields C of inputs 1 and 2.
1279 /// ],
1280 /// );
1281 ///
1282 /// // Technically the above produces `Output(A,B,B,C,A,C)` because the columns are concatenated.
1283 /// // A projection resolves this and produces the correct output.
1284 /// let result = joined.project(vec![0, 1, 3]);
1285 /// ```
1286 pub fn join(inputs: Vec<MirRelationExpr>, variables: Vec<Vec<(usize, usize)>>) -> Self {
1287 let input_mapper = join_input_mapper::JoinInputMapper::new(&inputs);
1288
1289 let equivalences = variables
1290 .into_iter()
1291 .map(|vs| {
1292 vs.into_iter()
1293 .map(|(r, c)| input_mapper.map_expr_to_global(MirScalarExpr::column(c), r))
1294 .collect::<Vec<_>>()
1295 })
1296 .collect::<Vec<_>>();
1297
1298 Self::join_scalars(inputs, equivalences)
1299 }
1300
1301 /// Constructs a join operator from inputs and required-equal scalar expressions.
1302 pub fn join_scalars(
1303 mut inputs: Vec<MirRelationExpr>,
1304 equivalences: Vec<Vec<MirScalarExpr>>,
1305 ) -> Self {
1306 // Remove all constant inputs that are the identity for join.
1307 // They neither introduce nor modify any column references.
1308 inputs.retain(|i| !i.is_constant_singleton());
1309 MirRelationExpr::Join {
1310 inputs,
1311 equivalences,
1312 implementation: JoinImplementation::Unimplemented,
1313 }
1314 }
1315
1316 /// Perform a key-wise reduction / aggregation.
1317 ///
1318 /// The `group_key` argument indicates columns in the input collection that should
1319 /// be grouped, and `aggregates` lists aggregation functions each of which produces
1320 /// one output column in addition to the keys.
1321 pub fn reduce(
1322 self,
1323 group_key: Vec<usize>,
1324 aggregates: Vec<AggregateExpr>,
1325 expected_group_size: Option<u64>,
1326 ) -> Self {
1327 MirRelationExpr::Reduce {
1328 input: Box::new(self),
1329 group_key: group_key.into_iter().map(MirScalarExpr::column).collect(),
1330 aggregates,
1331 monotonic: false,
1332 expected_group_size,
1333 }
1334 }
1335
1336 /// Perform a key-wise reduction order by and limit.
1337 ///
1338 /// The `group_key` argument indicates columns in the input collection that should
1339 /// be grouped, the `order_key` argument indicates columns that should be further
1340 /// used to order records within groups, and the `limit` argument constrains the
1341 /// total number of records that should be produced in each group.
1342 pub fn top_k(
1343 self,
1344 group_key: Vec<usize>,
1345 order_key: Vec<ColumnOrder>,
1346 limit: Option<MirScalarExpr>,
1347 offset: usize,
1348 expected_group_size: Option<u64>,
1349 ) -> Self {
1350 MirRelationExpr::TopK {
1351 input: Box::new(self),
1352 group_key,
1353 order_key,
1354 limit,
1355 offset,
1356 expected_group_size,
1357 monotonic: false,
1358 }
1359 }
1360
1361 /// Negates the occurrences of each row.
1362 pub fn negate(self) -> Self {
1363 if let MirRelationExpr::Negate { input } = self {
1364 *input
1365 } else {
1366 MirRelationExpr::Negate {
1367 input: Box::new(self),
1368 }
1369 }
1370 }
1371
1372 /// Removes all but the first occurrence of each row.
1373 pub fn distinct(self) -> Self {
1374 let arity = self.arity();
1375 self.distinct_by((0..arity).collect())
1376 }
1377
1378 /// Removes all but the first occurrence of each key. Columns not included
1379 /// in the `group_key` are discarded.
1380 pub fn distinct_by(self, group_key: Vec<usize>) -> Self {
1381 self.reduce(group_key, vec![], None)
1382 }
1383
1384 /// Discards rows with a negative frequency.
1385 pub fn threshold(self) -> Self {
1386 if let MirRelationExpr::Threshold { .. } = &self {
1387 self
1388 } else {
1389 MirRelationExpr::Threshold {
1390 input: Box::new(self),
1391 }
1392 }
1393 }
1394
1395 /// Unions together any number inputs.
1396 ///
1397 /// If `inputs` is empty, then an empty relation of type `typ` is
1398 /// constructed.
1399 pub fn union_many(mut inputs: Vec<Self>, typ: ReprRelationType) -> Self {
1400 // Deconstruct `inputs` as `Union`s and reconstitute.
1401 let mut flat_inputs = Vec::with_capacity(inputs.len());
1402 for input in inputs {
1403 if let MirRelationExpr::Union { base, inputs } = input {
1404 flat_inputs.push(*base);
1405 flat_inputs.extend(inputs);
1406 } else {
1407 flat_inputs.push(input);
1408 }
1409 }
1410 inputs = flat_inputs;
1411 if inputs.len() == 0 {
1412 MirRelationExpr::Constant {
1413 rows: Ok(vec![]),
1414 typ,
1415 }
1416 } else if inputs.len() == 1 {
1417 inputs.into_element()
1418 } else {
1419 MirRelationExpr::Union {
1420 base: Box::new(inputs.remove(0)),
1421 inputs,
1422 }
1423 }
1424 }
1425
1426 /// Produces one collection where each row is present with the sum of its frequencies in each input.
1427 pub fn union(self, other: Self) -> Self {
1428 // Deconstruct `self` and `other` as `Union`s and reconstitute.
1429 let mut flat_inputs = Vec::with_capacity(2);
1430 if let MirRelationExpr::Union { base, inputs } = self {
1431 flat_inputs.push(*base);
1432 flat_inputs.extend(inputs);
1433 } else {
1434 flat_inputs.push(self);
1435 }
1436 if let MirRelationExpr::Union { base, inputs } = other {
1437 flat_inputs.push(*base);
1438 flat_inputs.extend(inputs);
1439 } else {
1440 flat_inputs.push(other);
1441 }
1442
1443 MirRelationExpr::Union {
1444 base: Box::new(flat_inputs.remove(0)),
1445 inputs: flat_inputs,
1446 }
1447 }
1448
1449 /// Arranges the collection by the specified columns
1450 pub fn arrange_by(self, keys: &[Vec<MirScalarExpr>]) -> Self {
1451 MirRelationExpr::ArrangeBy {
1452 input: Box::new(self),
1453 keys: keys.to_owned(),
1454 }
1455 }
1456
1457 /// Indicates if this is a constant empty collection.
1458 ///
1459 /// A false value does not mean the collection is known to be non-empty,
1460 /// only that we cannot currently determine that it is statically empty.
1461 pub fn is_empty(&self) -> bool {
1462 if let Some((Ok(rows), ..)) = self.as_const() {
1463 rows.is_empty()
1464 } else {
1465 false
1466 }
1467 }
1468
1469 /// If the expression is a negated project, return the input and the projection.
1470 pub fn is_negated_project(&self) -> Option<(&MirRelationExpr, &[usize])> {
1471 if let MirRelationExpr::Negate { input } = self {
1472 if let MirRelationExpr::Project { input, outputs } = &**input {
1473 return Some((&**input, outputs));
1474 }
1475 }
1476 if let MirRelationExpr::Project { input, outputs } = self {
1477 if let MirRelationExpr::Negate { input } = &**input {
1478 return Some((&**input, outputs));
1479 }
1480 }
1481 None
1482 }
1483
1484 /// Pretty-print this [MirRelationExpr] to a string.
1485 pub fn pretty(&self) -> String {
1486 let config = ExplainConfig::default();
1487 self.debug_explain(&config, None)
1488 }
1489
1490 /// Pretty-print this [MirRelationExpr] to a string using a custom
1491 /// [ExplainConfig] and an optionally provided [ExprHumanizer].
1492 /// This is intended for debugging and tests, not users.
1493 pub fn debug_explain(
1494 &self,
1495 config: &ExplainConfig,
1496 humanizer: Option<&dyn ExprHumanizer>,
1497 ) -> String {
1498 text_string_at(self, || PlanRenderingContext {
1499 indent: Indent::default(),
1500 humanizer: humanizer.unwrap_or(&DummyHumanizer),
1501 annotations: BTreeMap::default(),
1502 config,
1503 ambiguous_ids: BTreeSet::default(),
1504 })
1505 }
1506
1507 /// Take ownership of `self`, leaving an empty `MirRelationExpr::Constant` with the optionally
1508 /// given scalar types. The given scalar types should be `base_eq` with the types that `typ()`
1509 /// would find. Keys and nullability are ignored in the given `SqlRelationType`, and instead we set
1510 /// the best possible key and nullability, since we are making an empty collection.
1511 ///
1512 /// If `typ` is not given, then this calls `.typ()` (which is possibly expensive) to determine
1513 /// the correct type.
1514 pub fn take_safely(&mut self, typ: Option<ReprRelationType>) -> MirRelationExpr {
1515 if let Some(typ) = &typ {
1516 let self_typ = self.typ();
1517 soft_assert_no_log!(
1518 self_typ
1519 .column_types
1520 .iter()
1521 .zip_eq(typ.column_types.iter())
1522 .all(|(t1, t2)| t1.scalar_type == t2.scalar_type)
1523 );
1524 }
1525 let mut typ = typ.unwrap_or_else(|| self.typ());
1526 typ.keys = vec![vec![]];
1527 for ct in typ.column_types.iter_mut() {
1528 ct.nullable = false;
1529 }
1530 std::mem::replace(
1531 self,
1532 MirRelationExpr::Constant {
1533 rows: Ok(vec![]),
1534 typ,
1535 },
1536 )
1537 }
1538
1539 /// Take ownership of `self`, leaving an empty `MirRelationExpr::Constant` with the given scalar
1540 /// types. Nullability is ignored in the given `SqlColumnType`s, and instead we set the best
1541 /// possible nullability, since we are making an empty collection.
1542 pub fn take_safely_with_sql_col_types(&mut self, typ: Vec<SqlColumnType>) -> MirRelationExpr {
1543 self.take_safely(Some(ReprRelationType::from(&SqlRelationType::new(typ))))
1544 }
1545
1546 /// Like [`Self::take_safely_with_col_types`], but accepts `Vec<ReprColumnType>`.
1547 ///
1548 /// This is the preferred entry point for optimizer transforms, where repr
1549 /// types are the native currency. Internally converts to [`SqlColumnType`]
1550 /// and delegates to [`Self::take_safely_with_col_types`].
1551 pub fn take_safely_with_col_types(&mut self, typ: Vec<ReprColumnType>) -> MirRelationExpr {
1552 self.take_safely(Some(ReprRelationType::new(typ)))
1553 }
1554
1555 /// Take ownership of `self`, leaving an empty `MirRelationExpr::Constant` with an **incorrect** type.
1556 ///
1557 /// This should only be used if `self` is about to be dropped or otherwise overwritten.
1558 pub fn take_dangerous(&mut self) -> MirRelationExpr {
1559 let empty = MirRelationExpr::Constant {
1560 rows: Ok(vec![]),
1561 typ: ReprRelationType::new(Vec::new()),
1562 };
1563 std::mem::replace(self, empty)
1564 }
1565
1566 /// Replaces `self` with some logic applied to `self`.
1567 pub fn replace_using<F>(&mut self, logic: F)
1568 where
1569 F: FnOnce(MirRelationExpr) -> MirRelationExpr,
1570 {
1571 let empty = MirRelationExpr::Constant {
1572 rows: Ok(vec![]),
1573 typ: ReprRelationType::new(Vec::new()),
1574 };
1575 let expr = std::mem::replace(self, empty);
1576 *self = logic(expr);
1577 }
1578
1579 /// Store `self` in a `Let` and pass the corresponding `Get` to `body`.
1580 pub fn let_in<Body, E>(self, id_gen: &mut IdGen, body: Body) -> Result<MirRelationExpr, E>
1581 where
1582 Body: FnOnce(&mut IdGen, MirRelationExpr) -> Result<MirRelationExpr, E>,
1583 {
1584 if let MirRelationExpr::Get { .. } = self {
1585 // already done
1586 body(id_gen, self)
1587 } else {
1588 let id = LocalId::new(id_gen.allocate_id());
1589 let get = MirRelationExpr::Get {
1590 id: Id::Local(id),
1591 typ: self.typ(),
1592 access_strategy: AccessStrategy::UnknownOrLocal,
1593 };
1594 let body = (body)(id_gen, get)?;
1595 Ok(MirRelationExpr::Let {
1596 id,
1597 value: Box::new(self),
1598 body: Box::new(body),
1599 })
1600 }
1601 }
1602
1603 /// Return every row in `self` that does not have a matching row in the first columns of `keys_and_values`, using `default` to fill in the remaining columns
1604 /// (If `default` is a row of nulls, this is the 'outer' part of LEFT OUTER JOIN)
1605 pub fn anti_lookup<E>(
1606 self,
1607 id_gen: &mut IdGen,
1608 keys_and_values: MirRelationExpr,
1609 default: Vec<(Datum, ReprScalarType)>,
1610 ) -> Result<MirRelationExpr, E> {
1611 let (data, column_types): (Vec<_>, Vec<_>) = default
1612 .into_iter()
1613 .map(|(datum, scalar_type)| {
1614 (
1615 datum,
1616 ReprColumnType {
1617 scalar_type,
1618 nullable: datum.is_null(),
1619 },
1620 )
1621 })
1622 .unzip();
1623 assert_eq!(keys_and_values.arity() - self.arity(), data.len());
1624 self.let_in(id_gen, |_id_gen, get_keys| {
1625 let get_keys_arity = get_keys.arity();
1626 Ok(MirRelationExpr::join(
1627 vec![
1628 // all the missing keys (with count 1)
1629 keys_and_values
1630 .distinct_by((0..get_keys_arity).collect())
1631 .negate()
1632 .union(get_keys.clone().distinct()),
1633 // join with keys to get the correct counts
1634 get_keys.clone(),
1635 ],
1636 (0..get_keys_arity).map(|i| vec![(0, i), (1, i)]).collect(),
1637 )
1638 // get rid of the extra copies of columns from keys
1639 .project((0..get_keys_arity).collect())
1640 // This join is logically equivalent to
1641 // `.map(<default_expr>)`, but using a join allows for
1642 // potential predicate pushdown and elision in the
1643 // optimizer.
1644 .product(MirRelationExpr::constant(
1645 vec![data],
1646 ReprRelationType::new(column_types),
1647 )))
1648 })
1649 }
1650
1651 /// Return:
1652 /// * every row in keys_and_values
1653 /// * every row in `self` that does not have a matching row in the first columns of
1654 /// `keys_and_values`, using `default` to fill in the remaining columns
1655 /// (This is LEFT OUTER JOIN if:
1656 /// 1) `default` is a row of null
1657 /// 2) matching rows in `keys_and_values` and `self` have the same multiplicity.)
1658 pub fn lookup<E>(
1659 self,
1660 id_gen: &mut IdGen,
1661 keys_and_values: MirRelationExpr,
1662 default: Vec<(Datum<'static>, ReprScalarType)>,
1663 ) -> Result<MirRelationExpr, E> {
1664 keys_and_values.let_in(id_gen, |id_gen, get_keys_and_values| {
1665 Ok(get_keys_and_values.clone().union(self.anti_lookup(
1666 id_gen,
1667 get_keys_and_values,
1668 default,
1669 )?))
1670 })
1671 }
1672
1673 /// True iff the expression contains a `NullaryFunc::MzLogicalTimestamp`.
1674 pub fn contains_temporal(&self) -> bool {
1675 let mut contains = false;
1676 self.visit_scalars(&mut |e| contains = contains || e.contains_temporal());
1677 contains
1678 }
1679
1680 /// Fallible visitor for the [`MirScalarExpr`]s directly owned by this relation expression.
1681 ///
1682 /// The `f` visitor should not recursively descend into owned [`MirRelationExpr`]s.
1683 pub fn try_visit_scalars_mut1<F, E>(&mut self, f: &mut F) -> Result<(), E>
1684 where
1685 F: FnMut(&mut MirScalarExpr) -> Result<(), E>,
1686 {
1687 use MirRelationExpr::*;
1688 match self {
1689 Map { scalars, .. } => {
1690 for s in scalars {
1691 f(s)?;
1692 }
1693 }
1694 Filter { predicates, .. } => {
1695 for p in predicates {
1696 f(p)?;
1697 }
1698 }
1699 FlatMap { exprs, .. } => {
1700 for expr in exprs {
1701 f(expr)?;
1702 }
1703 }
1704 Join {
1705 inputs: _,
1706 equivalences,
1707 implementation,
1708 } => {
1709 for equivalence in equivalences {
1710 for expr in equivalence {
1711 f(expr)?;
1712 }
1713 }
1714 match implementation {
1715 JoinImplementation::Differential((_, start_key, _), order) => {
1716 if let Some(start_key) = start_key {
1717 for k in start_key {
1718 f(k)?;
1719 }
1720 }
1721 for (_, lookup_key, _) in order {
1722 for k in lookup_key {
1723 f(k)?;
1724 }
1725 }
1726 }
1727 JoinImplementation::DeltaQuery(paths) => {
1728 for path in paths {
1729 for (_, lookup_key, _) in path {
1730 for k in lookup_key {
1731 f(k)?;
1732 }
1733 }
1734 }
1735 }
1736 JoinImplementation::IndexedFilter(_coll_id, _idx_id, index_key, _) => {
1737 for k in index_key {
1738 f(k)?;
1739 }
1740 }
1741 JoinImplementation::Unimplemented => {} // No scalar exprs
1742 }
1743 }
1744 ArrangeBy { keys, .. } => {
1745 for key in keys {
1746 for s in key {
1747 f(s)?;
1748 }
1749 }
1750 }
1751 Reduce {
1752 group_key,
1753 aggregates,
1754 ..
1755 } => {
1756 for s in group_key {
1757 f(s)?;
1758 }
1759 for agg in aggregates {
1760 f(&mut agg.expr)?;
1761 }
1762 }
1763 TopK { limit, .. } => {
1764 if let Some(s) = limit {
1765 f(s)?;
1766 }
1767 }
1768 Constant { .. }
1769 | Get { .. }
1770 | Let { .. }
1771 | LetRec { .. }
1772 | Project { .. }
1773 | Negate { .. }
1774 | Threshold { .. }
1775 | Union { .. } => (),
1776 }
1777 Ok(())
1778 }
1779
1780 /// Fallible mutable visitor for the [`MirScalarExpr`]s in the [`MirRelationExpr`] subtree
1781 /// rooted at `self`.
1782 ///
1783 /// Note that this does not recurse into [`MirRelationExpr`] subtrees within [`MirScalarExpr`]
1784 /// nodes.
1785 pub fn try_visit_scalars_mut<F, E>(&mut self, f: &mut F) -> Result<(), E>
1786 where
1787 F: FnMut(&mut MirScalarExpr) -> Result<(), E>,
1788 {
1789 self.try_visit_mut_post(&mut |expr| expr.try_visit_scalars_mut1(f))
1790 }
1791
1792 /// Infallible mutable visitor for the [`MirScalarExpr`]s in the [`MirRelationExpr`] subtree
1793 /// rooted at `self`.
1794 ///
1795 /// Note that this does not recurse into [`MirRelationExpr`] subtrees within [`MirScalarExpr`]
1796 /// nodes.
1797 pub fn visit_scalars_mut<F>(&mut self, f: &mut F)
1798 where
1799 F: FnMut(&mut MirScalarExpr),
1800 {
1801 self.try_visit_scalars_mut(&mut |s| {
1802 f(s);
1803 Ok::<_, RecursionLimitError>(())
1804 })
1805 .expect("Unexpected error in `visit_scalars_mut` call");
1806 }
1807
1808 /// Fallible visitor for the [`MirScalarExpr`]s directly owned by this relation expression.
1809 ///
1810 /// The `f` visitor should not recursively descend into owned [`MirRelationExpr`]s.
1811 pub fn try_visit_scalars_1<F, E>(&self, f: &mut F) -> Result<(), E>
1812 where
1813 F: FnMut(&MirScalarExpr) -> Result<(), E>,
1814 {
1815 use MirRelationExpr::*;
1816 match self {
1817 Map { scalars, .. } => {
1818 for s in scalars {
1819 f(s)?;
1820 }
1821 }
1822 Filter { predicates, .. } => {
1823 for p in predicates {
1824 f(p)?;
1825 }
1826 }
1827 FlatMap { exprs, .. } => {
1828 for expr in exprs {
1829 f(expr)?;
1830 }
1831 }
1832 Join {
1833 inputs: _,
1834 equivalences,
1835 implementation,
1836 } => {
1837 for equivalence in equivalences {
1838 for expr in equivalence {
1839 f(expr)?;
1840 }
1841 }
1842 match implementation {
1843 JoinImplementation::Differential((_, start_key, _), order) => {
1844 if let Some(start_key) = start_key {
1845 for k in start_key {
1846 f(k)?;
1847 }
1848 }
1849 for (_, lookup_key, _) in order {
1850 for k in lookup_key {
1851 f(k)?;
1852 }
1853 }
1854 }
1855 JoinImplementation::DeltaQuery(paths) => {
1856 for path in paths {
1857 for (_, lookup_key, _) in path {
1858 for k in lookup_key {
1859 f(k)?;
1860 }
1861 }
1862 }
1863 }
1864 JoinImplementation::IndexedFilter(_coll_id, _idx_id, index_key, _) => {
1865 for k in index_key {
1866 f(k)?;
1867 }
1868 }
1869 JoinImplementation::Unimplemented => {} // No scalar exprs
1870 }
1871 }
1872 ArrangeBy { keys, .. } => {
1873 for key in keys {
1874 for s in key {
1875 f(s)?;
1876 }
1877 }
1878 }
1879 Reduce {
1880 group_key,
1881 aggregates,
1882 ..
1883 } => {
1884 for s in group_key {
1885 f(s)?;
1886 }
1887 for agg in aggregates {
1888 f(&agg.expr)?;
1889 }
1890 }
1891 TopK { limit, .. } => {
1892 if let Some(s) = limit {
1893 f(s)?;
1894 }
1895 }
1896 Constant { .. }
1897 | Get { .. }
1898 | Let { .. }
1899 | LetRec { .. }
1900 | Project { .. }
1901 | Negate { .. }
1902 | Threshold { .. }
1903 | Union { .. } => (),
1904 }
1905 Ok(())
1906 }
1907
1908 /// Fallible immutable visitor for the [`MirScalarExpr`]s in the [`MirRelationExpr`] subtree
1909 /// rooted at `self`.
1910 ///
1911 /// Note that this does not recurse into [`MirRelationExpr`] subtrees within [`MirScalarExpr`]
1912 /// nodes.
1913 pub fn try_visit_scalars<F, E>(&self, f: &mut F) -> Result<(), E>
1914 where
1915 F: FnMut(&MirScalarExpr) -> Result<(), E>,
1916 {
1917 self.try_visit_post(&mut |expr| expr.try_visit_scalars_1(f))
1918 }
1919
1920 /// Infallible immutable visitor for the [`MirScalarExpr`]s in the [`MirRelationExpr`] subtree
1921 /// rooted at `self`.
1922 ///
1923 /// Note that this does not recurse into [`MirRelationExpr`] subtrees within [`MirScalarExpr`]
1924 /// nodes.
1925 pub fn visit_scalars<F>(&self, f: &mut F)
1926 where
1927 F: FnMut(&MirScalarExpr),
1928 {
1929 self.try_visit_scalars(&mut |s| {
1930 f(s);
1931 Ok::<_, RecursionLimitError>(())
1932 })
1933 .expect("Unexpected error in `visit_scalars` call");
1934 }
1935
1936 /// Clears the contents of `self` even if it's so deep that simply dropping it would cause a
1937 /// stack overflow in `drop_in_place`.
1938 ///
1939 /// Leaves `self` in an unusable state, so this should only be used if `self` is about to be
1940 /// dropped or otherwise overwritten.
1941 pub fn destroy_carefully(&mut self) {
1942 let mut todo = vec![self.take_dangerous()];
1943 while let Some(mut expr) = todo.pop() {
1944 for child in expr.children_mut() {
1945 todo.push(child.take_dangerous());
1946 }
1947 }
1948 }
1949
1950 /// Computes the size (total number of nodes) and maximum depth of a MirRelationExpr for
1951 /// debug printing purposes.
1952 pub fn debug_size_and_depth(&self) -> (usize, usize) {
1953 let mut size = 0;
1954 let mut max_depth = 0;
1955 let mut todo = vec![(self, 1)];
1956 while let Some((expr, depth)) = todo.pop() {
1957 size += 1;
1958 max_depth = max(max_depth, depth);
1959 todo.extend(expr.children().map(|c| (c, depth + 1)));
1960 }
1961 (size, max_depth)
1962 }
1963
1964 /// The MirRelationExpr is considered potentially expensive if and only if
1965 /// at least one of the following conditions is true:
1966 ///
1967 /// - It contains at least one MirScalarExpr with a function call.
1968 /// - It contains at least one FlatMap or a Reduce operator.
1969 /// - We run into a RecursionLimitError while analyzing the expression.
1970 ///
1971 /// !!!WARNING!!!: this method has an HirRelationExpr counterpart. The two
1972 /// should be kept in sync w.r.t. HIR ⇒ MIR lowering!
1973 pub fn could_run_expensive_function(&self) -> bool {
1974 let mut result = false;
1975 use MirRelationExpr::*;
1976 use MirScalarExpr::*;
1977 if let Err(_) = self.try_visit_scalars::<_, RecursionLimitError>(&mut |scalar| {
1978 result |= match scalar {
1979 Column(_, _) | Literal(_, _) | CallUnmaterializable(_) | If { .. } => false,
1980 // Function calls are considered expensive
1981 CallUnary { .. } | CallBinary { .. } | CallVariadic { .. } => true,
1982 };
1983 Ok(())
1984 }) {
1985 // Conservatively set `true` if on RecursionLimitError.
1986 result = true;
1987 }
1988 self.visit_pre(|e: &MirRelationExpr| {
1989 // FlatMap has a table function; Reduce has an aggregate function.
1990 // Other constructs use MirScalarExpr to run a function
1991 result |= matches!(e, FlatMap { .. } | Reduce { .. });
1992 });
1993 result
1994 }
1995
1996 /// Hash to an u64 using Rust's default Hasher. (Which is a somewhat slower, but better Hasher
1997 /// than what `Hashable::hashed` would give us.)
1998 pub fn hash_to_u64(&self) -> u64 {
1999 let mut h = DefaultHasher::new();
2000 self.hash(&mut h);
2001 h.finish()
2002 }
2003}
2004
2005// `LetRec` helpers
2006impl MirRelationExpr {
2007 /// True when `expr` contains a `LetRec` AST node.
2008 pub fn is_recursive(self: &MirRelationExpr) -> bool {
2009 let mut worklist = vec![self];
2010 while let Some(expr) = worklist.pop() {
2011 if let MirRelationExpr::LetRec { .. } = expr {
2012 return true;
2013 }
2014 worklist.extend(expr.children());
2015 }
2016 false
2017 }
2018
2019 /// Return the number of sub-expressions in the tree (including self).
2020 pub fn size(&self) -> usize {
2021 let mut size = 0;
2022 self.visit_pre(|_| size += 1);
2023 size
2024 }
2025
2026 /// Given the ids and values of a LetRec, it computes the subset of ids that are used across
2027 /// iterations. These are those ids that have a reference before they are defined, when reading
2028 /// all the bindings in order.
2029 ///
2030 /// For example:
2031 /// ```SQL
2032 /// WITH MUTUALLY RECURSIVE
2033 /// x(...) AS f(z),
2034 /// y(...) AS g(x),
2035 /// z(...) AS h(y)
2036 /// ...;
2037 /// ```
2038 /// Here, only `z` is returned, because `x` and `y` are referenced only within the same
2039 /// iteration.
2040 ///
2041 /// Note that if a binding references itself, that is also returned.
2042 pub fn recursive_ids(ids: &[LocalId], values: &[MirRelationExpr]) -> BTreeSet<LocalId> {
2043 let mut used_across_iterations = BTreeSet::new();
2044 let mut defined = BTreeSet::new();
2045 for (binding_id, value) in itertools::zip_eq(ids.iter(), values.iter()) {
2046 value.visit_pre(|expr| {
2047 if let MirRelationExpr::Get {
2048 id: Local(get_id), ..
2049 } = expr
2050 {
2051 // If we haven't seen a definition for it yet, then this will refer
2052 // to the previous iteration.
2053 // The `ids.contains` part of the condition is needed to exclude
2054 // those ids that are not really in this LetRec, but either an inner
2055 // or outer one.
2056 if !defined.contains(get_id) && ids.contains(get_id) {
2057 used_across_iterations.insert(*get_id);
2058 }
2059 }
2060 });
2061 defined.insert(*binding_id);
2062 }
2063 used_across_iterations
2064 }
2065
2066 /// Replaces `LetRec` nodes with a stack of `Let` nodes.
2067 ///
2068 /// In each `Let` binding, uses of `Get` in `value` that are not at strictly greater
2069 /// identifiers are rewritten to be the constant collection.
2070 /// This makes the computation perform exactly "one" iteration.
2071 ///
2072 /// This was used only temporarily while developing `LetRec`.
2073 pub fn make_nonrecursive(self: &mut MirRelationExpr) {
2074 let mut deadlist = BTreeSet::new();
2075 let mut worklist = vec![self];
2076 while let Some(expr) = worklist.pop() {
2077 if let MirRelationExpr::LetRec {
2078 ids,
2079 values,
2080 limits: _,
2081 body,
2082 } = expr
2083 {
2084 let ids_values = values
2085 .drain(..)
2086 .zip_eq(ids)
2087 .map(|(value, id)| (*id, value))
2088 .collect::<Vec<_>>();
2089 *expr = body.take_dangerous();
2090 for (id, mut value) in ids_values.into_iter().rev() {
2091 // Remove references to potentially recursive identifiers.
2092 deadlist.insert(id);
2093 value.visit_pre_mut(|e| {
2094 if let MirRelationExpr::Get {
2095 id: crate::Id::Local(id),
2096 typ,
2097 ..
2098 } = e
2099 {
2100 let typ = typ.clone();
2101 if deadlist.contains(id) {
2102 e.take_safely(Some(typ));
2103 }
2104 }
2105 });
2106 *expr = MirRelationExpr::Let {
2107 id,
2108 value: Box::new(value),
2109 body: Box::new(expr.take_dangerous()),
2110 };
2111 }
2112 worklist.push(expr);
2113 } else {
2114 worklist.extend(expr.children_mut().rev());
2115 }
2116 }
2117 }
2118
2119 /// For each Id `id'` referenced in `expr`, if it is larger or equal than `id`, then record in
2120 /// `expire_whens` that when `id'` is redefined, then we should expire the information that
2121 /// we are holding about `id`. Call `do_expirations` with `expire_whens` at each Id
2122 /// redefinition.
2123 ///
2124 /// IMPORTANT: Relies on the numbering of Ids to be what `renumber_bindings` gives.
2125 pub fn collect_expirations(
2126 id: LocalId,
2127 expr: &MirRelationExpr,
2128 expire_whens: &mut BTreeMap<LocalId, Vec<LocalId>>,
2129 ) {
2130 expr.visit_pre(|e| {
2131 if let MirRelationExpr::Get {
2132 id: Id::Local(referenced_id),
2133 ..
2134 } = e
2135 {
2136 // The following check needs `renumber_bindings` to have run recently
2137 if referenced_id >= &id {
2138 expire_whens
2139 .entry(*referenced_id)
2140 .or_insert_with(Vec::new)
2141 .push(id);
2142 }
2143 }
2144 });
2145 }
2146
2147 /// Call this function when `id` is redefined. It modifies `id_infos` by removing information
2148 /// about such Ids whose information depended on the earlier definition of `id`, according to
2149 /// `expire_whens`. Also modifies `expire_whens`: it removes the currently processed entry.
2150 pub fn do_expirations<I>(
2151 redefined_id: LocalId,
2152 expire_whens: &mut BTreeMap<LocalId, Vec<LocalId>>,
2153 id_infos: &mut BTreeMap<LocalId, I>,
2154 ) -> Vec<(LocalId, I)> {
2155 let mut expired_infos = Vec::new();
2156 if let Some(expirations) = expire_whens.remove(&redefined_id) {
2157 for expired_id in expirations.into_iter() {
2158 if let Some(offer) = id_infos.remove(&expired_id) {
2159 expired_infos.push((expired_id, offer));
2160 }
2161 }
2162 }
2163 expired_infos
2164 }
2165}
2166/// Augment non-nullability of columns, by observing either
2167/// 1. Predicates that explicitly test for null values, and
2168/// 2. Columns that if null would make a predicate be null.
2169pub fn non_nullable_columns(predicates: &[MirScalarExpr]) -> BTreeSet<usize> {
2170 let mut nonnull_required_columns = BTreeSet::new();
2171 for predicate in predicates {
2172 // Add any columns that being null would force the predicate to be null.
2173 // Should that happen, the row would be discarded.
2174 predicate.non_null_requirements(&mut nonnull_required_columns);
2175
2176 /*
2177 Test for explicit checks that a column is non-null.
2178
2179 This analysis is ad hoc, and will miss things:
2180
2181 materialize=> create table a(x int, y int);
2182 CREATE TABLE
2183 materialize=> explain with(types) select x from a where (y=x and y is not null) or x is not null;
2184 Optimized Plan
2185 --------------------------------------------------------------------------------------------------------
2186 Explained Query: +
2187 Project (#0) // { types: "(integer?)" } +
2188 Filter ((#0) IS NOT NULL OR ((#1) IS NOT NULL AND (#0 = #1))) // { types: "(integer?, integer?)" }+
2189 Get materialize.public.a // { types: "(integer?, integer?)" } +
2190 +
2191 Source materialize.public.a +
2192 filter=(((#0) IS NOT NULL OR ((#1) IS NOT NULL AND (#0 = #1)))) +
2193
2194 (1 row)
2195 */
2196
2197 if let MirScalarExpr::CallUnary {
2198 func: UnaryFunc::Not(scalar_func::Not),
2199 expr,
2200 } = predicate
2201 {
2202 if let MirScalarExpr::CallUnary {
2203 func: UnaryFunc::IsNull(scalar_func::IsNull),
2204 expr,
2205 } = &**expr
2206 {
2207 if let MirScalarExpr::Column(c, _name) = &**expr {
2208 nonnull_required_columns.insert(*c);
2209 }
2210 }
2211 }
2212 }
2213
2214 nonnull_required_columns
2215}
2216
2217impl CollectionPlan for MirRelationExpr {
2218 /// Collects the global collections that this MIR expression directly depends on, i.e., that it
2219 /// has a `Get` for. (It does _not_ traverse view definitions transitively.)
2220 ///
2221 /// !!!WARNING!!!: this method has an HirRelationExpr counterpart. The two
2222 /// should be kept in sync w.r.t. HIR ⇒ MIR lowering!
2223 fn depends_on_into(&self, out: &mut BTreeSet<GlobalId>) {
2224 if let MirRelationExpr::Get {
2225 id: Id::Global(id), ..
2226 } = self
2227 {
2228 out.insert(*id);
2229 }
2230 self.visit_children(|expr| expr.depends_on_into(out))
2231 }
2232}
2233
2234impl MirRelationExpr {
2235 /// Iterates through references to child expressions.
2236 pub fn children(&self) -> impl DoubleEndedIterator<Item = &Self> {
2237 let mut first = None;
2238 let mut second = None;
2239 let mut rest = None;
2240 let mut last = None;
2241
2242 use MirRelationExpr::*;
2243 match self {
2244 Constant { .. } | Get { .. } => (),
2245 Let { value, body, .. } => {
2246 first = Some(&**value);
2247 second = Some(&**body);
2248 }
2249 LetRec { values, body, .. } => {
2250 rest = Some(values);
2251 last = Some(&**body);
2252 }
2253 Project { input, .. }
2254 | Map { input, .. }
2255 | FlatMap { input, .. }
2256 | Filter { input, .. }
2257 | Reduce { input, .. }
2258 | TopK { input, .. }
2259 | Negate { input }
2260 | Threshold { input }
2261 | ArrangeBy { input, .. } => {
2262 first = Some(&**input);
2263 }
2264 Join { inputs, .. } => {
2265 rest = Some(inputs);
2266 }
2267 Union { base, inputs } => {
2268 first = Some(&**base);
2269 rest = Some(inputs);
2270 }
2271 }
2272
2273 first
2274 .into_iter()
2275 .chain(second)
2276 .chain(rest.into_iter().flatten())
2277 .chain(last)
2278 }
2279
2280 /// Iterates through mutable references to child expressions.
2281 pub fn children_mut(&mut self) -> impl DoubleEndedIterator<Item = &mut Self> {
2282 let mut first = None;
2283 let mut second = None;
2284 let mut rest = None;
2285 let mut last = None;
2286
2287 use MirRelationExpr::*;
2288 match self {
2289 Constant { .. } | Get { .. } => (),
2290 Let { value, body, .. } => {
2291 first = Some(&mut **value);
2292 second = Some(&mut **body);
2293 }
2294 LetRec { values, body, .. } => {
2295 rest = Some(values);
2296 last = Some(&mut **body);
2297 }
2298 Project { input, .. }
2299 | Map { input, .. }
2300 | FlatMap { input, .. }
2301 | Filter { input, .. }
2302 | Reduce { input, .. }
2303 | TopK { input, .. }
2304 | Negate { input }
2305 | Threshold { input }
2306 | ArrangeBy { input, .. } => {
2307 first = Some(&mut **input);
2308 }
2309 Join { inputs, .. } => {
2310 rest = Some(inputs);
2311 }
2312 Union { base, inputs } => {
2313 first = Some(&mut **base);
2314 rest = Some(inputs);
2315 }
2316 }
2317
2318 first
2319 .into_iter()
2320 .chain(second)
2321 .chain(rest.into_iter().flatten())
2322 .chain(last)
2323 }
2324
2325 /// Iterative pre-order visitor.
2326 pub fn visit_pre<'a, F: FnMut(&'a Self)>(&'a self, mut f: F) {
2327 let mut worklist = vec![self];
2328 while let Some(expr) = worklist.pop() {
2329 f(expr);
2330 worklist.extend(expr.children().rev());
2331 }
2332 }
2333
2334 /// Iterative pre-order visitor.
2335 pub fn visit_pre_mut<F: FnMut(&mut Self)>(&mut self, mut f: F) {
2336 let mut worklist = vec![self];
2337 while let Some(expr) = worklist.pop() {
2338 f(expr);
2339 worklist.extend(expr.children_mut().rev());
2340 }
2341 }
2342
2343 /// Return a vector of references to the subtrees of this expression
2344 /// in post-visit order (the last element is `&self`).
2345 pub fn post_order_vec(&self) -> Vec<&Self> {
2346 let mut stack = vec![self];
2347 let mut result = vec![];
2348 while let Some(expr) = stack.pop() {
2349 result.push(expr);
2350 stack.extend(expr.children());
2351 }
2352 result.reverse();
2353 result
2354 }
2355}
2356
2357impl VisitChildren<Self> for MirRelationExpr {
2358 fn visit_children<F>(&self, mut f: F)
2359 where
2360 F: FnMut(&Self),
2361 {
2362 for child in self.children() {
2363 f(child)
2364 }
2365 }
2366
2367 fn visit_mut_children<F>(&mut self, mut f: F)
2368 where
2369 F: FnMut(&mut Self),
2370 {
2371 for child in self.children_mut() {
2372 f(child)
2373 }
2374 }
2375
2376 fn try_visit_children<F, E>(&self, mut f: F) -> Result<(), E>
2377 where
2378 F: FnMut(&Self) -> Result<(), E>,
2379 {
2380 for child in self.children() {
2381 f(child)?
2382 }
2383 Ok(())
2384 }
2385
2386 fn try_visit_mut_children<F, E>(&mut self, mut f: F) -> Result<(), E>
2387 where
2388 F: FnMut(&mut Self) -> Result<(), E>,
2389 {
2390 for child in self.children_mut() {
2391 f(child)?
2392 }
2393 Ok(())
2394 }
2395
2396 fn children<'a>(&'a self) -> impl DoubleEndedIterator<Item = &'a MirRelationExpr>
2397 where
2398 Self: 'a,
2399 {
2400 self.children()
2401 }
2402
2403 fn children_mut<'a>(&'a mut self) -> impl DoubleEndedIterator<Item = &'a mut MirRelationExpr>
2404 where
2405 Self: 'a,
2406 {
2407 self.children_mut()
2408 }
2409}
2410
2411/// Specification for an ordering by a column.
2412#[derive(
2413 Debug,
2414 Clone,
2415 Copy,
2416 Eq,
2417 PartialEq,
2418 Ord,
2419 PartialOrd,
2420 Serialize,
2421 Deserialize,
2422 Hash
2423)]
2424pub struct ColumnOrder {
2425 /// The column index.
2426 pub column: usize,
2427 /// Whether to sort in descending order.
2428 #[serde(default)]
2429 pub desc: bool,
2430 /// Whether to sort nulls last.
2431 #[serde(default)]
2432 pub nulls_last: bool,
2433}
2434
2435impl Columnation for ColumnOrder {
2436 type InnerRegion = CopyRegion<Self>;
2437}
2438
2439impl<'a, M> fmt::Display for HumanizedExpr<'a, ColumnOrder, M>
2440where
2441 M: HumanizerMode,
2442{
2443 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
2444 // If you modify this, then please also attend to Display for ColumnOrderWithExpr!
2445 write!(
2446 f,
2447 "{} {} {}",
2448 self.child(&self.expr.column),
2449 if self.expr.desc { "desc" } else { "asc" },
2450 if self.expr.nulls_last {
2451 "nulls_last"
2452 } else {
2453 "nulls_first"
2454 },
2455 )
2456 }
2457}
2458
2459/// Describes an aggregation expression.
2460#[derive(
2461 Clone,
2462 Debug,
2463 Eq,
2464 PartialEq,
2465 Ord,
2466 PartialOrd,
2467 Serialize,
2468 Deserialize,
2469 Hash
2470)]
2471pub struct AggregateExpr {
2472 /// Names the aggregation function.
2473 pub func: AggregateFunc,
2474 /// An expression which extracts from each row the input to `func`.
2475 pub expr: MirScalarExpr,
2476 /// Should the aggregation be applied only to distinct results in each group.
2477 #[serde(default)]
2478 pub distinct: bool,
2479}
2480
2481impl AggregateExpr {
2482 /// Computes the type of this `AggregateExpr`.
2483 pub fn sql_typ(&self, column_types: &[SqlColumnType]) -> SqlColumnType {
2484 self.func.output_sql_type(self.expr.sql_typ(column_types))
2485 }
2486
2487 /// Computes the type of this `AggregateExpr`.
2488 pub fn typ(&self, column_types: &[ReprColumnType]) -> ReprColumnType {
2489 self.func.output_type(self.expr.typ(column_types))
2490 }
2491
2492 /// Returns whether the expression has a constant result.
2493 pub fn is_constant(&self) -> bool {
2494 match self.func {
2495 AggregateFunc::MaxNumeric
2496 | AggregateFunc::MaxInt16
2497 | AggregateFunc::MaxInt32
2498 | AggregateFunc::MaxInt64
2499 | AggregateFunc::MaxUInt16
2500 | AggregateFunc::MaxUInt32
2501 | AggregateFunc::MaxUInt64
2502 | AggregateFunc::MaxMzTimestamp
2503 | AggregateFunc::MaxFloat32
2504 | AggregateFunc::MaxFloat64
2505 | AggregateFunc::MaxBool
2506 | AggregateFunc::MaxString
2507 | AggregateFunc::MaxDate
2508 | AggregateFunc::MaxTimestamp
2509 | AggregateFunc::MaxTimestampTz
2510 | AggregateFunc::MaxInterval
2511 | AggregateFunc::MaxTime
2512 | AggregateFunc::MinNumeric
2513 | AggregateFunc::MinInt16
2514 | AggregateFunc::MinInt32
2515 | AggregateFunc::MinInt64
2516 | AggregateFunc::MinUInt16
2517 | AggregateFunc::MinUInt32
2518 | AggregateFunc::MinUInt64
2519 | AggregateFunc::MinMzTimestamp
2520 | AggregateFunc::MinFloat32
2521 | AggregateFunc::MinFloat64
2522 | AggregateFunc::MinBool
2523 | AggregateFunc::MinString
2524 | AggregateFunc::MinDate
2525 | AggregateFunc::MinTimestamp
2526 | AggregateFunc::MinTimestampTz
2527 | AggregateFunc::MinInterval
2528 | AggregateFunc::MinTime
2529 | AggregateFunc::Any
2530 | AggregateFunc::All
2531 | AggregateFunc::Dummy => self.expr.is_literal(),
2532 AggregateFunc::Count => self.expr.is_literal_null(),
2533 AggregateFunc::SumInt16
2534 | AggregateFunc::SumInt32
2535 | AggregateFunc::SumInt64
2536 | AggregateFunc::SumUInt16
2537 | AggregateFunc::SumUInt32
2538 | AggregateFunc::SumUInt64
2539 | AggregateFunc::SumFloat32
2540 | AggregateFunc::SumFloat64
2541 | AggregateFunc::SumNumeric
2542 | AggregateFunc::SumInterval
2543 | AggregateFunc::JsonbAgg { .. }
2544 | AggregateFunc::JsonbObjectAgg { .. }
2545 | AggregateFunc::MapAgg { .. }
2546 | AggregateFunc::ArrayConcat { .. }
2547 | AggregateFunc::ListConcat { .. }
2548 | AggregateFunc::StringAgg { .. }
2549 | AggregateFunc::RowNumber { .. }
2550 | AggregateFunc::Rank { .. }
2551 | AggregateFunc::DenseRank { .. }
2552 | AggregateFunc::LagLead { .. }
2553 | AggregateFunc::FirstValue { .. }
2554 | AggregateFunc::LastValue { .. }
2555 | AggregateFunc::FusedValueWindowFunc { .. }
2556 | AggregateFunc::WindowAggregate { .. }
2557 | AggregateFunc::FusedWindowAggregate { .. } => self.expr.is_literal_err(),
2558 }
2559 }
2560
2561 /// Returns an expression that computes `self` on a group that has exactly one row.
2562 /// Instead of performing a `Reduce` with `self`, one can perform a `Map` with the expression
2563 /// returned by `on_unique`, which is cheaper. (See `ReduceElision`.)
2564 pub fn on_unique(&self, input_type: &[ReprColumnType]) -> MirScalarExpr {
2565 match &self.func {
2566 // Count is one if non-null, and zero if null.
2567 AggregateFunc::Count => self
2568 .expr
2569 .clone()
2570 .call_unary(UnaryFunc::IsNull(crate::func::IsNull))
2571 .if_then_else(
2572 MirScalarExpr::literal_ok(Datum::Int64(0), ReprScalarType::Int64),
2573 MirScalarExpr::literal_ok(Datum::Int64(1), ReprScalarType::Int64),
2574 ),
2575
2576 // SumInt16 takes Int16s as input, but outputs Int64s.
2577 AggregateFunc::SumInt16 => self
2578 .expr
2579 .clone()
2580 .call_unary(UnaryFunc::CastInt16ToInt64(scalar_func::CastInt16ToInt64)),
2581
2582 // SumInt32 takes Int32s as input, but outputs Int64s.
2583 AggregateFunc::SumInt32 => self
2584 .expr
2585 .clone()
2586 .call_unary(UnaryFunc::CastInt32ToInt64(scalar_func::CastInt32ToInt64)),
2587
2588 // SumInt64 takes Int64s as input, but outputs numerics.
2589 AggregateFunc::SumInt64 => self.expr.clone().call_unary(UnaryFunc::CastInt64ToNumeric(
2590 scalar_func::CastInt64ToNumeric(Some(NumericMaxScale::ZERO)),
2591 )),
2592
2593 // SumUInt16 takes UInt16s as input, but outputs UInt64s.
2594 AggregateFunc::SumUInt16 => self.expr.clone().call_unary(
2595 UnaryFunc::CastUint16ToUint64(scalar_func::CastUint16ToUint64),
2596 ),
2597
2598 // SumUInt32 takes UInt32s as input, but outputs UInt64s.
2599 AggregateFunc::SumUInt32 => self.expr.clone().call_unary(
2600 UnaryFunc::CastUint32ToUint64(scalar_func::CastUint32ToUint64),
2601 ),
2602
2603 // SumUInt64 takes UInt64s as input, but outputs numerics.
2604 AggregateFunc::SumUInt64 => {
2605 self.expr.clone().call_unary(UnaryFunc::CastUint64ToNumeric(
2606 scalar_func::CastUint64ToNumeric(Some(NumericMaxScale::ZERO)),
2607 ))
2608 }
2609
2610 // JsonbAgg takes _anything_ as input, but must output a Jsonb array.
2611 AggregateFunc::JsonbAgg { .. } => MirScalarExpr::call_variadic(
2612 JsonbBuildArray,
2613 vec![
2614 self.expr
2615 .clone()
2616 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0))),
2617 ],
2618 ),
2619
2620 // JsonbAgg takes _anything_ as input, but must output a Jsonb object.
2621 AggregateFunc::JsonbObjectAgg { .. } => {
2622 let record = self
2623 .expr
2624 .clone()
2625 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2626 MirScalarExpr::call_variadic(
2627 JsonbBuildObject,
2628 (0..2)
2629 .map(|i| {
2630 record
2631 .clone()
2632 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(i)))
2633 })
2634 .collect(),
2635 )
2636 }
2637
2638 AggregateFunc::MapAgg { value_type, .. } => {
2639 let record = self
2640 .expr
2641 .clone()
2642 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2643 MirScalarExpr::call_variadic(
2644 MapBuild {
2645 value_type: value_type.clone(),
2646 },
2647 (0..2)
2648 .map(|i| {
2649 record
2650 .clone()
2651 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(i)))
2652 })
2653 .collect(),
2654 )
2655 }
2656
2657 // StringAgg takes nested records of strings and outputs a string
2658 AggregateFunc::StringAgg { .. } => self
2659 .expr
2660 .clone()
2661 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)))
2662 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0))),
2663
2664 // ListConcat and ArrayConcat take a single level of records and output a list containing exactly 1 element
2665 AggregateFunc::ListConcat { .. } | AggregateFunc::ArrayConcat { .. } => self
2666 .expr
2667 .clone()
2668 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0))),
2669
2670 // RowNumber, Rank, DenseRank take a list of records and output a list containing exactly 1 element
2671 AggregateFunc::RowNumber { .. } => {
2672 self.on_unique_ranking_window_funcs(input_type, "?row_number?")
2673 }
2674 AggregateFunc::Rank { .. } => self.on_unique_ranking_window_funcs(input_type, "?rank?"),
2675 AggregateFunc::DenseRank { .. } => {
2676 self.on_unique_ranking_window_funcs(input_type, "?dense_rank?")
2677 }
2678
2679 // The input type for LagLead is ((OriginalRow, (InputValue, Offset, Default)), OrderByExprs...)
2680 AggregateFunc::LagLead { lag_lead, .. } => {
2681 let tuple = self
2682 .expr
2683 .clone()
2684 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2685
2686 // Get the overall return type
2687 let return_type_with_orig_row = self
2688 .typ(input_type)
2689 .scalar_type
2690 .unwrap_list_element_type()
2691 .clone();
2692 let lag_lead_return_type =
2693 return_type_with_orig_row.unwrap_record_element_type()[0].clone();
2694
2695 // Extract the original row
2696 let original_row = tuple
2697 .clone()
2698 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2699
2700 // Extract the encoded args
2701 let encoded_args =
2702 tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2703
2704 let (result_expr, column_name) =
2705 Self::on_unique_lag_lead(lag_lead, encoded_args, lag_lead_return_type.clone());
2706
2707 MirScalarExpr::call_variadic(
2708 ListCreate {
2709 elem_type: SqlScalarType::from_repr(&return_type_with_orig_row),
2710 },
2711 vec![MirScalarExpr::call_variadic(
2712 RecordCreate {
2713 field_names: vec![column_name, ColumnName::from("?record?")],
2714 },
2715 vec![result_expr, original_row],
2716 )],
2717 )
2718 }
2719
2720 // The input type for FirstValue is ((OriginalRow, InputValue), OrderByExprs...)
2721 AggregateFunc::FirstValue { window_frame, .. } => {
2722 let tuple = self
2723 .expr
2724 .clone()
2725 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2726
2727 // Get the overall return type
2728 let return_type_with_orig_row = self
2729 .typ(input_type)
2730 .scalar_type
2731 .unwrap_list_element_type()
2732 .clone();
2733 let first_value_return_type =
2734 return_type_with_orig_row.unwrap_record_element_type()[0].clone();
2735
2736 // Extract the original row
2737 let original_row = tuple
2738 .clone()
2739 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2740
2741 // Extract the input value
2742 let arg = tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2743
2744 let (result_expr, column_name) = Self::on_unique_first_value_last_value(
2745 window_frame,
2746 arg,
2747 first_value_return_type,
2748 );
2749
2750 MirScalarExpr::call_variadic(
2751 ListCreate {
2752 elem_type: SqlScalarType::from_repr(&return_type_with_orig_row),
2753 },
2754 vec![MirScalarExpr::call_variadic(
2755 RecordCreate {
2756 field_names: vec![column_name, ColumnName::from("?record?")],
2757 },
2758 vec![result_expr, original_row],
2759 )],
2760 )
2761 }
2762
2763 // The input type for LastValue is ((OriginalRow, InputValue), OrderByExprs...)
2764 AggregateFunc::LastValue { window_frame, .. } => {
2765 let tuple = self
2766 .expr
2767 .clone()
2768 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2769
2770 // Get the overall return type
2771 let return_type_with_orig_row = self
2772 .typ(input_type)
2773 .scalar_type
2774 .unwrap_list_element_type()
2775 .clone();
2776 let last_value_return_type =
2777 return_type_with_orig_row.unwrap_record_element_type()[0].clone();
2778
2779 // Extract the original row
2780 let original_row = tuple
2781 .clone()
2782 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2783
2784 // Extract the input value
2785 let arg = tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2786
2787 let (result_expr, column_name) = Self::on_unique_first_value_last_value(
2788 window_frame,
2789 arg,
2790 last_value_return_type,
2791 );
2792
2793 MirScalarExpr::call_variadic(
2794 ListCreate {
2795 elem_type: SqlScalarType::from_repr(&return_type_with_orig_row),
2796 },
2797 vec![MirScalarExpr::call_variadic(
2798 RecordCreate {
2799 field_names: vec![column_name, ColumnName::from("?record?")],
2800 },
2801 vec![result_expr, original_row],
2802 )],
2803 )
2804 }
2805
2806 // The input type for window aggs is ((OriginalRow, InputValue), OrderByExprs...)
2807 // See an example MIR in `window_func_applied_to`.
2808 AggregateFunc::WindowAggregate {
2809 wrapped_aggregate,
2810 window_frame,
2811 order_by: _,
2812 } => {
2813 // TODO: deduplicate code between the various window function cases.
2814
2815 let tuple = self
2816 .expr
2817 .clone()
2818 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2819
2820 // Get the overall return type
2821 let return_type = self
2822 .typ(input_type)
2823 .scalar_type
2824 .unwrap_list_element_type()
2825 .clone();
2826 let window_agg_return_type = return_type.unwrap_record_element_type()[0].clone();
2827
2828 // Extract the original row
2829 let original_row = tuple
2830 .clone()
2831 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2832
2833 // Extract the input value
2834 let arg_expr = tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2835
2836 let (result, column_name) = Self::on_unique_window_agg(
2837 window_frame,
2838 arg_expr,
2839 input_type,
2840 window_agg_return_type,
2841 wrapped_aggregate,
2842 );
2843
2844 MirScalarExpr::call_variadic(
2845 ListCreate {
2846 elem_type: SqlScalarType::from_repr(&return_type),
2847 },
2848 vec![MirScalarExpr::call_variadic(
2849 RecordCreate {
2850 field_names: vec![column_name, ColumnName::from("?record?")],
2851 },
2852 vec![result, original_row],
2853 )],
2854 )
2855 }
2856
2857 // The input type is ((OriginalRow, (Arg1, Arg2, ...)), OrderByExprs...)
2858 AggregateFunc::FusedWindowAggregate {
2859 wrapped_aggregates,
2860 order_by: _,
2861 window_frame,
2862 } => {
2863 // Throw away OrderByExprs
2864 let tuple = self
2865 .expr
2866 .clone()
2867 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2868
2869 // Extract the original row
2870 let original_row = tuple
2871 .clone()
2872 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2873
2874 // Extract the args of the fused call
2875 let all_args = tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2876
2877 let return_type_with_orig_row = self
2878 .typ(input_type)
2879 .scalar_type
2880 .unwrap_list_element_type()
2881 .clone();
2882
2883 let all_func_return_types =
2884 return_type_with_orig_row.unwrap_record_element_type()[0].clone();
2885 let mut func_result_exprs = Vec::new();
2886 let mut col_names = Vec::new();
2887 for (idx, wrapped_aggr) in wrapped_aggregates.iter().enumerate() {
2888 let arg = all_args
2889 .clone()
2890 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(idx)));
2891 let return_type =
2892 all_func_return_types.unwrap_record_element_type()[idx].clone();
2893 let (result, column_name) = Self::on_unique_window_agg(
2894 window_frame,
2895 arg,
2896 input_type,
2897 return_type,
2898 wrapped_aggr,
2899 );
2900 func_result_exprs.push(result);
2901 col_names.push(column_name);
2902 }
2903
2904 MirScalarExpr::call_variadic(
2905 ListCreate {
2906 elem_type: SqlScalarType::from_repr(&return_type_with_orig_row),
2907 },
2908 vec![MirScalarExpr::call_variadic(
2909 RecordCreate {
2910 field_names: vec![
2911 ColumnName::from("?fused_window_aggr?"),
2912 ColumnName::from("?record?"),
2913 ],
2914 },
2915 vec![
2916 MirScalarExpr::call_variadic(
2917 RecordCreate {
2918 field_names: col_names,
2919 },
2920 func_result_exprs,
2921 ),
2922 original_row,
2923 ],
2924 )],
2925 )
2926 }
2927
2928 // The input type is ((OriginalRow, (Args1, Args2, ...)), OrderByExprs...)
2929 AggregateFunc::FusedValueWindowFunc {
2930 funcs,
2931 order_by: outer_order_by,
2932 } => {
2933 // Throw away OrderByExprs
2934 let tuple = self
2935 .expr
2936 .clone()
2937 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2938
2939 // Extract the original row
2940 let original_row = tuple
2941 .clone()
2942 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
2943
2944 // Extract the encoded args of the fused call
2945 let all_encoded_args =
2946 tuple.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
2947
2948 let return_type_with_orig_row = self
2949 .typ(input_type)
2950 .scalar_type
2951 .unwrap_list_element_type()
2952 .clone();
2953
2954 let all_func_return_types =
2955 return_type_with_orig_row.unwrap_record_element_type()[0].clone();
2956 let mut func_result_exprs = Vec::new();
2957 let mut col_names = Vec::new();
2958 for (idx, func) in funcs.iter().enumerate() {
2959 let args_for_func = all_encoded_args
2960 .clone()
2961 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(idx)));
2962 let return_type_for_func =
2963 all_func_return_types.unwrap_record_element_type()[idx].clone();
2964 let (result, column_name) = match func {
2965 AggregateFunc::LagLead {
2966 lag_lead,
2967 order_by,
2968 ignore_nulls: _,
2969 } => {
2970 assert_eq!(order_by, outer_order_by);
2971 Self::on_unique_lag_lead(lag_lead, args_for_func, return_type_for_func)
2972 }
2973 AggregateFunc::FirstValue {
2974 window_frame,
2975 order_by,
2976 } => {
2977 assert_eq!(order_by, outer_order_by);
2978 Self::on_unique_first_value_last_value(
2979 window_frame,
2980 args_for_func,
2981 return_type_for_func,
2982 )
2983 }
2984 AggregateFunc::LastValue {
2985 window_frame,
2986 order_by,
2987 } => {
2988 assert_eq!(order_by, outer_order_by);
2989 Self::on_unique_first_value_last_value(
2990 window_frame,
2991 args_for_func,
2992 return_type_for_func,
2993 )
2994 }
2995 _ => panic!("unknown function in FusedValueWindowFunc"),
2996 };
2997 func_result_exprs.push(result);
2998 col_names.push(column_name);
2999 }
3000
3001 MirScalarExpr::call_variadic(
3002 ListCreate {
3003 elem_type: SqlScalarType::from_repr(&return_type_with_orig_row),
3004 },
3005 vec![MirScalarExpr::call_variadic(
3006 RecordCreate {
3007 field_names: vec![
3008 ColumnName::from("?fused_value_window_func?"),
3009 ColumnName::from("?record?"),
3010 ],
3011 },
3012 vec![
3013 MirScalarExpr::call_variadic(
3014 RecordCreate {
3015 field_names: col_names,
3016 },
3017 func_result_exprs,
3018 ),
3019 original_row,
3020 ],
3021 )],
3022 )
3023 }
3024
3025 // All other variants should return the argument to the aggregation.
3026 AggregateFunc::MaxNumeric
3027 | AggregateFunc::MaxInt16
3028 | AggregateFunc::MaxInt32
3029 | AggregateFunc::MaxInt64
3030 | AggregateFunc::MaxUInt16
3031 | AggregateFunc::MaxUInt32
3032 | AggregateFunc::MaxUInt64
3033 | AggregateFunc::MaxMzTimestamp
3034 | AggregateFunc::MaxFloat32
3035 | AggregateFunc::MaxFloat64
3036 | AggregateFunc::MaxBool
3037 | AggregateFunc::MaxString
3038 | AggregateFunc::MaxDate
3039 | AggregateFunc::MaxTimestamp
3040 | AggregateFunc::MaxTimestampTz
3041 | AggregateFunc::MaxInterval
3042 | AggregateFunc::MaxTime
3043 | AggregateFunc::MinNumeric
3044 | AggregateFunc::MinInt16
3045 | AggregateFunc::MinInt32
3046 | AggregateFunc::MinInt64
3047 | AggregateFunc::MinUInt16
3048 | AggregateFunc::MinUInt32
3049 | AggregateFunc::MinUInt64
3050 | AggregateFunc::MinMzTimestamp
3051 | AggregateFunc::MinFloat32
3052 | AggregateFunc::MinFloat64
3053 | AggregateFunc::MinBool
3054 | AggregateFunc::MinString
3055 | AggregateFunc::MinDate
3056 | AggregateFunc::MinTimestamp
3057 | AggregateFunc::MinTimestampTz
3058 | AggregateFunc::MinInterval
3059 | AggregateFunc::MinTime
3060 | AggregateFunc::SumFloat32
3061 | AggregateFunc::SumFloat64
3062 | AggregateFunc::SumNumeric
3063 | AggregateFunc::SumInterval
3064 | AggregateFunc::Any
3065 | AggregateFunc::All
3066 | AggregateFunc::Dummy => self.expr.clone(),
3067 }
3068 }
3069
3070 /// `on_unique` for ROW_NUMBER, RANK, DENSE_RANK
3071 fn on_unique_ranking_window_funcs(
3072 &self,
3073 input_type: &[ReprColumnType],
3074 col_name: &str,
3075 ) -> MirScalarExpr {
3076 let sql_input_type: Vec<SqlColumnType> =
3077 input_type.iter().map(SqlColumnType::from_repr).collect();
3078 let list = self
3079 .expr
3080 .clone()
3081 // extract the list within the record
3082 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
3083
3084 // extract the expression within the list
3085 let record = MirScalarExpr::call_variadic(
3086 ListIndex,
3087 vec![
3088 list,
3089 MirScalarExpr::literal_ok(Datum::Int64(1), ReprScalarType::Int64),
3090 ],
3091 );
3092
3093 MirScalarExpr::call_variadic(
3094 ListCreate {
3095 elem_type: self
3096 .sql_typ(&sql_input_type)
3097 .scalar_type
3098 .unwrap_list_element_type()
3099 .clone(),
3100 },
3101 vec![MirScalarExpr::call_variadic(
3102 RecordCreate {
3103 field_names: vec![ColumnName::from(col_name), ColumnName::from("?record?")],
3104 },
3105 vec![
3106 MirScalarExpr::literal_ok(Datum::Int64(1), ReprScalarType::Int64),
3107 record,
3108 ],
3109 )],
3110 )
3111 }
3112
3113 /// `on_unique` for `lag` and `lead`
3114 fn on_unique_lag_lead(
3115 lag_lead: &LagLeadType,
3116 encoded_args: MirScalarExpr,
3117 return_type: ReprScalarType,
3118 ) -> (MirScalarExpr, ColumnName) {
3119 let expr = encoded_args
3120 .clone()
3121 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(0)));
3122 let offset = encoded_args
3123 .clone()
3124 .call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(1)));
3125 let default_value =
3126 encoded_args.call_unary(UnaryFunc::RecordGet(scalar_func::RecordGet(2)));
3127
3128 // In this case, the window always has only one element, so if the offset is not null and
3129 // not zero, the default value should be returned instead.
3130 let value = offset
3131 .clone()
3132 .call_binary(
3133 MirScalarExpr::literal_ok(Datum::Int32(0), ReprScalarType::Int32),
3134 crate::func::Eq,
3135 )
3136 .if_then_else(expr, default_value);
3137 let result_expr = offset
3138 .call_unary(UnaryFunc::IsNull(crate::func::IsNull))
3139 .if_then_else(MirScalarExpr::literal_null(return_type), value);
3140
3141 let column_name = ColumnName::from(match lag_lead {
3142 LagLeadType::Lag => "?lag?",
3143 LagLeadType::Lead => "?lead?",
3144 });
3145
3146 (result_expr, column_name)
3147 }
3148
3149 /// `on_unique` for `first_value` and `last_value`
3150 fn on_unique_first_value_last_value(
3151 window_frame: &WindowFrame,
3152 arg: MirScalarExpr,
3153 return_type: ReprScalarType,
3154 ) -> (MirScalarExpr, ColumnName) {
3155 // If the window frame includes the current (single) row, return its value, null otherwise
3156 let result_expr = if window_frame.includes_current_row() {
3157 arg
3158 } else {
3159 MirScalarExpr::literal_null(return_type)
3160 };
3161 (result_expr, ColumnName::from("?first_value?"))
3162 }
3163
3164 /// `on_unique` for window aggregations
3165 fn on_unique_window_agg(
3166 window_frame: &WindowFrame,
3167 arg_expr: MirScalarExpr,
3168 input_type: &[ReprColumnType],
3169 return_type: ReprScalarType,
3170 wrapped_aggr: &AggregateFunc,
3171 ) -> (MirScalarExpr, ColumnName) {
3172 // If the window frame includes the current (single) row, evaluate the wrapped aggregate on
3173 // that row. Otherwise, return the default value for the aggregate.
3174 let result_expr = if window_frame.includes_current_row() {
3175 AggregateExpr {
3176 func: wrapped_aggr.clone(),
3177 expr: arg_expr,
3178 distinct: false, // We have just one input element; DISTINCT doesn't matter.
3179 }
3180 .on_unique(input_type)
3181 } else {
3182 MirScalarExpr::literal_ok(wrapped_aggr.default(), return_type)
3183 };
3184 (result_expr, ColumnName::from("?window_agg?"))
3185 }
3186
3187 /// Returns whether the expression is COUNT(*) or not. Note that
3188 /// when we define the count builtin in sql::func, we convert
3189 /// COUNT(*) to COUNT(true), making it indistinguishable from
3190 /// literal COUNT(true), but we prefer to consider this as the
3191 /// former.
3192 ///
3193 /// (HIR has the same `is_count_asterisk`.)
3194 pub fn is_count_asterisk(&self) -> bool {
3195 self.func == AggregateFunc::Count && self.expr.is_literal_true() && !self.distinct
3196 }
3197}
3198
3199/// Describe a join implementation in dataflow.
3200#[derive(
3201 Clone,
3202 Debug,
3203 Eq,
3204 PartialEq,
3205 Ord,
3206 PartialOrd,
3207 Serialize,
3208 Deserialize,
3209 Hash
3210)]
3211pub enum JoinImplementation {
3212 /// Perform a sequence of binary differential dataflow joins.
3213 ///
3214 /// The first argument indicates
3215 /// 1) the index of the starting collection,
3216 /// 2) if it should be arranged, the keys to arrange it by, and
3217 /// 3) the characteristics of the starting collection (for EXPLAINing).
3218 /// The sequence that follows lists other relation indexes, and the key for
3219 /// the arrangement we should use when joining it in.
3220 /// The JoinInputCharacteristics are for EXPLAINing the characteristics that
3221 /// were used for join ordering.
3222 ///
3223 /// Each collection index should occur exactly once, either as the starting collection
3224 /// or somewhere in the list.
3225 Differential(
3226 (
3227 usize,
3228 Option<Vec<MirScalarExpr>>,
3229 Option<JoinInputCharacteristics>,
3230 ),
3231 Vec<(usize, Vec<MirScalarExpr>, Option<JoinInputCharacteristics>)>,
3232 ),
3233 /// Perform independent delta query dataflows for each input.
3234 ///
3235 /// The argument is a sequence of plans, for the input collections in order.
3236 /// Each plan starts from the corresponding index, and then in sequence joins
3237 /// against collections identified by index and with the specified arrangement key.
3238 /// The JoinInputCharacteristics are for EXPLAINing the characteristics that were
3239 /// used for join ordering.
3240 DeltaQuery(Vec<Vec<(usize, Vec<MirScalarExpr>, Option<JoinInputCharacteristics>)>>),
3241 /// Join a user-created index with a constant collection to speed up the evaluation of a
3242 /// predicate such as `(f1 = 3 AND f2 = 5) OR (f1 = 7 AND f2 = 9)`.
3243 /// This gets translated to a Differential join during MIR -> LIR lowering, but we still want
3244 /// to represent it in MIR, because the fast path detection wants to match on this.
3245 ///
3246 /// Consists of (`<coll_id>`, `<index_id>`, `<index_key>`, `<constants>`)
3247 IndexedFilter(GlobalId, GlobalId, Vec<MirScalarExpr>, Vec<Row>),
3248 /// No implementation yet selected.
3249 Unimplemented,
3250}
3251
3252impl Default for JoinImplementation {
3253 fn default() -> Self {
3254 JoinImplementation::Unimplemented
3255 }
3256}
3257
3258impl JoinImplementation {
3259 /// Returns `true` iff the value is not [`JoinImplementation::Unimplemented`].
3260 pub fn is_implemented(&self) -> bool {
3261 match self {
3262 Self::Unimplemented => false,
3263 _ => true,
3264 }
3265 }
3266
3267 /// Returns an optional implementation name if the value is not [`JoinImplementation::Unimplemented`].
3268 pub fn name(&self) -> Option<&'static str> {
3269 match self {
3270 Self::Differential(..) => Some("differential"),
3271 Self::DeltaQuery(..) => Some("delta"),
3272 Self::IndexedFilter(..) => Some("indexed_filter"),
3273 Self::Unimplemented => None,
3274 }
3275 }
3276}
3277
3278/// Characteristics of a join order candidate collection.
3279///
3280/// A candidate is described by a collection and a key, and may have various liabilities.
3281/// Primarily, the candidate may risk substantial inflation of records, which is something
3282/// that concerns us greatly. Additionally, the candidate may be unarranged, and we would
3283/// prefer candidates that do not require additional memory. Finally, we prefer lower id
3284/// collections in the interest of consistent tie-breaking. For more characteristics, see
3285/// comments on individual fields.
3286///
3287/// This has more than one version. `new` instantiates the appropriate version based on a
3288/// feature flag.
3289#[derive(
3290 Eq,
3291 PartialEq,
3292 Ord,
3293 PartialOrd,
3294 Debug,
3295 Clone,
3296 Serialize,
3297 Deserialize,
3298 Hash
3299)]
3300pub enum JoinInputCharacteristics {
3301 /// Old version, with `enable_join_prioritize_arranged` turned off.
3302 V1(JoinInputCharacteristicsV1),
3303 /// Newer version, with `enable_join_prioritize_arranged` turned on.
3304 V2(JoinInputCharacteristicsV2),
3305}
3306
3307impl JoinInputCharacteristics {
3308 /// Creates a new instance with the given characteristics.
3309 pub fn new(
3310 unique_key: bool,
3311 key_length: usize,
3312 arranged: bool,
3313 cardinality: Option<usize>,
3314 filters: FilterCharacteristics,
3315 input: usize,
3316 enable_join_prioritize_arranged: bool,
3317 ) -> Self {
3318 if enable_join_prioritize_arranged {
3319 Self::V2(JoinInputCharacteristicsV2::new(
3320 unique_key,
3321 key_length,
3322 arranged,
3323 cardinality,
3324 filters,
3325 input,
3326 ))
3327 } else {
3328 Self::V1(JoinInputCharacteristicsV1::new(
3329 unique_key,
3330 key_length,
3331 arranged,
3332 cardinality,
3333 filters,
3334 input,
3335 ))
3336 }
3337 }
3338
3339 /// Turns the instance into a String to be printed in EXPLAIN.
3340 pub fn explain(&self) -> String {
3341 match self {
3342 Self::V1(jic) => jic.explain(),
3343 Self::V2(jic) => jic.explain(),
3344 }
3345 }
3346
3347 /// Whether the join input described by `self` is arranged.
3348 pub fn arranged(&self) -> bool {
3349 match self {
3350 Self::V1(jic) => jic.arranged,
3351 Self::V2(jic) => jic.arranged,
3352 }
3353 }
3354
3355 /// Returns the `FilterCharacteristics` for the join input described by `self`.
3356 pub fn filters(&mut self) -> &mut FilterCharacteristics {
3357 match self {
3358 Self::V1(jic) => &mut jic.filters,
3359 Self::V2(jic) => &mut jic.filters,
3360 }
3361 }
3362}
3363
3364/// Newer version of `JoinInputCharacteristics`, with `enable_join_prioritize_arranged` turned on.
3365#[derive(
3366 Eq,
3367 PartialEq,
3368 Ord,
3369 PartialOrd,
3370 Debug,
3371 Clone,
3372 Serialize,
3373 Deserialize,
3374 Hash
3375)]
3376pub struct JoinInputCharacteristicsV2 {
3377 /// An excellent indication that record count will not increase.
3378 pub unique_key: bool,
3379 /// Cross joins are bad.
3380 /// (`key_length > 0` also implies that it is not a cross join. However, we need to note cross
3381 /// joins in a separate field, because not being a cross join is more important than `arranged`,
3382 /// but otherwise `key_length` is less important than `arranged`.)
3383 pub not_cross: bool,
3384 /// Indicates that there will be no additional in-memory footprint.
3385 pub arranged: bool,
3386 /// A weaker signal that record count will not increase.
3387 pub key_length: usize,
3388 /// Estimated cardinality (lower is better)
3389 pub cardinality: Option<std::cmp::Reverse<usize>>,
3390 /// Characteristics of the filter that is applied at this input.
3391 pub filters: FilterCharacteristics,
3392 /// We want to prefer input earlier in the input list, for stability of ordering.
3393 pub input: std::cmp::Reverse<usize>,
3394}
3395
3396impl JoinInputCharacteristicsV2 {
3397 /// Creates a new instance with the given characteristics.
3398 pub fn new(
3399 unique_key: bool,
3400 key_length: usize,
3401 arranged: bool,
3402 cardinality: Option<usize>,
3403 filters: FilterCharacteristics,
3404 input: usize,
3405 ) -> Self {
3406 Self {
3407 unique_key,
3408 not_cross: key_length > 0,
3409 arranged,
3410 key_length,
3411 cardinality: cardinality.map(std::cmp::Reverse),
3412 filters,
3413 input: std::cmp::Reverse(input),
3414 }
3415 }
3416
3417 /// Turns the instance into a String to be printed in EXPLAIN.
3418 pub fn explain(&self) -> String {
3419 let mut e = "".to_owned();
3420 if self.unique_key {
3421 e.push_str("U");
3422 }
3423 // Don't need to print `not_cross`, because that is visible in the printed key.
3424 // if !self.not_cross {
3425 // e.push_str("C");
3426 // }
3427 for _ in 0..self.key_length {
3428 e.push_str("K");
3429 }
3430 if self.arranged {
3431 e.push_str("A");
3432 }
3433 if let Some(std::cmp::Reverse(cardinality)) = self.cardinality {
3434 e.push_str(&format!("|{cardinality}|"));
3435 }
3436 e.push_str(&self.filters.explain());
3437 e
3438 }
3439}
3440
3441/// Old version of `JoinInputCharacteristics`, with `enable_join_prioritize_arranged` turned off.
3442#[derive(
3443 Eq,
3444 PartialEq,
3445 Ord,
3446 PartialOrd,
3447 Debug,
3448 Clone,
3449 Serialize,
3450 Deserialize,
3451 Hash
3452)]
3453pub struct JoinInputCharacteristicsV1 {
3454 /// An excellent indication that record count will not increase.
3455 pub unique_key: bool,
3456 /// A weaker signal that record count will not increase.
3457 pub key_length: usize,
3458 /// Indicates that there will be no additional in-memory footprint.
3459 pub arranged: bool,
3460 /// Estimated cardinality (lower is better)
3461 pub cardinality: Option<std::cmp::Reverse<usize>>,
3462 /// Characteristics of the filter that is applied at this input.
3463 pub filters: FilterCharacteristics,
3464 /// We want to prefer input earlier in the input list, for stability of ordering.
3465 pub input: std::cmp::Reverse<usize>,
3466}
3467
3468impl JoinInputCharacteristicsV1 {
3469 /// Creates a new instance with the given characteristics.
3470 pub fn new(
3471 unique_key: bool,
3472 key_length: usize,
3473 arranged: bool,
3474 cardinality: Option<usize>,
3475 filters: FilterCharacteristics,
3476 input: usize,
3477 ) -> Self {
3478 Self {
3479 unique_key,
3480 key_length,
3481 arranged,
3482 cardinality: cardinality.map(std::cmp::Reverse),
3483 filters,
3484 input: std::cmp::Reverse(input),
3485 }
3486 }
3487
3488 /// Turns the instance into a String to be printed in EXPLAIN.
3489 pub fn explain(&self) -> String {
3490 let mut e = "".to_owned();
3491 if self.unique_key {
3492 e.push_str("U");
3493 }
3494 for _ in 0..self.key_length {
3495 e.push_str("K");
3496 }
3497 if self.arranged {
3498 e.push_str("A");
3499 }
3500 if let Some(std::cmp::Reverse(cardinality)) = self.cardinality {
3501 e.push_str(&format!("|{cardinality}|"));
3502 }
3503 e.push_str(&self.filters.explain());
3504 e
3505 }
3506}
3507
3508/// Instructions for finishing the result of a query.
3509///
3510/// The primary reason for the existence of this structure and attendant code
3511/// is that SQL's ORDER BY requires sorting rows (as already implied by the
3512/// keywords), whereas much of the rest of SQL is defined in terms of unordered
3513/// multisets. But as it turns out, the same idea can be used to optimize
3514/// trivial peeks.
3515///
3516/// The generic parameters are for accommodating prepared statement parameters in
3517/// `limit` and `offset`: the planner can hold these fields as HirScalarExpr long enough to call
3518/// `bind_parameters` on them.
3519#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
3520pub struct RowSetFinishing<L = NonNeg<i64>, O = usize> {
3521 /// Order rows by the given columns.
3522 pub order_by: Vec<ColumnOrder>,
3523 /// Include only as many rows (after offset).
3524 pub limit: Option<L>,
3525 /// Omit as many rows.
3526 pub offset: O,
3527 /// Include only given columns.
3528 pub project: Vec<usize>,
3529}
3530
3531impl<L> RowSetFinishing<L> {
3532 /// Returns a trivial finishing, i.e., that does nothing to the result set.
3533 pub fn trivial(arity: usize) -> RowSetFinishing<L> {
3534 RowSetFinishing {
3535 order_by: Vec::new(),
3536 limit: None,
3537 offset: 0,
3538 project: (0..arity).collect(),
3539 }
3540 }
3541 /// True if the finishing does nothing to any result set.
3542 pub fn is_trivial(&self, arity: usize) -> bool {
3543 self.limit.is_none()
3544 && self.order_by.is_empty()
3545 && self.offset == 0
3546 && self.project.iter().copied().eq(0..arity)
3547 }
3548 /// True if the finishing does not require an ORDER BY.
3549 ///
3550 /// LIMIT and OFFSET without an ORDER BY _are_ streamable: without an
3551 /// explicit ordering we will skip an arbitrary bag of elements and return
3552 /// the first arbitrary elements in the remaining bag. The result semantics
3553 /// are still correct but maybe surprising for some users.
3554 pub fn is_streamable(&self, arity: usize) -> bool {
3555 self.order_by.is_empty() && self.project.iter().copied().eq(0..arity)
3556 }
3557}
3558
3559impl RowSetFinishing<NonNeg<i64>, usize> {
3560 /// The number of rows needed from before the finishing to evaluate the finishing:
3561 /// offset + limit.
3562 ///
3563 /// If it returns None, then we need all the rows.
3564 pub fn num_rows_needed(&self) -> Option<usize> {
3565 self.limit
3566 .as_ref()
3567 .map(|l| usize::cast_from(u64::from(l.clone())) + self.offset)
3568 }
3569}
3570
3571impl RowSetFinishing {
3572 /// Applies finishing actions to a [`RowCollection`], and reports the total
3573 /// time it took to run.
3574 ///
3575 /// Returns a [`RowCollectionIter`] that contains all of the response data, as
3576 /// well as the size of the response in bytes.
3577 pub fn finish(
3578 &self,
3579 rows: RowCollection,
3580 max_result_size: u64,
3581 max_returned_query_size: Option<u64>,
3582 duration_histogram: &Histogram,
3583 ) -> Result<(RowCollectionIter, usize), String> {
3584 let now = Instant::now();
3585 let result = self.finish_inner(rows, max_result_size, max_returned_query_size);
3586 let duration = now.elapsed();
3587 duration_histogram.observe(duration.as_secs_f64());
3588
3589 result
3590 }
3591
3592 /// Implementation for [`RowSetFinishing::finish`].
3593 fn finish_inner(
3594 &self,
3595 rows: RowCollection,
3596 max_result_size: u64,
3597 max_returned_query_size: Option<u64>,
3598 ) -> Result<(RowCollectionIter, usize), String> {
3599 // Bail if the already-materialized collection is larger than the cap.
3600 // `byte_len` includes the per-row offset metadata, and finishing does
3601 // not allocate any additional per-entry structure.
3602 if rows.byte_len() > usize::cast_from(max_result_size) {
3603 let max_bytes = ByteSize::b(max_result_size);
3604 return Err(format!("result exceeds max size of {max_bytes}",));
3605 }
3606
3607 let sorted_view = rows;
3608 let mut iter = sorted_view
3609 .into_row_iter()
3610 .apply_offset(self.offset)
3611 .with_projection(self.project.clone());
3612
3613 if let Some(limit) = self.limit {
3614 let limit = u64::from(limit);
3615 let limit = usize::cast_from(limit);
3616 iter = iter.with_limit(limit);
3617 };
3618
3619 // TODO(parkmycar): Re-think how we can calculate the total response size without
3620 // having to iterate through the entire collection of Rows, while still
3621 // respecting the LIMIT, OFFSET, and projections.
3622 //
3623 // Note: It feels a bit bad always calculating the response size, but we almost
3624 // always need it to either check the `max_returned_query_size`, or for reporting
3625 // in the query history.
3626 let response_size: usize = iter.clone().map(|row| row.data().len()).sum();
3627
3628 // Bail if we would end up returning more data to the client than they can support.
3629 if let Some(max) = max_returned_query_size {
3630 if response_size > usize::cast_from(max) {
3631 let max_bytes = ByteSize::b(max);
3632 return Err(format!("result exceeds max size of {max_bytes}"));
3633 }
3634 }
3635
3636 Ok((iter, response_size))
3637 }
3638}
3639
3640/// A [RowSetFinishing] that can be repeatedly applied to batches of updates (in
3641/// a [RowCollection]) and keeps track of the remaining limit, offset, and cap
3642/// on query result size.
3643#[derive(Debug)]
3644pub struct RowSetFinishingIncremental {
3645 /// Include only as many rows (after offset).
3646 pub remaining_limit: Option<usize>,
3647 /// Omit as many rows.
3648 pub remaining_offset: usize,
3649 /// The maximum allowed result size, as requested by the client.
3650 pub max_returned_query_size: Option<u64>,
3651 /// Tracks our remaining allowed budget for result size.
3652 pub remaining_max_returned_query_size: Option<u64>,
3653 /// Include only given columns.
3654 pub project: Vec<usize>,
3655}
3656
3657impl RowSetFinishingIncremental {
3658 /// Turns the given [RowSetFinishing] into a [RowSetFinishingIncremental].
3659 /// Can only be used when [is_streamable](RowSetFinishing::is_streamable) is
3660 /// `true`.
3661 ///
3662 /// # Panics
3663 ///
3664 /// Panics if the result is not streamable, that is it has an ORDER BY.
3665 pub fn new(
3666 offset: usize,
3667 limit: Option<NonNeg<i64>>,
3668 project: Vec<usize>,
3669 max_returned_query_size: Option<u64>,
3670 ) -> Self {
3671 let limit = limit.map(|l| {
3672 let l = u64::from(l);
3673 let l = usize::cast_from(l);
3674 l
3675 });
3676
3677 RowSetFinishingIncremental {
3678 remaining_limit: limit,
3679 remaining_offset: offset,
3680 max_returned_query_size,
3681 remaining_max_returned_query_size: max_returned_query_size,
3682 project,
3683 }
3684 }
3685
3686 /// Applies finishing actions to the given [`RowCollection`], and reports
3687 /// the total time it took to run.
3688 ///
3689 /// Returns a [`RowCollectionIter`] that contains all of the response
3690 /// data.
3691 pub fn finish_incremental(
3692 &mut self,
3693 rows: RowCollection,
3694 max_result_size: u64,
3695 duration_histogram: &Histogram,
3696 ) -> Result<RowCollectionIter, String> {
3697 let now = Instant::now();
3698 let result = self.finish_incremental_inner(rows, max_result_size);
3699 let duration = now.elapsed();
3700 duration_histogram.observe(duration.as_secs_f64());
3701
3702 result
3703 }
3704
3705 fn finish_incremental_inner(
3706 &mut self,
3707 rows: RowCollection,
3708 max_result_size: u64,
3709 ) -> Result<RowCollectionIter, String> {
3710 // Bail if the already-materialized collection is larger than the cap.
3711 // `byte_len` includes the per-row offset metadata, and finishing does
3712 // not allocate any additional per-entry structure.
3713 if rows.byte_len() > usize::cast_from(max_result_size) {
3714 let max_bytes = ByteSize::b(max_result_size);
3715 return Err(format!("total result exceeds max size of {max_bytes}",));
3716 }
3717
3718 let batch_num_rows = rows.count();
3719
3720 let sorted_view = rows;
3721 let mut iter = sorted_view
3722 .into_row_iter()
3723 .apply_offset(self.remaining_offset)
3724 .with_projection(self.project.clone());
3725
3726 if let Some(limit) = self.remaining_limit {
3727 iter = iter.with_limit(limit);
3728 };
3729
3730 self.remaining_offset = self.remaining_offset.saturating_sub(batch_num_rows);
3731 if let Some(remaining_limit) = self.remaining_limit.as_mut() {
3732 *remaining_limit -= iter.count();
3733 }
3734
3735 // TODO(parkmycar): Re-think how we can calculate the total response size without
3736 // having to iterate through the entire collection of Rows, while still
3737 // respecting the LIMIT, OFFSET, and projections.
3738 //
3739 // Note: It feels a bit bad always calculating the response size, but we almost
3740 // always need it to either check the `max_returned_query_size`, or for reporting
3741 // in the query history.
3742 let response_size: usize = iter.clone().map(|row| row.data().len()).sum();
3743
3744 // Bail if we would end up returning more data to the client than they can support.
3745 if let Some(max) = &mut self.remaining_max_returned_query_size {
3746 if let Some(remaining) = max.checked_sub(response_size.cast_into()) {
3747 *max = remaining;
3748 } else {
3749 let max_bytes = ByteSize::b(self.max_returned_query_size.expect("known to exist"));
3750 return Err(format!("total result exceeds max size of {max_bytes}"));
3751 }
3752 }
3753
3754 Ok(iter)
3755 }
3756}
3757
3758/// Compares two rows columnwise, using [compare_columns].
3759///
3760/// Compared to the naive implementation, this allows sharing some memory and implements some
3761/// optimizations that avoid unnecessary row unpacking.
3762#[derive(Debug, Clone)]
3763pub struct RowComparator<O: AsRef<[ColumnOrder]> = Vec<ColumnOrder>> {
3764 order: O,
3765 /// Invariant: all column references in the order are less than this limit.
3766 /// This allows for partial unpacking of rows.
3767 limit: usize,
3768 left_vec: RefCell<DatumVec>,
3769 right_vec: RefCell<DatumVec>,
3770}
3771
3772impl<O: AsRef<[ColumnOrder]>> RowComparator<O> {
3773 /// Create a new row comparator from the given column ordering.
3774 pub fn new(order: O) -> Self {
3775 let limit = order
3776 .as_ref()
3777 .iter()
3778 .map(|o| o.column + 1)
3779 .max()
3780 .unwrap_or(0);
3781 Self {
3782 order,
3783 limit,
3784 left_vec: Default::default(),
3785 right_vec: Default::default(),
3786 }
3787 }
3788
3789 /// Compare two (references to) rows.
3790 pub fn compare_rows(
3791 &self,
3792 left_row: &RowRef,
3793 right_row: &RowRef,
3794 tiebreaker: impl Fn() -> Ordering,
3795 ) -> Ordering {
3796 let order = if self.limit == 0 {
3797 Ordering::Equal
3798 } else {
3799 // These borrows should never fail, since this struct is non-sync and this function
3800 // is non-recursive.
3801 let mut left_ref = self.left_vec.borrow_mut();
3802 let mut right_ref = self.right_vec.borrow_mut();
3803 let left_cols = left_ref.borrow_with_limit(left_row, self.limit);
3804 let right_cols = right_ref.borrow_with_limit(right_row, self.limit);
3805 compare_columns(self.order.as_ref(), &left_cols, &right_cols, || {
3806 Ordering::Equal
3807 })
3808 };
3809 // Tiebreak without the vecs borrowed, in case that recursively invokes this function.
3810 order.then_with(tiebreaker)
3811 }
3812}
3813
3814/// Compare `left` and `right` using `order`. If that doesn't produce a strict
3815/// ordering, call `tiebreaker`.
3816pub fn compare_columns<F>(
3817 order: &[ColumnOrder],
3818 left: &[Datum],
3819 right: &[Datum],
3820 tiebreaker: F,
3821) -> Ordering
3822where
3823 F: Fn() -> Ordering,
3824{
3825 for order in order {
3826 let cmp = match (&left[order.column], &right[order.column]) {
3827 (Datum::Null, Datum::Null) => Ordering::Equal,
3828 (Datum::Null, _) => {
3829 if order.nulls_last {
3830 Ordering::Greater
3831 } else {
3832 Ordering::Less
3833 }
3834 }
3835 (_, Datum::Null) => {
3836 if order.nulls_last {
3837 Ordering::Less
3838 } else {
3839 Ordering::Greater
3840 }
3841 }
3842 (lval, rval) => {
3843 if order.desc {
3844 rval.cmp(lval)
3845 } else {
3846 lval.cmp(rval)
3847 }
3848 }
3849 };
3850 if cmp != Ordering::Equal {
3851 return cmp;
3852 }
3853 }
3854 tiebreaker()
3855}
3856
3857/// Describe a window frame, e.g. `RANGE UNBOUNDED PRECEDING` or
3858/// `ROWS BETWEEN 5 PRECEDING AND CURRENT ROW`.
3859///
3860/// Window frames define a subset of the partition , and only a subset of
3861/// window functions make use of the window frame.
3862#[derive(
3863 Debug,
3864 Clone,
3865 Eq,
3866 PartialEq,
3867 Ord,
3868 PartialOrd,
3869 Serialize,
3870 Deserialize,
3871 Hash
3872)]
3873pub struct WindowFrame {
3874 /// ROWS, RANGE or GROUPS
3875 pub units: WindowFrameUnits,
3876 /// Where the frame starts
3877 pub start_bound: WindowFrameBound,
3878 /// Where the frame ends
3879 pub end_bound: WindowFrameBound,
3880}
3881
3882impl Display for WindowFrame {
3883 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
3884 write!(
3885 f,
3886 "{} between {} and {}",
3887 self.units, self.start_bound, self.end_bound
3888 )
3889 }
3890}
3891
3892impl WindowFrame {
3893 /// Return the default window frame used when one is not explicitly defined
3894 pub fn default() -> Self {
3895 WindowFrame {
3896 units: WindowFrameUnits::Range,
3897 start_bound: WindowFrameBound::UnboundedPreceding,
3898 end_bound: WindowFrameBound::CurrentRow,
3899 }
3900 }
3901
3902 fn includes_current_row(&self) -> bool {
3903 use WindowFrameBound::*;
3904 match self.start_bound {
3905 UnboundedPreceding => match self.end_bound {
3906 UnboundedPreceding => false,
3907 OffsetPreceding(0) => true,
3908 OffsetPreceding(_) => false,
3909 CurrentRow => true,
3910 OffsetFollowing(_) => true,
3911 UnboundedFollowing => true,
3912 },
3913 OffsetPreceding(0) => match self.end_bound {
3914 UnboundedPreceding => unreachable!(),
3915 OffsetPreceding(0) => true,
3916 // Any nonzero offsets here will create an empty window
3917 OffsetPreceding(_) => false,
3918 CurrentRow => true,
3919 OffsetFollowing(_) => true,
3920 UnboundedFollowing => true,
3921 },
3922 OffsetPreceding(_) => match self.end_bound {
3923 UnboundedPreceding => unreachable!(),
3924 // Window ends at the current row
3925 OffsetPreceding(0) => true,
3926 OffsetPreceding(_) => false,
3927 CurrentRow => true,
3928 OffsetFollowing(_) => true,
3929 UnboundedFollowing => true,
3930 },
3931 CurrentRow => true,
3932 OffsetFollowing(0) => match self.end_bound {
3933 UnboundedPreceding => unreachable!(),
3934 OffsetPreceding(_) => unreachable!(),
3935 CurrentRow => unreachable!(),
3936 OffsetFollowing(_) => true,
3937 UnboundedFollowing => true,
3938 },
3939 OffsetFollowing(_) => match self.end_bound {
3940 UnboundedPreceding => unreachable!(),
3941 OffsetPreceding(_) => unreachable!(),
3942 CurrentRow => unreachable!(),
3943 OffsetFollowing(_) => false,
3944 UnboundedFollowing => false,
3945 },
3946 UnboundedFollowing => false,
3947 }
3948 }
3949}
3950
3951/// Describe how frame bounds are interpreted
3952#[derive(
3953 Debug,
3954 Clone,
3955 Eq,
3956 PartialEq,
3957 Ord,
3958 PartialOrd,
3959 Serialize,
3960 Deserialize,
3961 Hash
3962)]
3963pub enum WindowFrameUnits {
3964 /// Each row is treated as the unit of work for bounds
3965 Rows,
3966 /// Each peer group is treated as the unit of work for bounds,
3967 /// and offset-based bounds use the value of the ORDER BY expression
3968 Range,
3969 /// Each peer group is treated as the unit of work for bounds.
3970 /// Groups is currently not supported, and it is rejected during planning.
3971 Groups,
3972}
3973
3974impl Display for WindowFrameUnits {
3975 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
3976 match self {
3977 WindowFrameUnits::Rows => write!(f, "rows"),
3978 WindowFrameUnits::Range => write!(f, "range"),
3979 WindowFrameUnits::Groups => write!(f, "groups"),
3980 }
3981 }
3982}
3983
3984/// Specifies [WindowFrame]'s `start_bound` and `end_bound`
3985///
3986/// The order between frame bounds is significant, as Postgres enforces
3987/// some restrictions there.
3988#[derive(
3989 Debug,
3990 Clone,
3991 Serialize,
3992 Deserialize,
3993 PartialEq,
3994 Eq,
3995 Hash,
3996 PartialOrd,
3997 Ord
3998)]
3999pub enum WindowFrameBound {
4000 /// `UNBOUNDED PRECEDING`
4001 UnboundedPreceding,
4002 /// `<N> PRECEDING`
4003 OffsetPreceding(u64),
4004 /// `CURRENT ROW`
4005 CurrentRow,
4006 /// `<N> FOLLOWING`
4007 OffsetFollowing(u64),
4008 /// `UNBOUNDED FOLLOWING`.
4009 UnboundedFollowing,
4010}
4011
4012impl Display for WindowFrameBound {
4013 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
4014 match self {
4015 WindowFrameBound::UnboundedPreceding => write!(f, "unbounded preceding"),
4016 WindowFrameBound::OffsetPreceding(offset) => write!(f, "{} preceding", offset),
4017 WindowFrameBound::CurrentRow => write!(f, "current row"),
4018 WindowFrameBound::OffsetFollowing(offset) => write!(f, "{} following", offset),
4019 WindowFrameBound::UnboundedFollowing => write!(f, "unbounded following"),
4020 }
4021 }
4022}
4023
4024/// Maximum iterations for a LetRec.
4025#[derive(
4026 Debug,
4027 Clone,
4028 Copy,
4029 PartialEq,
4030 Eq,
4031 PartialOrd,
4032 Ord,
4033 Hash,
4034 Serialize,
4035 Deserialize
4036)]
4037pub struct LetRecLimit {
4038 /// Maximum number of iterations to evaluate.
4039 pub max_iters: NonZeroU64,
4040 /// Whether to throw an error when reaching the above limit.
4041 /// If true, we simply use the current contents of each Id as the final result.
4042 pub return_at_limit: bool,
4043}
4044
4045impl LetRecLimit {
4046 /// Compute the smallest limit from a Vec of `LetRecLimit`s.
4047 pub fn min_max_iter(limits: &Vec<Option<LetRecLimit>>) -> Option<u64> {
4048 limits
4049 .iter()
4050 .filter_map(|l| l.as_ref().map(|l| l.max_iters.get()))
4051 .min()
4052 }
4053
4054 /// The default value of `LetRecLimit::return_at_limit` when using the RECURSION LIMIT option of
4055 /// WMR without ERROR AT or RETURN AT.
4056 pub const RETURN_AT_LIMIT_DEFAULT: bool = false;
4057}
4058
4059impl Display for LetRecLimit {
4060 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
4061 write!(f, "[recursion_limit={}", self.max_iters)?;
4062 if self.return_at_limit != LetRecLimit::RETURN_AT_LIMIT_DEFAULT {
4063 write!(f, ", return_at_limit")?;
4064 }
4065 write!(f, "]")
4066 }
4067}
4068
4069/// For a global Get, this indicates whether we are going to read from Persist or from an index.
4070/// (See comment in MirRelationExpr::Get.)
4071#[derive(
4072 Clone,
4073 Debug,
4074 Eq,
4075 PartialEq,
4076 Ord,
4077 PartialOrd,
4078 Serialize,
4079 Deserialize,
4080 Hash
4081)]
4082pub enum AccessStrategy {
4083 /// It's either a local Get (a CTE), or unknown at the time.
4084 /// `prune_and_annotate_dataflow_index_imports` decides it for global Gets, and thus switches to
4085 /// one of the other variants.
4086 UnknownOrLocal,
4087 /// The Get will read from Persist.
4088 Persist,
4089 /// The Get will read from an index or indexes: (index id, how the index will be used).
4090 Index(Vec<(GlobalId, IndexUsageType)>),
4091 /// The Get will read a collection that is computed by the same dataflow, but in a different
4092 /// `BuildDesc` in `objects_to_build`.
4093 SameDataflow,
4094}
4095
4096#[cfg(test)]
4097mod tests {
4098 use std::num::NonZeroUsize;
4099
4100 use mz_repr::explain::text::text_string_at;
4101
4102 use crate::explain::HumanizedExplain;
4103
4104 use super::*;
4105
4106 #[mz_ore::test]
4107 fn test_row_set_finishing_as_text() {
4108 let finishing = RowSetFinishing {
4109 order_by: vec![ColumnOrder {
4110 column: 4,
4111 desc: true,
4112 nulls_last: true,
4113 }],
4114 limit: Some(NonNeg::try_from(7).unwrap()),
4115 offset: Default::default(),
4116 project: vec![1, 3, 4, 5],
4117 };
4118
4119 let mode = HumanizedExplain::new(false);
4120 let expr = mode.expr(&finishing, None);
4121
4122 let act = text_string_at(&expr, mz_ore::str::Indent::default);
4123
4124 let exp = {
4125 use mz_ore::fmt::FormatBuffer;
4126 let mut s = String::new();
4127 write!(&mut s, "Finish");
4128 write!(&mut s, " order_by=[#4 desc nulls_last]");
4129 write!(&mut s, " limit=7");
4130 write!(&mut s, " output=[#1, #3..=#5]");
4131 writeln!(&mut s, "");
4132 s
4133 };
4134
4135 assert_eq!(act, exp);
4136 }
4137
4138 #[mz_ore::test]
4139 fn test_row_set_finishing_incremental_max_returned_query_size() {
4140 let row = Row::pack_slice(&[Datum::String("hello")]);
4141 let row_size = u64::cast_from(row.data().len());
4142 let diff = NonZeroUsize::new(1).unwrap();
4143 let batch = RowCollection::new(vec![(row, diff)], &[]);
4144
4145 // Set max_returned_query_size to hold exactly 2 batches worth of rows.
4146 let mut finishing = RowSetFinishingIncremental::new(0, None, vec![0], Some(row_size * 2));
4147
4148 let max_result_size = u64::MAX;
4149
4150 let r = finishing.finish_incremental_inner(batch.clone(), max_result_size);
4151 assert!(r.is_ok());
4152 assert_eq!(finishing.remaining_max_returned_query_size, Some(row_size));
4153
4154 let r = finishing.finish_incremental_inner(batch.clone(), max_result_size);
4155 assert!(r.is_ok());
4156 assert_eq!(finishing.remaining_max_returned_query_size, Some(0));
4157
4158 let r = finishing.finish_incremental_inner(batch, max_result_size);
4159 assert!(r.unwrap_err().contains("total result exceeds max size"));
4160 }
4161}
4162
4163/// An iterator over AST structures, which calls out nodes in difference.
4164///
4165/// The iterators visit two ASTs in tandem, continuing as long as the AST node data matches,
4166/// and yielding an output pair as soon as the AST nodes do not match. Their intent is to call
4167/// attention to the moments in the ASTs where they differ, and incidentally a stack-free way
4168/// to compare two ASTs.
4169mod structured_diff {
4170
4171 use super::MirRelationExpr;
4172 use itertools::Itertools;
4173
4174 /// An iterator over structured differences between two `MirRelationExpr` instances.
4175 pub struct MreDiff<'a> {
4176 /// Pairs of expressions that must still be compared.
4177 todo: Vec<(&'a MirRelationExpr, &'a MirRelationExpr)>,
4178 }
4179
4180 impl<'a> MreDiff<'a> {
4181 /// Create a new `MirRelationExpr` structured difference.
4182 pub fn new(expr1: &'a MirRelationExpr, expr2: &'a MirRelationExpr) -> Self {
4183 MreDiff {
4184 todo: vec![(expr1, expr2)],
4185 }
4186 }
4187 }
4188
4189 impl<'a> Iterator for MreDiff<'a> {
4190 // Pairs of expressions that do not match.
4191 type Item = (&'a MirRelationExpr, &'a MirRelationExpr);
4192
4193 fn next(&mut self) -> Option<Self::Item> {
4194 while let Some((expr1, expr2)) = self.todo.pop() {
4195 match (expr1, expr2) {
4196 (
4197 MirRelationExpr::Constant {
4198 rows: rows1,
4199 typ: typ1,
4200 },
4201 MirRelationExpr::Constant {
4202 rows: rows2,
4203 typ: typ2,
4204 },
4205 ) => {
4206 if rows1 != rows2 || typ1 != typ2 {
4207 return Some((expr1, expr2));
4208 }
4209 }
4210 (
4211 MirRelationExpr::Get {
4212 id: id1,
4213 typ: typ1,
4214 access_strategy: as1,
4215 },
4216 MirRelationExpr::Get {
4217 id: id2,
4218 typ: typ2,
4219 access_strategy: as2,
4220 },
4221 ) => {
4222 if id1 != id2 || typ1 != typ2 || as1 != as2 {
4223 return Some((expr1, expr2));
4224 }
4225 }
4226 (
4227 MirRelationExpr::Let {
4228 id: id1,
4229 body: body1,
4230 value: value1,
4231 },
4232 MirRelationExpr::Let {
4233 id: id2,
4234 body: body2,
4235 value: value2,
4236 },
4237 ) => {
4238 if id1 != id2 {
4239 return Some((expr1, expr2));
4240 } else {
4241 self.todo.push((body1, body2));
4242 self.todo.push((value1, value2));
4243 }
4244 }
4245 (
4246 MirRelationExpr::LetRec {
4247 ids: ids1,
4248 body: body1,
4249 values: values1,
4250 limits: limits1,
4251 },
4252 MirRelationExpr::LetRec {
4253 ids: ids2,
4254 body: body2,
4255 values: values2,
4256 limits: limits2,
4257 },
4258 ) => {
4259 if ids1 != ids2 || values1.len() != values2.len() || limits1 != limits2 {
4260 return Some((expr1, expr2));
4261 } else {
4262 self.todo.push((body1, body2));
4263 self.todo.extend(values1.iter().zip_eq(values2.iter()));
4264 }
4265 }
4266 (
4267 MirRelationExpr::Project {
4268 outputs: outputs1,
4269 input: input1,
4270 },
4271 MirRelationExpr::Project {
4272 outputs: outputs2,
4273 input: input2,
4274 },
4275 ) => {
4276 if outputs1 != outputs2 {
4277 return Some((expr1, expr2));
4278 } else {
4279 self.todo.push((input1, input2));
4280 }
4281 }
4282 (
4283 MirRelationExpr::Map {
4284 scalars: scalars1,
4285 input: input1,
4286 },
4287 MirRelationExpr::Map {
4288 scalars: scalars2,
4289 input: input2,
4290 },
4291 ) => {
4292 if scalars1 != scalars2 {
4293 return Some((expr1, expr2));
4294 } else {
4295 self.todo.push((input1, input2));
4296 }
4297 }
4298 (
4299 MirRelationExpr::Filter {
4300 predicates: predicates1,
4301 input: input1,
4302 },
4303 MirRelationExpr::Filter {
4304 predicates: predicates2,
4305 input: input2,
4306 },
4307 ) => {
4308 if predicates1 != predicates2 {
4309 return Some((expr1, expr2));
4310 } else {
4311 self.todo.push((input1, input2));
4312 }
4313 }
4314 (
4315 MirRelationExpr::FlatMap {
4316 input: input1,
4317 func: func1,
4318 exprs: exprs1,
4319 },
4320 MirRelationExpr::FlatMap {
4321 input: input2,
4322 func: func2,
4323 exprs: exprs2,
4324 },
4325 ) => {
4326 if func1 != func2 || exprs1 != exprs2 {
4327 return Some((expr1, expr2));
4328 } else {
4329 self.todo.push((input1, input2));
4330 }
4331 }
4332 (
4333 MirRelationExpr::Join {
4334 inputs: inputs1,
4335 equivalences: eq1,
4336 implementation: impl1,
4337 },
4338 MirRelationExpr::Join {
4339 inputs: inputs2,
4340 equivalences: eq2,
4341 implementation: impl2,
4342 },
4343 ) => {
4344 if inputs1.len() != inputs2.len() || eq1 != eq2 || impl1 != impl2 {
4345 return Some((expr1, expr2));
4346 } else {
4347 self.todo.extend(inputs1.iter().zip_eq(inputs2.iter()));
4348 }
4349 }
4350 (
4351 MirRelationExpr::Reduce {
4352 aggregates: aggregates1,
4353 input: inputs1,
4354 group_key: gk1,
4355 monotonic: m1,
4356 expected_group_size: egs1,
4357 },
4358 MirRelationExpr::Reduce {
4359 aggregates: aggregates2,
4360 input: inputs2,
4361 group_key: gk2,
4362 monotonic: m2,
4363 expected_group_size: egs2,
4364 },
4365 ) => {
4366 if aggregates1 != aggregates2 || gk1 != gk2 || m1 != m2 || egs1 != egs2 {
4367 return Some((expr1, expr2));
4368 } else {
4369 self.todo.push((inputs1, inputs2));
4370 }
4371 }
4372 (
4373 MirRelationExpr::TopK {
4374 group_key: gk1,
4375 order_key: order1,
4376 input: input1,
4377 limit: l1,
4378 offset: o1,
4379 monotonic: m1,
4380 expected_group_size: egs1,
4381 },
4382 MirRelationExpr::TopK {
4383 group_key: gk2,
4384 order_key: order2,
4385 input: input2,
4386 limit: l2,
4387 offset: o2,
4388 monotonic: m2,
4389 expected_group_size: egs2,
4390 },
4391 ) => {
4392 if order1 != order2
4393 || gk1 != gk2
4394 || l1 != l2
4395 || o1 != o2
4396 || m1 != m2
4397 || egs1 != egs2
4398 {
4399 return Some((expr1, expr2));
4400 } else {
4401 self.todo.push((input1, input2));
4402 }
4403 }
4404 (
4405 MirRelationExpr::Negate { input: input1 },
4406 MirRelationExpr::Negate { input: input2 },
4407 ) => {
4408 self.todo.push((input1, input2));
4409 }
4410 (
4411 MirRelationExpr::Threshold { input: input1 },
4412 MirRelationExpr::Threshold { input: input2 },
4413 ) => {
4414 self.todo.push((input1, input2));
4415 }
4416 (
4417 MirRelationExpr::Union {
4418 base: base1,
4419 inputs: inputs1,
4420 },
4421 MirRelationExpr::Union {
4422 base: base2,
4423 inputs: inputs2,
4424 },
4425 ) => {
4426 if inputs1.len() != inputs2.len() {
4427 return Some((expr1, expr2));
4428 } else {
4429 self.todo.push((base1, base2));
4430 self.todo.extend(inputs1.iter().zip_eq(inputs2.iter()));
4431 }
4432 }
4433 (
4434 MirRelationExpr::ArrangeBy {
4435 keys: keys1,
4436 input: input1,
4437 },
4438 MirRelationExpr::ArrangeBy {
4439 keys: keys2,
4440 input: input2,
4441 },
4442 ) => {
4443 if keys1 != keys2 {
4444 return Some((expr1, expr2));
4445 } else {
4446 self.todo.push((input1, input2));
4447 }
4448 }
4449 _ => {
4450 return Some((expr1, expr2));
4451 }
4452 }
4453 }
4454 None
4455 }
4456 }
4457}