Skip to main content

mz_expr/relation/
func.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#![allow(missing_docs)]
11
12use std::cmp::{max, min};
13use std::iter::Sum;
14use std::ops::Deref;
15use std::str::FromStr;
16use std::{fmt, iter};
17
18use chrono::{DateTime, NaiveDateTime, NaiveTime, Utc};
19use dec::OrderedDecimal;
20use itertools::{Either, Itertools};
21use mz_ore::cast::CastFrom;
22
23use mz_ore::str::separated;
24use mz_ore::{soft_assert_eq_no_log, soft_assert_or_log};
25use mz_repr::adt::array::ArrayDimension;
26use mz_repr::adt::date::Date;
27use mz_repr::adt::interval::Interval;
28use mz_repr::adt::numeric::{self, Numeric, NumericMaxScale};
29use mz_repr::adt::regex::{Regex as ReprRegex, RegexCompilationError};
30use mz_repr::adt::timestamp::{CheckedTimestamp, TimestampLike};
31use mz_repr::{
32    ColumnName, Datum, Diff, ReprColumnType, ReprRelationType, Row, RowArena, RowPacker, SharedRow,
33    SqlColumnType, SqlRelationType, SqlScalarType, datum_size,
34};
35use num::{CheckedAdd, Integer, Signed, ToPrimitive};
36use ordered_float::OrderedFloat;
37use regex::Regex;
38use serde::{Deserialize, Serialize};
39use smallvec::SmallVec;
40
41use crate::EvalError;
42use crate::WindowFrameBound::{
43    CurrentRow, OffsetFollowing, OffsetPreceding, UnboundedFollowing, UnboundedPreceding,
44};
45use crate::WindowFrameUnits::{Groups, Range, Rows};
46use crate::explain::{HumanizedExpr, HumanizerMode};
47use crate::relation::{
48    ColumnOrder, WindowFrame, WindowFrameBound, WindowFrameUnits, compare_columns,
49};
50use crate::scalar::func::{add_timestamp_months, jsonb_stringify};
51
52// TODO(jamii) be careful about overflow in sum/avg
53// see https://timely.zulipchat.com/#narrow/stream/186635-engineering/topic/additional.20work/near/163507435
54
55fn max_string<'a, I>(datums: I) -> Datum<'a>
56where
57    I: IntoIterator<Item = Datum<'a>>,
58{
59    match datums
60        .into_iter()
61        .filter(|d| !d.is_null())
62        .max_by(|a, b| a.unwrap_str().cmp(b.unwrap_str()))
63    {
64        Some(datum) => datum,
65        None => Datum::Null,
66    }
67}
68
69fn max_datum<'a, I, DatumType>(datums: I) -> Datum<'a>
70where
71    I: IntoIterator<Item = Datum<'a>>,
72    DatumType: TryFrom<Datum<'a>> + Ord,
73    <DatumType as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
74    Datum<'a>: From<Option<DatumType>>,
75{
76    let x: Option<DatumType> = datums
77        .into_iter()
78        .filter(|d| !d.is_null())
79        .map(|d| DatumType::try_from(d).expect("unexpected type"))
80        .max();
81
82    x.into()
83}
84
85fn min_datum<'a, I, DatumType>(datums: I) -> Datum<'a>
86where
87    I: IntoIterator<Item = Datum<'a>>,
88    DatumType: TryFrom<Datum<'a>> + Ord,
89    <DatumType as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
90    Datum<'a>: From<Option<DatumType>>,
91{
92    let x: Option<DatumType> = datums
93        .into_iter()
94        .filter(|d| !d.is_null())
95        .map(|d| DatumType::try_from(d).expect("unexpected type"))
96        .min();
97
98    x.into()
99}
100
101fn min_string<'a, I>(datums: I) -> Datum<'a>
102where
103    I: IntoIterator<Item = Datum<'a>>,
104{
105    match datums
106        .into_iter()
107        .filter(|d| !d.is_null())
108        .min_by(|a, b| a.unwrap_str().cmp(b.unwrap_str()))
109    {
110        Some(datum) => datum,
111        None => Datum::Null,
112    }
113}
114
115fn sum_datum<'a, I, DatumType, ResultType>(datums: I) -> Datum<'a>
116where
117    I: IntoIterator<Item = Datum<'a>>,
118    DatumType: TryFrom<Datum<'a>>,
119    <DatumType as TryFrom<Datum<'a>>>::Error: std::fmt::Debug,
120    ResultType: From<DatumType> + Sum + Into<Datum<'a>>,
121{
122    let mut datums = datums.into_iter().filter(|d| !d.is_null()).peekable();
123    if datums.peek().is_none() {
124        Datum::Null
125    } else {
126        let x = datums
127            .map(|d| ResultType::from(DatumType::try_from(d).expect("unexpected type")))
128            .sum::<ResultType>();
129        x.into()
130    }
131}
132
133/// Count-aware signed-integer sum. Accumulates `Σ value·diff` in `i128`, which
134/// matches the width of the dataflow's `Accum::SimpleNumber` accumulator (see
135/// `build_accumulable` and `finalize_accum` in `mz_compute::render::reduce`);
136/// `narrow` then reproduces that variant's `finalize_accum` arm. Unlike
137/// `expand_counts`, this consumes the multiplicity directly, so it is linear in
138/// the number of distinct values and correct for negative diffs (retractions),
139/// which `expand_counts` would silently drop.
140///
141/// Returns `Datum::Null` when no non-null value was accumulated, matching
142/// `finalize_accum`'s null handling: its `is_zero` check on `SimpleNumber`
143/// requires both a zero running sum and a zero non-null count.
144fn sum_signed_int_counted<'a, I, N>(datums: I, narrow: N) -> Datum<'a>
145where
146    I: IntoIterator<Item = (Datum<'a>, Diff)>,
147    N: FnOnce(i128) -> Datum<'a>,
148{
149    let mut accum: i128 = 0;
150    let mut non_nulls = Diff::ZERO;
151    for (datum, diff) in datums {
152        if datum.is_null() {
153            continue;
154        }
155        let value = match datum {
156            Datum::Int16(i) => i128::from(i),
157            Datum::Int32(i) => i128::from(i),
158            Datum::Int64(i) => i128::from(i),
159            other => panic!("unexpected non-integer datum in signed sum: {other:?}"),
160        };
161        // The dataflow accumulates `value * diff` in an `Overflowing<i128>`; we
162        // mirror that. Genuine i128 overflow would require summands far beyond
163        // any realistic input, so wrapping matches the dataflow's production
164        // behavior.
165        accum = accum.wrapping_add(value.wrapping_mul(i128::from(diff.into_inner())));
166        non_nulls += diff;
167    }
168    if accum == 0 && non_nulls.is_zero() {
169        Datum::Null
170    } else {
171        narrow(accum)
172    }
173}
174
175fn sum_numeric<'a, I>(datums: I) -> Datum<'a>
176where
177    I: IntoIterator<Item = Datum<'a>>,
178{
179    let mut cx = numeric::cx_datum();
180    let mut sum = Numeric::zero();
181    let mut empty = true;
182    for d in datums {
183        if !d.is_null() {
184            empty = false;
185            cx.add(&mut sum, &d.unwrap_numeric().0);
186        }
187    }
188    match empty {
189        true => Datum::Null,
190        false => Datum::from(sum),
191    }
192}
193
194/// Sums intervals component-wise, as PostgreSQL's interval addition does:
195/// months, days, and microseconds each accumulate on their own, and nothing is
196/// carried from a coarser component into a finer one.
197fn sum_interval<'a, I>(datums: I) -> Datum<'a>
198where
199    I: IntoIterator<Item = Datum<'a>>,
200{
201    sum_interval_counted(datums.into_iter().map(|datum| (datum, Diff::ONE)))
202}
203
204/// Count-aware interval sum. Accumulates `Σ value·diff` per interval component
205/// in `i128`, which matches `Accum::Interval` in `mz_compute::render::reduce`;
206/// the narrowing back to the `Interval` field widths reproduces that variant's
207/// `finalize_accum` arm. Consuming the multiplicity directly keeps this linear
208/// in the number of distinct values and correct for negative diffs
209/// (retractions), which `expand_counts` would silently drop.
210///
211/// A sum whose components exceed the `Interval` field widths is a query error
212/// only where the aggregate renders as an accumulable reduce, because that is
213/// the only rendering that pairs `finalize_accum` with `AccumulableErrorCheck`.
214/// This function and `finalize_accum` both return a bare `Datum`, so every
215/// other caller (constant folding, and window aggregates, which arrive through
216/// `eval` for whole-partition frames and through `AccumulableOneByOneAggr`
217/// otherwise) yields `Datum::Null` on overflow, which a reader cannot tell
218/// apart from an all-null input. `SumUInt*` has the same split today for
219/// negative accumulations. Closing it means giving aggregate evaluation an
220/// error channel.
221///
222/// Returns `Datum::Null` when no non-null value was accumulated, matching
223/// `finalize_accum`'s null handling: its `is_zero` check on `Accum::Interval`
224/// requires all three running sums and the non-null count to be zero.
225fn sum_interval_counted<'a, I>(datums: I) -> Datum<'a>
226where
227    I: IntoIterator<Item = (Datum<'a>, Diff)>,
228{
229    let (mut months, mut days, mut micros) = (0i128, 0i128, 0i128);
230    let mut non_nulls = Diff::ZERO;
231    for (datum, diff) in datums {
232        if datum.is_null() {
233            continue;
234        }
235        let interval = datum.unwrap_interval();
236        // The dataflow accumulates each component in an `Overflowing<i128>`; we
237        // mirror that. Genuine i128 overflow would require summands far beyond
238        // any realistic input, so wrapping matches the dataflow's production
239        // behavior.
240        let scale = i128::from(diff.into_inner());
241        months = months.wrapping_add(i128::from(interval.months).wrapping_mul(scale));
242        days = days.wrapping_add(i128::from(interval.days).wrapping_mul(scale));
243        micros = micros.wrapping_add(i128::from(interval.micros).wrapping_mul(scale));
244        non_nulls += diff;
245    }
246    if months == 0 && days == 0 && micros == 0 && non_nulls.is_zero() {
247        return Datum::Null;
248    }
249    match Interval::try_new(months, days, micros) {
250        Some(interval) => Datum::Interval(interval),
251        None => Datum::Null,
252    }
253}
254
255fn count<'a, I>(datums: I) -> Datum<'a>
256where
257    I: IntoIterator<Item = (Datum<'a>, Diff)>,
258{
259    // Count is accumulable: rather than expand each `(datum, diff)` into `diff`
260    // copies and count them, we sum the diffs directly. A net-negative count is
261    // possible (the surface does not define behavior in that case) and surfaces
262    // here as a negative result.
263    // TODO(jkosh44) This should error when the count can't fit inside of an `i64` instead of returning a negative result.
264    let mut count = Diff::ZERO;
265    for (datum, diff) in datums {
266        if !datum.is_null() {
267            count += diff;
268        }
269    }
270    Datum::from(count.into_inner())
271}
272
273fn any<'a, I>(datums: I) -> Datum<'a>
274where
275    I: IntoIterator<Item = Datum<'a>>,
276{
277    datums
278        .into_iter()
279        .fold(Datum::False, |state, next| match (state, next) {
280            (Datum::True, _) | (_, Datum::True) => Datum::True,
281            (Datum::Null, _) | (_, Datum::Null) => Datum::Null,
282            _ => Datum::False,
283        })
284}
285
286fn all<'a, I>(datums: I) -> Datum<'a>
287where
288    I: IntoIterator<Item = Datum<'a>>,
289{
290    datums
291        .into_iter()
292        .fold(Datum::True, |state, next| match (state, next) {
293            (Datum::False, _) | (_, Datum::False) => Datum::False,
294            (Datum::Null, _) | (_, Datum::Null) => Datum::Null,
295            _ => Datum::True,
296        })
297}
298
299fn string_agg<'a, I>(datums: I, temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
300where
301    I: IntoIterator<Item = Datum<'a>>,
302{
303    const EMPTY_SEP: &str = "";
304
305    let datums = order_aggregate_datums(datums, order_by);
306    let mut sep_value_pairs = datums.into_iter().filter_map(|d| {
307        if d.is_null() {
308            return None;
309        }
310        let mut value_sep = d.unwrap_list().iter();
311        match (value_sep.next().unwrap(), value_sep.next().unwrap()) {
312            (Datum::Null, _) => None,
313            (Datum::String(val), Datum::Null) => Some((EMPTY_SEP, val)),
314            (Datum::String(val), Datum::String(sep)) => Some((sep, val)),
315            _ => unreachable!(),
316        }
317    });
318
319    let mut s = String::default();
320    match sep_value_pairs.next() {
321        // First value not prefixed by its separator
322        Some((_, value)) => s.push_str(value),
323        // If no non-null values sent, return NULL.
324        None => return Datum::Null,
325    }
326
327    for (sep, value) in sep_value_pairs {
328        s.push_str(sep);
329        s.push_str(value);
330    }
331
332    Datum::String(temp_storage.push_string(s))
333}
334
335fn jsonb_agg<'a, I>(datums: I, temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
336where
337    I: IntoIterator<Item = Datum<'a>>,
338{
339    let datums = order_aggregate_datums(datums, order_by);
340    temp_storage.make_datum(|packer| {
341        packer.push_list(datums.into_iter().filter(|d| !d.is_null()));
342    })
343}
344
345fn dict_agg<'a, I>(datums: I, temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
346where
347    I: IntoIterator<Item = Datum<'a>>,
348{
349    let datums = order_aggregate_datums(datums, order_by);
350    temp_storage.make_datum(|packer| {
351        let mut datums: Vec<_> = datums
352            .into_iter()
353            .filter_map(|d| {
354                if d.is_null() {
355                    return None;
356                }
357                let mut list = d.unwrap_list().iter();
358                let key = list.next().unwrap();
359                let val = list.next().unwrap();
360                if key.is_null() {
361                    // TODO(benesch): this should produce an error, but
362                    // aggregate functions cannot presently produce errors.
363                    None
364                } else {
365                    Some((key.unwrap_str(), val))
366                }
367            })
368            .collect();
369        // datums are ordered by any ORDER BY clause now, and we want to preserve
370        // the last entry for each key, but we also need to present unique and sorted
371        // keys to push_dict. Use sort_by here, which is stable, and so will preserve
372        // the ORDER BY order. Then reverse and dedup to retain the last of each
373        // key. Reverse again so we're back in push_dict order.
374        datums.sort_by_key(|(k, _v)| *k);
375        datums.reverse();
376        datums.dedup_by_key(|(k, _v)| *k);
377        datums.reverse();
378        packer.push_dict(datums);
379    })
380}
381
382/// Assuming datums is a List, sort them by the 2nd through Nth elements
383/// corresponding to order_by, then return the 1st element.
384///
385/// Near the usages of this function, we sometimes want to produce Datums with a shorter lifetime
386/// than 'a. We have to actually perform the shortening of the lifetime here, inside this function,
387/// because if we were to simply return `impl Iterator<Item = Datum<'a>>`, that wouldn't be
388/// covariant in the item type, because opaque types are always invariant. (Contrast this with how
389/// we perform the shortening _inside_ this function: the input of the `map` is known to
390/// specifically be `std::vec::IntoIter`, which is known to be covariant.)
391pub fn order_aggregate_datums<'a: 'b, 'b, I>(
392    datums: I,
393    order_by: &[ColumnOrder],
394) -> impl Iterator<Item = Datum<'b>>
395where
396    I: IntoIterator<Item = Datum<'a>>,
397{
398    order_aggregate_datums_with_rank_inner(datums, order_by)
399        .into_iter()
400        // (`payload` is coerced here to `Datum<'b>` in the argument of the closure)
401        .map(|(payload, _order_datums)| payload)
402}
403
404/// Assuming datums is a List, sort them by the 2nd through Nth elements
405/// corresponding to order_by, then return the 1st element and computed order by expression.
406fn order_aggregate_datums_with_rank<'a, I>(
407    datums: I,
408    order_by: &[ColumnOrder],
409) -> impl Iterator<Item = (Datum<'a>, Row)>
410where
411    I: IntoIterator<Item = Datum<'a>>,
412{
413    order_aggregate_datums_with_rank_inner(datums, order_by)
414        .into_iter()
415        .map(|(payload, order_by_datums)| (payload, Row::pack(order_by_datums)))
416}
417
418fn order_aggregate_datums_with_rank_inner<'a, I>(
419    datums: I,
420    order_by: &[ColumnOrder],
421) -> Vec<(Datum<'a>, Vec<Datum<'a>>)>
422where
423    I: IntoIterator<Item = Datum<'a>>,
424{
425    let mut decoded: Vec<(Datum, Vec<Datum>)> = datums
426        .into_iter()
427        .map(|d| {
428            let list = d.unwrap_list();
429            let mut list_it = list.iter();
430            let payload = list_it.next().unwrap();
431
432            // We decode the order_by Datums here instead of the comparison function, because the
433            // comparison function is expected to be called `O(log n)` times on each input row.
434            // The only downside is that the decoded data might be bigger, but I think that's fine,
435            // because:
436            // - if we have a window partition so big that this would create a memory problem, then
437            //   the non-incrementalness of window functions will create a serious CPU problem
438            //   anyway,
439            // - and anyhow various other parts of the window function code already do decoding
440            //   upfront.
441            let mut order_by_datums = Vec::with_capacity(order_by.len());
442            for _ in 0..order_by.len() {
443                order_by_datums.push(
444                    list_it
445                        .next()
446                        .expect("must have exactly the same number of Datums as `order_by`"),
447                );
448            }
449
450            (payload, order_by_datums)
451        })
452        .collect();
453
454    let mut sort_by =
455        |(payload_left, left_order_by_datums): &(Datum, Vec<Datum>),
456         (payload_right, right_order_by_datums): &(Datum, Vec<Datum>)| {
457            compare_columns(
458                order_by,
459                left_order_by_datums,
460                right_order_by_datums,
461                || payload_left.cmp(payload_right),
462            )
463        };
464    // `sort_unstable_by` can be faster and uses less memory than `sort_by`. An unstable sort is
465    // enough here, because if two elements are equal in our `compare` function, then the elements
466    // are actually binary-equal (because of the `tiebreaker` given to `compare_columns`), so it
467    // doesn't matter what order they end up in.
468    decoded.sort_unstable_by(&mut sort_by);
469    decoded
470}
471
472fn array_concat<'a, I>(datums: I, temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
473where
474    I: IntoIterator<Item = Datum<'a>>,
475{
476    let datums = order_aggregate_datums(datums, order_by);
477    let datums: Vec<_> = datums
478        .into_iter()
479        .map(|d| d.unwrap_array().elements().iter())
480        .flatten()
481        .collect();
482    let dims = ArrayDimension {
483        lower_bound: 1,
484        length: datums.len(),
485    };
486    temp_storage.make_datum(|packer| {
487        packer.try_push_array(&[dims], datums).unwrap();
488    })
489}
490
491fn list_concat<'a, I>(datums: I, temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
492where
493    I: IntoIterator<Item = Datum<'a>>,
494{
495    let datums = order_aggregate_datums(datums, order_by);
496    temp_storage.make_datum(|packer| {
497        packer.push_list(datums.into_iter().map(|d| d.unwrap_list().iter()).flatten());
498    })
499}
500
501/// The expected input is in the format of `[((OriginalRow, [EncodedArgs]), OrderByExprs...)]`
502/// The output is in the format of `[result_value, original_row]`.
503/// See an example at `lag_lead`, where the input-output formats are similar.
504fn row_number<'a, I>(
505    datums: I,
506    callers_temp_storage: &'a RowArena,
507    order_by: &[ColumnOrder],
508) -> Datum<'a>
509where
510    I: IntoIterator<Item = Datum<'a>>,
511{
512    // We want to use our own temp_storage here, to avoid flooding `callers_temp_storage` with a
513    // large number of new datums. This is because we don't want to make an assumption about
514    // whether the caller creates a new temp_storage between window partitions.
515    let temp_storage = RowArena::new();
516    let datums = row_number_no_list(datums, &temp_storage, order_by);
517
518    callers_temp_storage.make_datum(|packer| {
519        packer.push_list(datums);
520    })
521}
522
523/// Like `row_number`, but doesn't perform the final wrapping in a list, returning an Iterator
524/// instead.
525fn row_number_no_list<'a: 'b, 'b, I>(
526    datums: I,
527    callers_temp_storage: &'b RowArena,
528    order_by: &[ColumnOrder],
529) -> impl Iterator<Item = Datum<'b>>
530where
531    I: IntoIterator<Item = Datum<'a>>,
532{
533    let datums = order_aggregate_datums(datums, order_by);
534
535    callers_temp_storage.reserve(datums.size_hint().0);
536    #[allow(clippy::disallowed_methods)]
537    datums
538        .into_iter()
539        .map(|d| d.unwrap_list().iter())
540        .flatten()
541        .zip(1i64..)
542        .map(|(d, i)| {
543            callers_temp_storage.make_datum(|packer| {
544                packer.push_list_with(|packer| {
545                    packer.push(Datum::Int64(i));
546                    packer.push(d);
547                });
548            })
549        })
550}
551
552/// The expected input is in the format of `[((OriginalRow, [EncodedArgs]), OrderByExprs...)]`
553/// The output is in the format of `[result_value, original_row]`.
554/// See an example at `lag_lead`, where the input-output formats are similar.
555fn rank<'a, I>(datums: I, callers_temp_storage: &'a RowArena, order_by: &[ColumnOrder]) -> Datum<'a>
556where
557    I: IntoIterator<Item = Datum<'a>>,
558{
559    let temp_storage = RowArena::new();
560    let datums = rank_no_list(datums, &temp_storage, order_by);
561
562    callers_temp_storage.make_datum(|packer| {
563        packer.push_list(datums);
564    })
565}
566
567/// Like `rank`, but doesn't perform the final wrapping in a list, returning an Iterator
568/// instead.
569fn rank_no_list<'a: 'b, 'b, I>(
570    datums: I,
571    callers_temp_storage: &'b RowArena,
572    order_by: &[ColumnOrder],
573) -> impl Iterator<Item = Datum<'b>>
574where
575    I: IntoIterator<Item = Datum<'a>>,
576{
577    // Keep the row used for ordering around, as it is used to determine the rank
578    let datums = order_aggregate_datums_with_rank(datums, order_by);
579
580    let mut datums = datums
581        .into_iter()
582        .map(|(d0, order_row)| {
583            d0.unwrap_list()
584                .iter()
585                .map(move |d1| (d1, order_row.clone()))
586        })
587        .flatten();
588
589    callers_temp_storage.reserve(datums.size_hint().0);
590    datums
591        .next()
592        .map_or(vec![], |(first_datum, first_order_row)| {
593            // Folding with (last order_by row, last assigned rank,
594            // row number, output vec)
595            datums.fold(
596                (first_order_row, 1, 1, vec![(first_datum, 1)]),
597                |mut acc, (next_datum, next_order_row)| {
598                let (ref mut acc_row, ref mut acc_rank, ref mut acc_row_num, ref mut output) = acc;
599                *acc_row_num += 1;
600                // Identity is based on the order_by expression
601                if *acc_row != next_order_row {
602                    *acc_rank = *acc_row_num;
603                    *acc_row = next_order_row;
604                }
605
606                (*output).push((next_datum, *acc_rank));
607                acc
608            })
609        }.3).into_iter().map(|(d, i)| {
610        callers_temp_storage.make_datum(|packer| {
611            packer.push_list_with(|packer| {
612                packer.push(Datum::Int64(i));
613                packer.push(d);
614            });
615        })
616    })
617}
618
619/// The expected input is in the format of `[((OriginalRow, [EncodedArgs]), OrderByExprs...)]`
620/// The output is in the format of `[result_value, original_row]`.
621/// See an example at `lag_lead`, where the input-output formats are similar.
622fn dense_rank<'a, I>(
623    datums: I,
624    callers_temp_storage: &'a RowArena,
625    order_by: &[ColumnOrder],
626) -> Datum<'a>
627where
628    I: IntoIterator<Item = Datum<'a>>,
629{
630    let temp_storage = RowArena::new();
631    let datums = dense_rank_no_list(datums, &temp_storage, order_by);
632
633    callers_temp_storage.make_datum(|packer| {
634        packer.push_list(datums);
635    })
636}
637
638/// Like `dense_rank`, but doesn't perform the final wrapping in a list, returning an Iterator
639/// instead.
640fn dense_rank_no_list<'a: 'b, 'b, I>(
641    datums: I,
642    callers_temp_storage: &'b RowArena,
643    order_by: &[ColumnOrder],
644) -> impl Iterator<Item = Datum<'b>>
645where
646    I: IntoIterator<Item = Datum<'a>>,
647{
648    // Keep the row used for ordering around, as it is used to determine the rank
649    let datums = order_aggregate_datums_with_rank(datums, order_by);
650
651    let mut datums = datums
652        .into_iter()
653        .map(|(d0, order_row)| {
654            d0.unwrap_list()
655                .iter()
656                .map(move |d1| (d1, order_row.clone()))
657        })
658        .flatten();
659
660    callers_temp_storage.reserve(datums.size_hint().0);
661    datums
662        .next()
663        .map_or(vec![], |(first_datum, first_order_row)| {
664            // Folding with (last order_by row, last assigned rank,
665            // output vec)
666            datums.fold(
667                (first_order_row, 1, vec![(first_datum, 1)]),
668                |mut acc, (next_datum, next_order_row)| {
669                let (ref mut acc_row, ref mut acc_rank, ref mut output) = acc;
670                // Identity is based on the order_by expression
671                if *acc_row != next_order_row {
672                    *acc_rank += 1;
673                    *acc_row = next_order_row;
674                }
675
676                (*output).push((next_datum, *acc_rank));
677                acc
678            })
679        }.2).into_iter().map(|(d, i)| {
680        callers_temp_storage.make_datum(|packer| {
681            packer.push_list_with(|packer| {
682                packer.push(Datum::Int64(i));
683                packer.push(d);
684            });
685        })
686    })
687}
688
689/// The expected input is in the format of `[((OriginalRow, EncodedArgs), OrderByExprs...)]`
690/// For example,
691///
692/// lag(x*y, 1, null) over (partition by x+y order by x-y, x/y)
693///
694/// list of:
695/// row(
696///   row(
697///     row(#0, #1),
698///     row((#0 * #1), 1, null)
699///   ),
700///   (#0 - #1),
701///   (#0 / #1)
702/// )
703///
704/// The output is in the format of `[result_value, original_row]`, e.g.
705/// list of:
706/// row(
707///   42,
708///   row(7, 8)
709/// )
710fn lag_lead<'a, I>(
711    datums: I,
712    callers_temp_storage: &'a RowArena,
713    order_by: &[ColumnOrder],
714    lag_lead_type: &LagLeadType,
715    ignore_nulls: &bool,
716) -> Datum<'a>
717where
718    I: IntoIterator<Item = Datum<'a>>,
719{
720    let temp_storage = RowArena::new();
721    let iter = lag_lead_no_list(datums, &temp_storage, order_by, lag_lead_type, ignore_nulls);
722    callers_temp_storage.make_datum(|packer| {
723        packer.push_list(iter);
724    })
725}
726
727/// Like `lag_lead`, but doesn't perform the final wrapping in a list, returning an Iterator
728/// instead.
729fn lag_lead_no_list<'a: 'b, 'b, I>(
730    datums: I,
731    callers_temp_storage: &'b RowArena,
732    order_by: &[ColumnOrder],
733    lag_lead_type: &LagLeadType,
734    ignore_nulls: &bool,
735) -> impl Iterator<Item = Datum<'b>>
736where
737    I: IntoIterator<Item = Datum<'a>>,
738{
739    // Sort the datums according to the ORDER BY expressions and return the (OriginalRow, EncodedArgs) record
740    let datums = order_aggregate_datums(datums, order_by);
741
742    // Take the (OriginalRow, EncodedArgs) records and unwrap them into separate datums.
743    // EncodedArgs = (InputValue, Offset, DefaultValue) for Lag/Lead
744    // (`OriginalRow` is kept in a record form, as we don't need to look inside that.)
745    let (orig_rows, unwrapped_args): (Vec<_>, Vec<_>) = datums
746        .into_iter()
747        .map(|d| {
748            let mut iter = d.unwrap_list().iter();
749            let original_row = iter.next().unwrap();
750            let (input_value, offset, default_value) =
751                unwrap_lag_lead_encoded_args(iter.next().unwrap());
752            (original_row, (input_value, offset, default_value))
753        })
754        .unzip();
755
756    let result = lag_lead_inner(unwrapped_args, lag_lead_type, ignore_nulls);
757
758    callers_temp_storage.reserve(result.len());
759    result
760        .into_iter()
761        .zip_eq(orig_rows)
762        .map(|(result_value, original_row)| {
763            callers_temp_storage.make_datum(|packer| {
764                packer.push_list_with(|packer| {
765                    packer.push(result_value);
766                    packer.push(original_row);
767                });
768            })
769        })
770}
771
772/// lag/lead's arguments are in a record. This function unwraps this record.
773fn unwrap_lag_lead_encoded_args(encoded_args: Datum) -> (Datum, Datum, Datum) {
774    let mut encoded_args_iter = encoded_args.unwrap_list().iter();
775    let (input_value, offset, default_value) = (
776        encoded_args_iter.next().unwrap(),
777        encoded_args_iter.next().unwrap(),
778        encoded_args_iter.next().unwrap(),
779    );
780    (input_value, offset, default_value)
781}
782
783/// Each element of `args` has the 3 arguments evaluated for a single input row.
784/// Returns the results for each input row.
785fn lag_lead_inner<'a>(
786    args: Vec<(Datum<'a>, Datum<'a>, Datum<'a>)>,
787    lag_lead_type: &LagLeadType,
788    ignore_nulls: &bool,
789) -> Vec<Datum<'a>> {
790    if *ignore_nulls {
791        lag_lead_inner_ignore_nulls(args, lag_lead_type)
792    } else {
793        lag_lead_inner_respect_nulls(args, lag_lead_type)
794    }
795}
796
797fn lag_lead_inner_respect_nulls<'a>(
798    args: Vec<(Datum<'a>, Datum<'a>, Datum<'a>)>,
799    lag_lead_type: &LagLeadType,
800) -> Vec<Datum<'a>> {
801    let mut result: Vec<Datum> = Vec::with_capacity(args.len());
802    for (idx, (_, offset, default_value)) in args.iter().enumerate() {
803        // Null offsets are acceptable, and always return null
804        if offset.is_null() {
805            result.push(Datum::Null);
806            continue;
807        }
808
809        let idx = i64::try_from(idx).expect("Array index does not fit in i64");
810        let offset = i64::from(offset.unwrap_int32());
811        let offset = match lag_lead_type {
812            LagLeadType::Lag => -offset,
813            LagLeadType::Lead => offset,
814        };
815
816        // Get a Datum from `datums`. Return None if index is out of range.
817        let datums_get = |i: i64| -> Option<Datum> {
818            match u64::try_from(i) {
819                Ok(i) => args
820                    .get(usize::cast_from(i))
821                    .map(|d| Some(d.0)) // succeeded in getting a Datum from the vec
822                    .unwrap_or(None), // overindexing
823                Err(_) => None, // underindexing (negative index)
824            }
825        };
826
827        let lagged_value = datums_get(idx + offset).unwrap_or(*default_value);
828
829        result.push(lagged_value);
830    }
831
832    result
833}
834
835// `i64` indexes get involved in this function because it's convenient to allow negative indexes and
836// have `datums_get` fail on them, and thus handle the beginning and end of the input vector
837// uniformly, rather than checking underflow separately during index manipulations.
838#[allow(clippy::as_conversions)]
839fn lag_lead_inner_ignore_nulls<'a>(
840    args: Vec<(Datum<'a>, Datum<'a>, Datum<'a>)>,
841    lag_lead_type: &LagLeadType,
842) -> Vec<Datum<'a>> {
843    // We check here once that even the largest index fits in `i64`, and then do silent `as`
844    // conversions from `usize` indexes to `i64` indexes throughout this function.
845    if i64::try_from(args.len()).is_err() {
846        panic!("window partition way too big")
847    }
848    // Preparation: Make sure we can jump over a run of nulls in constant time, i.e., regardless of
849    // how many nulls the run has. The following skip tables will point to the next non-null index.
850    let mut skip_nulls_backward = vec![None; args.len()];
851    let mut last_non_null: i64 = -1;
852    let pairs = args
853        .iter()
854        .enumerate()
855        .zip_eq(skip_nulls_backward.iter_mut());
856    for ((i, (d, _, _)), slot) in pairs {
857        if d.is_null() {
858            *slot = Some(last_non_null);
859        } else {
860            last_non_null = i as i64;
861        }
862    }
863    let mut skip_nulls_forward = vec![None; args.len()];
864    let mut last_non_null: i64 = args.len() as i64;
865    let pairs = args
866        .iter()
867        .enumerate()
868        .rev()
869        .zip_eq(skip_nulls_forward.iter_mut().rev());
870    for ((i, (d, _, _)), slot) in pairs {
871        if d.is_null() {
872            *slot = Some(last_non_null);
873        } else {
874            last_non_null = i as i64;
875        }
876    }
877
878    // The actual computation.
879    let mut result: Vec<Datum> = Vec::with_capacity(args.len());
880    for (idx, (_, offset, default_value)) in args.iter().enumerate() {
881        // Null offsets are acceptable, and always return null
882        if offset.is_null() {
883            result.push(Datum::Null);
884            continue;
885        }
886
887        let idx = idx as i64; // checked at the beginning of the function that len() fits
888        let offset = i64::cast_from(offset.unwrap_int32());
889        let offset = match lag_lead_type {
890            LagLeadType::Lag => -offset,
891            LagLeadType::Lead => offset,
892        };
893        let increment = offset.signum();
894
895        // Get a Datum from `datums`. Return None if index is out of range.
896        let datums_get = |i: i64| -> Option<Datum> {
897            match u64::try_from(i) {
898                Ok(i) => args
899                    .get(usize::cast_from(i))
900                    .map(|d| Some(d.0)) // succeeded in getting a Datum from the vec
901                    .unwrap_or(None), // overindexing
902                Err(_) => None, // underindexing (negative index)
903            }
904        };
905
906        let lagged_value = if increment != 0 {
907            // We start j from idx, and step j until we have seen an abs(offset) number of non-null
908            // values or reach the beginning or end of the partition.
909            //
910            // If offset is big, then this is slow: Considering the entire function, it's
911            // `O(partition_size * offset)`.
912            // However, a common use case is an offset of 1, for which this doesn't matter.
913            // TODO: For larger offsets, we could have a completely different implementation
914            // that starts the inner loop from the index where we found the previous result:
915            // https://github.com/MaterializeInc/materialize/pull/29287#discussion_r1738695174
916            let mut j = idx;
917            for _ in 0..num::abs(offset) {
918                j += increment;
919                // Jump over a run of nulls
920                if datums_get(j).is_some_and(|d| d.is_null()) {
921                    let ju = j as usize; // `j >= 0` because of the above `is_some_and`
922                    if increment > 0 {
923                        j = skip_nulls_forward[ju].expect("checked above that it's null");
924                    } else {
925                        j = skip_nulls_backward[ju].expect("checked above that it's null");
926                    }
927                }
928                if datums_get(j).is_none() {
929                    break;
930                }
931            }
932            match datums_get(j) {
933                Some(datum) => datum,
934                None => *default_value,
935            }
936        } else {
937            assert_eq!(offset, 0);
938            let datum = datums_get(idx).expect("known to exist");
939            if !datum.is_null() {
940                datum
941            } else {
942                // Not clear what should the semantics be here. See
943                // https://github.com/MaterializeInc/database-issues/issues/8497
944                // (We used to run into an infinite loop in this case, so panicking is
945                // better.)
946                panic!("0 offset in lag/lead IGNORE NULLS");
947            }
948        };
949
950        result.push(lagged_value);
951    }
952
953    result
954}
955
956/// The expected input is in the format of [((OriginalRow, InputValue), OrderByExprs...)]
957fn first_value<'a, I>(
958    datums: I,
959    callers_temp_storage: &'a RowArena,
960    order_by: &[ColumnOrder],
961    window_frame: &WindowFrame,
962) -> Datum<'a>
963where
964    I: IntoIterator<Item = Datum<'a>>,
965{
966    let temp_storage = RowArena::new();
967    let iter = first_value_no_list(datums, &temp_storage, order_by, window_frame);
968    callers_temp_storage.make_datum(|packer| {
969        packer.push_list(iter);
970    })
971}
972
973/// Like `first_value`, but doesn't perform the final wrapping in a list, returning an Iterator
974/// instead.
975fn first_value_no_list<'a: 'b, 'b, I>(
976    datums: I,
977    callers_temp_storage: &'b RowArena,
978    order_by: &[ColumnOrder],
979    window_frame: &WindowFrame,
980) -> impl Iterator<Item = Datum<'b>>
981where
982    I: IntoIterator<Item = Datum<'a>>,
983{
984    // Sort the datums according to the ORDER BY expressions and return the (OriginalRow, InputValue) record
985    let datums = order_aggregate_datums(datums, order_by);
986
987    // Decode the input (OriginalRow, InputValue) into separate datums
988    let (orig_rows, args): (Vec<_>, Vec<_>) = datums
989        .into_iter()
990        .map(|d| {
991            let mut iter = d.unwrap_list().iter();
992            let original_row = iter.next().unwrap();
993            let arg = iter.next().unwrap();
994
995            (original_row, arg)
996        })
997        .unzip();
998
999    let results = first_value_inner(args, window_frame);
1000
1001    callers_temp_storage.reserve(results.len());
1002    results
1003        .into_iter()
1004        .zip_eq(orig_rows)
1005        .map(|(result_value, original_row)| {
1006            callers_temp_storage.make_datum(|packer| {
1007                packer.push_list_with(|packer| {
1008                    packer.push(result_value);
1009                    packer.push(original_row);
1010                });
1011            })
1012        })
1013}
1014
1015fn first_value_inner<'a>(datums: Vec<Datum<'a>>, window_frame: &WindowFrame) -> Vec<Datum<'a>> {
1016    let length = datums.len();
1017    let mut result: Vec<Datum> = Vec::with_capacity(length);
1018    for (idx, current_datum) in datums.iter().enumerate() {
1019        let first_value = match &window_frame.start_bound {
1020            // Always return the current value
1021            WindowFrameBound::CurrentRow => *current_datum,
1022            WindowFrameBound::UnboundedPreceding => {
1023                if let WindowFrameBound::OffsetPreceding(end_offset) = &window_frame.end_bound {
1024                    let end_offset = usize::cast_from(*end_offset);
1025
1026                    // If the frame ends before the first row, return null
1027                    if idx < end_offset {
1028                        Datum::Null
1029                    } else {
1030                        datums[0]
1031                    }
1032                } else {
1033                    datums[0]
1034                }
1035            }
1036            WindowFrameBound::OffsetPreceding(offset) => {
1037                let start_offset = usize::cast_from(*offset);
1038                let start_idx = idx.saturating_sub(start_offset);
1039                if let WindowFrameBound::OffsetPreceding(end_offset) = &window_frame.end_bound {
1040                    let end_offset = usize::cast_from(*end_offset);
1041
1042                    // If the frame is empty or ends before the first row, return null
1043                    if start_offset < end_offset || idx < end_offset {
1044                        Datum::Null
1045                    } else {
1046                        datums[start_idx]
1047                    }
1048                } else {
1049                    datums[start_idx]
1050                }
1051            }
1052            WindowFrameBound::OffsetFollowing(offset) => {
1053                let start_offset = usize::cast_from(*offset);
1054                let start_idx = idx.saturating_add(start_offset);
1055                if let WindowFrameBound::OffsetFollowing(end_offset) = &window_frame.end_bound {
1056                    // If the frame is empty or starts after the last row, return null
1057                    if offset > end_offset || start_idx >= length {
1058                        Datum::Null
1059                    } else {
1060                        datums[start_idx]
1061                    }
1062                } else {
1063                    datums
1064                        .get(start_idx)
1065                        .map(|d| d.clone())
1066                        .unwrap_or(Datum::Null)
1067                }
1068            }
1069            // Forbidden during planning
1070            WindowFrameBound::UnboundedFollowing => unreachable!(),
1071        };
1072        result.push(first_value);
1073    }
1074    result
1075}
1076
1077/// The expected input is in the format of [((OriginalRow, InputValue), OrderByExprs...)]
1078fn last_value<'a, I>(
1079    datums: I,
1080    callers_temp_storage: &'a RowArena,
1081    order_by: &[ColumnOrder],
1082    window_frame: &WindowFrame,
1083) -> Datum<'a>
1084where
1085    I: IntoIterator<Item = Datum<'a>>,
1086{
1087    let temp_storage = RowArena::new();
1088    let iter = last_value_no_list(datums, &temp_storage, order_by, window_frame);
1089    callers_temp_storage.make_datum(|packer| {
1090        packer.push_list(iter);
1091    })
1092}
1093
1094/// Like `last_value`, but doesn't perform the final wrapping in a list, returning an Iterator
1095/// instead.
1096fn last_value_no_list<'a: 'b, 'b, I>(
1097    datums: I,
1098    callers_temp_storage: &'b RowArena,
1099    order_by: &[ColumnOrder],
1100    window_frame: &WindowFrame,
1101) -> impl Iterator<Item = Datum<'b>>
1102where
1103    I: IntoIterator<Item = Datum<'a>>,
1104{
1105    // Sort the datums according to the ORDER BY expressions and return the ((OriginalRow, InputValue), OrderByRow) record
1106    // The OrderByRow is kept around because it is required to compute the peer groups in RANGE mode
1107    let datums = order_aggregate_datums_with_rank(datums, order_by);
1108
1109    // Decode the input (OriginalRow, InputValue) into separate datums, while keeping the OrderByRow
1110    let size_hint = datums.size_hint().0;
1111    let mut args = Vec::with_capacity(size_hint);
1112    let mut original_rows = Vec::with_capacity(size_hint);
1113    let mut order_by_rows = Vec::with_capacity(size_hint);
1114    for (d, order_by_row) in datums.into_iter() {
1115        let mut iter = d.unwrap_list().iter();
1116        let original_row = iter.next().unwrap();
1117        let arg = iter.next().unwrap();
1118        order_by_rows.push(order_by_row);
1119        original_rows.push(original_row);
1120        args.push(arg);
1121    }
1122
1123    let results = last_value_inner(args, &order_by_rows, window_frame);
1124
1125    callers_temp_storage.reserve(results.len());
1126    results
1127        .into_iter()
1128        .zip_eq(original_rows)
1129        .map(|(result_value, original_row)| {
1130            callers_temp_storage.make_datum(|packer| {
1131                packer.push_list_with(|packer| {
1132                    packer.push(result_value);
1133                    packer.push(original_row);
1134                });
1135            })
1136        })
1137}
1138
1139fn last_value_inner<'a>(
1140    args: Vec<Datum<'a>>,
1141    order_by_rows: &Vec<Row>,
1142    window_frame: &WindowFrame,
1143) -> Vec<Datum<'a>> {
1144    let length = args.len();
1145    let mut results: Vec<Datum> = Vec::with_capacity(length);
1146    for (idx, (current_datum, order_by_row)) in args.iter().zip_eq(order_by_rows).enumerate() {
1147        let last_value = match &window_frame.end_bound {
1148            WindowFrameBound::CurrentRow => match &window_frame.units {
1149                // Always return the current value when in ROWS mode
1150                WindowFrameUnits::Rows => *current_datum,
1151                WindowFrameUnits::Range => {
1152                    // When in RANGE mode, return the last value of the peer group
1153                    // The peer group is the group of rows with the same ORDER BY value
1154                    // Note: Range is only supported for the default window frame (RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW),
1155                    // which is why it does not appear in the other branches
1156                    let target_idx = order_by_rows[idx..]
1157                        .iter()
1158                        .enumerate()
1159                        .take_while(|(_, row)| *row == order_by_row)
1160                        .last()
1161                        .unwrap()
1162                        .0
1163                        + idx;
1164                    args[target_idx]
1165                }
1166                // GROUPS is not supported, and forbidden during planning
1167                WindowFrameUnits::Groups => unreachable!(),
1168            },
1169            WindowFrameBound::UnboundedFollowing => {
1170                if let WindowFrameBound::OffsetFollowing(start_offset) = &window_frame.start_bound {
1171                    let start_offset = usize::cast_from(*start_offset);
1172
1173                    // If the frame starts after the last row of the window, return null
1174                    if idx + start_offset > length - 1 {
1175                        Datum::Null
1176                    } else {
1177                        args[length - 1]
1178                    }
1179                } else {
1180                    args[length - 1]
1181                }
1182            }
1183            WindowFrameBound::OffsetFollowing(offset) => {
1184                let end_offset = usize::cast_from(*offset);
1185                let end_idx = idx.saturating_add(end_offset);
1186                if let WindowFrameBound::OffsetFollowing(start_offset) = &window_frame.start_bound {
1187                    let start_offset = usize::cast_from(*start_offset);
1188                    let start_idx = idx.saturating_add(start_offset);
1189
1190                    // If the frame is empty or starts after the last row of the window, return null
1191                    if end_offset < start_offset || start_idx >= length {
1192                        Datum::Null
1193                    } else {
1194                        // Return the last valid element in the window
1195                        args.get(end_idx).unwrap_or(&args[length - 1]).clone()
1196                    }
1197                } else {
1198                    args.get(end_idx).unwrap_or(&args[length - 1]).clone()
1199                }
1200            }
1201            WindowFrameBound::OffsetPreceding(offset) => {
1202                let end_offset = usize::cast_from(*offset);
1203                let end_idx = idx.saturating_sub(end_offset);
1204                if idx < end_offset {
1205                    // If the frame ends before the first row, return null
1206                    Datum::Null
1207                } else if let WindowFrameBound::OffsetPreceding(start_offset) =
1208                    &window_frame.start_bound
1209                {
1210                    // If the frame is empty, return null
1211                    if offset > start_offset {
1212                        Datum::Null
1213                    } else {
1214                        args[end_idx]
1215                    }
1216                } else {
1217                    args[end_idx]
1218                }
1219            }
1220            // Forbidden during planning
1221            WindowFrameBound::UnboundedPreceding => unreachable!(),
1222        };
1223        results.push(last_value);
1224    }
1225    results
1226}
1227
1228/// Executes `FusedValueWindowFunc` on a reduction group.
1229/// The expected input is in the format of `[((OriginalRow, (Args1, Args2, ...)), OrderByExprs...)]`
1230/// where `Args1`, `Args2`, are the arguments of each of the fused functions. For functions that
1231/// have only a single argument (first_value/last_value), these are simple values. For functions
1232/// that have multiple arguments (lag/lead), these are also records.
1233fn fused_value_window_func<'a, I>(
1234    input_datums: I,
1235    callers_temp_storage: &'a RowArena,
1236    funcs: &Vec<AggregateFunc>,
1237    order_by: &Vec<ColumnOrder>,
1238) -> Datum<'a>
1239where
1240    I: IntoIterator<Item = Datum<'a>>,
1241{
1242    let temp_storage = RowArena::new();
1243    let iter = fused_value_window_func_no_list(input_datums, &temp_storage, funcs, order_by);
1244    callers_temp_storage.make_datum(|packer| {
1245        packer.push_list(iter);
1246    })
1247}
1248
1249/// Like `fused_value_window_func`, but doesn't perform the final wrapping in a list, returning an
1250/// Iterator instead.
1251fn fused_value_window_func_no_list<'a: 'b, 'b, I>(
1252    input_datums: I,
1253    callers_temp_storage: &'b RowArena,
1254    funcs: &Vec<AggregateFunc>,
1255    order_by: &Vec<ColumnOrder>,
1256) -> impl Iterator<Item = Datum<'b>>
1257where
1258    I: IntoIterator<Item = Datum<'a>>,
1259{
1260    let has_last_value = funcs
1261        .iter()
1262        .any(|f| matches!(f, AggregateFunc::LastValue { .. }));
1263
1264    let input_datums_with_ranks = order_aggregate_datums_with_rank(input_datums, order_by);
1265
1266    let size_hint = input_datums_with_ranks.size_hint().0;
1267    let mut encoded_argsss = vec![Vec::with_capacity(size_hint); funcs.len()];
1268    let mut original_rows = Vec::with_capacity(size_hint);
1269    let mut order_by_rows = Vec::with_capacity(size_hint);
1270    for (d, order_by_row) in input_datums_with_ranks {
1271        let mut iter = d.unwrap_list().iter();
1272        let original_row = iter.next().unwrap();
1273        original_rows.push(original_row);
1274        let mut argss_iter = iter.next().unwrap().unwrap_list().iter();
1275        for i in 0..funcs.len() {
1276            let encoded_args = argss_iter.next().unwrap();
1277            encoded_argsss[i].push(encoded_args);
1278        }
1279        if has_last_value {
1280            order_by_rows.push(order_by_row);
1281        }
1282    }
1283
1284    let mut results_per_row = vec![Vec::with_capacity(funcs.len()); original_rows.len()];
1285    for (func, encoded_argss) in funcs.iter().zip_eq(encoded_argsss) {
1286        let results = match func {
1287            AggregateFunc::LagLead {
1288                order_by: inner_order_by,
1289                lag_lead,
1290                ignore_nulls,
1291            } => {
1292                assert_eq!(order_by, inner_order_by);
1293                let unwrapped_argss = encoded_argss
1294                    .into_iter()
1295                    .map(|encoded_args| unwrap_lag_lead_encoded_args(encoded_args))
1296                    .collect();
1297                lag_lead_inner(unwrapped_argss, lag_lead, ignore_nulls)
1298            }
1299            AggregateFunc::FirstValue {
1300                order_by: inner_order_by,
1301                window_frame,
1302            } => {
1303                assert_eq!(order_by, inner_order_by);
1304                // (No unwrapping to do on the args here, because there is only 1 arg, so it's not
1305                // wrapped into a record.)
1306                first_value_inner(encoded_argss, window_frame)
1307            }
1308            AggregateFunc::LastValue {
1309                order_by: inner_order_by,
1310                window_frame,
1311            } => {
1312                assert_eq!(order_by, inner_order_by);
1313                // (No unwrapping to do on the args here, because there is only 1 arg, so it's not
1314                // wrapped into a record.)
1315                last_value_inner(encoded_argss, &order_by_rows, window_frame)
1316            }
1317            _ => panic!("unknown window function in FusedValueWindowFunc"),
1318        };
1319        for (results, result) in results_per_row.iter_mut().zip_eq(results) {
1320            results.push(result);
1321        }
1322    }
1323
1324    callers_temp_storage.reserve(2 * original_rows.len());
1325    results_per_row
1326        .into_iter()
1327        .enumerate()
1328        .map(move |(i, results)| {
1329            callers_temp_storage.make_datum(|packer| {
1330                packer.push_list_with(|packer| {
1331                    packer
1332                        .push(callers_temp_storage.make_datum(|packer| packer.push_list(results)));
1333                    packer.push(original_rows[i]);
1334                });
1335            })
1336        })
1337}
1338
1339/// `input_datums` is an entire window partition.
1340/// The expected input is in the format of `[((OriginalRow, InputValue), OrderByExprs...)]`
1341/// See also in the comment in `window_func_applied_to`.
1342///
1343/// `wrapped_aggregate`: e.g., for `sum(...) OVER (...)`, this is the `sum(...)`.
1344///
1345/// Note that this `order_by` doesn't have expressions, only `ColumnOrder`s. For an explanation,
1346/// see the comment on `WindowExprType`.
1347fn window_aggr<'a, I, A>(
1348    input_datums: I,
1349    callers_temp_storage: &'a RowArena,
1350    wrapped_aggregate: &AggregateFunc,
1351    order_by: &[ColumnOrder],
1352    window_frame: &WindowFrame,
1353) -> Datum<'a>
1354where
1355    I: IntoIterator<Item = Datum<'a>>,
1356    A: OneByOneAggr,
1357{
1358    let temp_storage = RowArena::new();
1359    let iter = window_aggr_no_list::<I, A>(
1360        input_datums,
1361        &temp_storage,
1362        wrapped_aggregate,
1363        order_by,
1364        window_frame,
1365    );
1366    callers_temp_storage.make_datum(|packer| {
1367        packer.push_list(iter);
1368    })
1369}
1370
1371/// Like `window_aggr`, but doesn't perform the final wrapping in a list, returning an Iterator
1372/// instead.
1373fn window_aggr_no_list<'a: 'b, 'b, I, A>(
1374    input_datums: I,
1375    callers_temp_storage: &'b RowArena,
1376    wrapped_aggregate: &AggregateFunc,
1377    order_by: &[ColumnOrder],
1378    window_frame: &WindowFrame,
1379) -> impl Iterator<Item = Datum<'b>>
1380where
1381    I: IntoIterator<Item = Datum<'a>>,
1382    A: OneByOneAggr,
1383{
1384    // Sort the datums according to the ORDER BY expressions and return the ((OriginalRow, InputValue), OrderByRow) record
1385    // The OrderByRow is kept around because it is required to compute the peer groups in RANGE mode
1386    let datums = order_aggregate_datums_with_rank(input_datums, order_by);
1387
1388    // Decode the input (OriginalRow, InputValue) into separate datums, while keeping the OrderByRow
1389    let size_hint = datums.size_hint().0;
1390    let mut args: Vec<Datum> = Vec::with_capacity(size_hint);
1391    let mut original_rows: Vec<Datum> = Vec::with_capacity(size_hint);
1392    let mut order_by_rows = Vec::with_capacity(size_hint);
1393    for (d, order_by_row) in datums.into_iter() {
1394        let mut iter = d.unwrap_list().iter();
1395        let original_row = iter.next().unwrap();
1396        let arg = iter.next().unwrap();
1397        order_by_rows.push(order_by_row);
1398        original_rows.push(original_row);
1399        args.push(arg);
1400    }
1401
1402    let results = window_aggr_inner::<A>(
1403        args,
1404        &order_by_rows,
1405        wrapped_aggregate,
1406        order_by,
1407        window_frame,
1408        callers_temp_storage,
1409    );
1410
1411    callers_temp_storage.reserve(results.len());
1412    results
1413        .into_iter()
1414        .zip_eq(original_rows)
1415        .map(|(result_value, original_row)| {
1416            callers_temp_storage.make_datum(|packer| {
1417                packer.push_list_with(|packer| {
1418                    packer.push(result_value);
1419                    packer.push(original_row);
1420                });
1421            })
1422        })
1423}
1424
1425fn window_aggr_inner<'a, A>(
1426    mut args: Vec<Datum<'a>>,
1427    order_by_rows: &Vec<Row>,
1428    wrapped_aggregate: &AggregateFunc,
1429    order_by: &[ColumnOrder],
1430    window_frame: &WindowFrame,
1431    temp_storage: &'a RowArena,
1432) -> Vec<Datum<'a>>
1433where
1434    A: OneByOneAggr,
1435{
1436    let length = args.len();
1437    let mut result: Vec<Datum> = Vec::with_capacity(length);
1438
1439    // In this degenerate case, all results would be `wrapped_aggregate.default()` (usually null).
1440    // However, this currently can't happen, because
1441    // - Groups frame mode is currently not supported;
1442    // - Range frame mode is currently supported only for the default frame, which includes the
1443    //   current row.
1444    soft_assert_or_log!(
1445        !((matches!(window_frame.units, WindowFrameUnits::Groups)
1446            || matches!(window_frame.units, WindowFrameUnits::Range))
1447            && !window_frame.includes_current_row()),
1448        "window frame without current row"
1449    );
1450
1451    if (matches!(
1452        window_frame.start_bound,
1453        WindowFrameBound::UnboundedPreceding
1454    ) && matches!(window_frame.end_bound, WindowFrameBound::UnboundedFollowing))
1455        || (order_by.is_empty()
1456            && (matches!(window_frame.units, WindowFrameUnits::Groups)
1457                || matches!(window_frame.units, WindowFrameUnits::Range))
1458            && window_frame.includes_current_row())
1459    {
1460        // Either
1461        //  - UNBOUNDED frame in both directions, or
1462        //  - There is no ORDER BY and the frame is such that the current peer group is included.
1463        //    (The current peer group will be the whole partition if there is no ORDER BY.)
1464        // We simply need to compute the aggregate once, on the entire partition, and each input
1465        // row will get this one aggregate value as result.
1466        let result_value =
1467            wrapped_aggregate.eval(args.into_iter().map(|d| (d, Diff::ONE)), temp_storage);
1468        // Every row will get the above aggregate as result.
1469        for _ in 0..length {
1470            result.push(result_value);
1471        }
1472    } else {
1473        fn rows_between_unbounded_preceding_and_current_row<'a, A>(
1474            args: Vec<Datum<'a>>,
1475            result: &mut Vec<Datum<'a>>,
1476            mut one_by_one_aggr: A,
1477            temp_storage: &'a RowArena,
1478        ) where
1479            A: OneByOneAggr,
1480        {
1481            for current_arg in args.into_iter() {
1482                one_by_one_aggr.give(&current_arg);
1483                let result_value = one_by_one_aggr.get_current_aggregate(temp_storage);
1484                result.push(result_value);
1485            }
1486        }
1487
1488        fn groups_between_unbounded_preceding_and_current_row<'a, A>(
1489            args: Vec<Datum<'a>>,
1490            order_by_rows: &Vec<Row>,
1491            result: &mut Vec<Datum<'a>>,
1492            mut one_by_one_aggr: A,
1493            temp_storage: &'a RowArena,
1494        ) where
1495            A: OneByOneAggr,
1496        {
1497            let mut peer_group_start = 0;
1498            while peer_group_start < args.len() {
1499                // Find the boundaries of the current peer group.
1500                // peer_group_start will point to the first element of the peer group,
1501                // peer_group_end will point to _just after_ the last element of the peer group.
1502                let mut peer_group_end = peer_group_start + 1;
1503                while peer_group_end < args.len()
1504                    && order_by_rows[peer_group_start] == order_by_rows[peer_group_end]
1505                {
1506                    // The peer group goes on while the OrderByRows not differ.
1507                    peer_group_end += 1;
1508                }
1509                // Let's compute the aggregate (which will be the same for all records in this
1510                // peer group).
1511                for current_arg in args[peer_group_start..peer_group_end].iter() {
1512                    one_by_one_aggr.give(current_arg);
1513                }
1514                let agg_for_peer_group = one_by_one_aggr.get_current_aggregate(temp_storage);
1515                // Put the above aggregate into each record in the peer group.
1516                for _ in args[peer_group_start..peer_group_end].iter() {
1517                    result.push(agg_for_peer_group);
1518                }
1519                // Point to the start of the next peer group.
1520                peer_group_start = peer_group_end;
1521            }
1522        }
1523
1524        fn rows_between_offset_and_offset<'a>(
1525            args: Vec<Datum<'a>>,
1526            result: &mut Vec<Datum<'a>>,
1527            wrapped_aggregate: &AggregateFunc,
1528            temp_storage: &'a RowArena,
1529            offset_start: i64,
1530            offset_end: i64,
1531        ) {
1532            let len = args
1533                .len()
1534                .to_i64()
1535                .expect("window partition's len should fit into i64");
1536            for i in 0..len {
1537                let i = i.to_i64().expect("window partition shouldn't be super big");
1538                // Trim the start of the frame to make it not reach over the start of the window
1539                // partition.
1540                let frame_start = max(i + offset_start, 0)
1541                    .to_usize()
1542                    .expect("The max made sure it's not negative");
1543                // Trim the end of the frame to make it not reach over the end of the window
1544                // partition.
1545                let frame_end = min(i + offset_end, len - 1).to_usize();
1546                match frame_end {
1547                    Some(frame_end) => {
1548                        if frame_start <= frame_end {
1549                            // Compute the aggregate on the frame.
1550                            // TODO:
1551                            // This implementation is quite slow if the frame is large: we do an
1552                            // inner loop over the entire frame, and compute the aggregate from
1553                            // scratch. We could do better:
1554                            //  - For invertible aggregations we could do a rolling aggregation.
1555                            //  - There are various tricks for min/max as well, making use of either
1556                            //    the fixed size of the window, or that we are not retracting
1557                            //    arbitrary elements but doing queue operations. E.g., see
1558                            //    http://codercareer.blogspot.com/2012/02/no-33-maximums-in-sliding-windows.html
1559                            let frame_values = args[frame_start..=frame_end]
1560                                .iter()
1561                                .map(|d| (*d, Diff::ONE));
1562                            let result_value = wrapped_aggregate.eval(frame_values, temp_storage);
1563                            result.push(result_value);
1564                        } else {
1565                            // frame_start > frame_end, so this is an empty frame.
1566                            let result_value = wrapped_aggregate.default();
1567                            result.push(result_value);
1568                        }
1569                    }
1570                    None => {
1571                        // frame_end would be negative, so this is an empty frame.
1572                        let result_value = wrapped_aggregate.default();
1573                        result.push(result_value);
1574                    }
1575                }
1576            }
1577        }
1578
1579        match (
1580            &window_frame.units,
1581            &window_frame.start_bound,
1582            &window_frame.end_bound,
1583        ) {
1584            // Cases where one edge of the frame is CurrentRow.
1585            // Note that these cases could be merged into the more general cases below where one
1586            // edge is some offset (with offset = 0), but the CurrentRow cases probably cover 95%
1587            // of user queries, so let's make this simple and fast.
1588            (Rows, UnboundedPreceding, CurrentRow) => {
1589                rows_between_unbounded_preceding_and_current_row::<A>(
1590                    args,
1591                    &mut result,
1592                    A::new(wrapped_aggregate, false),
1593                    temp_storage,
1594                );
1595            }
1596            (Rows, CurrentRow, UnboundedFollowing) => {
1597                // Same as above, but reverse.
1598                args.reverse();
1599                rows_between_unbounded_preceding_and_current_row::<A>(
1600                    args,
1601                    &mut result,
1602                    A::new(wrapped_aggregate, true),
1603                    temp_storage,
1604                );
1605                result.reverse();
1606            }
1607            (Range, UnboundedPreceding, CurrentRow) => {
1608                // Note that for the default frame, the RANGE frame mode is identical to the GROUPS
1609                // frame mode.
1610                groups_between_unbounded_preceding_and_current_row::<A>(
1611                    args,
1612                    order_by_rows,
1613                    &mut result,
1614                    A::new(wrapped_aggregate, false),
1615                    temp_storage,
1616                );
1617            }
1618            // The next several cases all call `rows_between_offset_and_offset`. Note that the
1619            // offset passed to `rows_between_offset_and_offset` should be negated when it's
1620            // PRECEDING.
1621            (Rows, OffsetPreceding(start_prec), OffsetPreceding(end_prec)) => {
1622                let start_prec = start_prec.to_i64().expect(
1623                    "window frame start OFFSET shouldn't be super big (the planning ensured this)",
1624                );
1625                let end_prec = end_prec.to_i64().expect(
1626                    "window frame end OFFSET shouldn't be super big (the planning ensured this)",
1627                );
1628                rows_between_offset_and_offset(
1629                    args,
1630                    &mut result,
1631                    wrapped_aggregate,
1632                    temp_storage,
1633                    -start_prec,
1634                    -end_prec,
1635                );
1636            }
1637            (Rows, OffsetPreceding(start_prec), OffsetFollowing(end_fol)) => {
1638                let start_prec = start_prec.to_i64().expect(
1639                    "window frame start OFFSET shouldn't be super big (the planning ensured this)",
1640                );
1641                let end_fol = end_fol.to_i64().expect(
1642                    "window frame end OFFSET shouldn't be super big (the planning ensured this)",
1643                );
1644                rows_between_offset_and_offset(
1645                    args,
1646                    &mut result,
1647                    wrapped_aggregate,
1648                    temp_storage,
1649                    -start_prec,
1650                    end_fol,
1651                );
1652            }
1653            (Rows, OffsetFollowing(start_fol), OffsetFollowing(end_fol)) => {
1654                let start_fol = start_fol.to_i64().expect(
1655                    "window frame start OFFSET shouldn't be super big (the planning ensured this)",
1656                );
1657                let end_fol = end_fol.to_i64().expect(
1658                    "window frame end OFFSET shouldn't be super big (the planning ensured this)",
1659                );
1660                rows_between_offset_and_offset(
1661                    args,
1662                    &mut result,
1663                    wrapped_aggregate,
1664                    temp_storage,
1665                    start_fol,
1666                    end_fol,
1667                );
1668            }
1669            (Rows, OffsetFollowing(_), OffsetPreceding(_)) => {
1670                unreachable!() // The planning ensured that this nonsensical case can't happen
1671            }
1672            (Rows, OffsetPreceding(start_prec), CurrentRow) => {
1673                let start_prec = start_prec.to_i64().expect(
1674                    "window frame start OFFSET shouldn't be super big (the planning ensured this)",
1675                );
1676                let end_fol = 0;
1677                rows_between_offset_and_offset(
1678                    args,
1679                    &mut result,
1680                    wrapped_aggregate,
1681                    temp_storage,
1682                    -start_prec,
1683                    end_fol,
1684                );
1685            }
1686            (Rows, CurrentRow, OffsetFollowing(end_fol)) => {
1687                let start_fol = 0;
1688                let end_fol = end_fol.to_i64().expect(
1689                    "window frame end OFFSET shouldn't be super big (the planning ensured this)",
1690                );
1691                rows_between_offset_and_offset(
1692                    args,
1693                    &mut result,
1694                    wrapped_aggregate,
1695                    temp_storage,
1696                    start_fol,
1697                    end_fol,
1698                );
1699            }
1700            (Rows, CurrentRow, CurrentRow) => {
1701                // We could have a more efficient implementation for this, but this is probably
1702                // super rare. (Might be more common with RANGE or GROUPS frame mode, though!)
1703                let start_fol = 0;
1704                let end_fol = 0;
1705                rows_between_offset_and_offset(
1706                    args,
1707                    &mut result,
1708                    wrapped_aggregate,
1709                    temp_storage,
1710                    start_fol,
1711                    end_fol,
1712                );
1713            }
1714            (Rows, CurrentRow, OffsetPreceding(_))
1715            | (Rows, UnboundedFollowing, _)
1716            | (Rows, _, UnboundedPreceding)
1717            | (Rows, OffsetFollowing(..), CurrentRow) => {
1718                unreachable!() // The planning ensured that these nonsensical cases can't happen
1719            }
1720            (Rows, UnboundedPreceding, UnboundedFollowing) => {
1721                // This is handled by the complicated if condition near the beginning of this
1722                // function.
1723                unreachable!()
1724            }
1725            (Rows, UnboundedPreceding, OffsetPreceding(_))
1726            | (Rows, UnboundedPreceding, OffsetFollowing(_))
1727            | (Rows, OffsetPreceding(..), UnboundedFollowing)
1728            | (Rows, OffsetFollowing(..), UnboundedFollowing) => {
1729                // Unsupported. Bail in the planner.
1730                // https://github.com/MaterializeInc/database-issues/issues/6720
1731                unreachable!()
1732            }
1733            (Range, _, _) => {
1734                // Unsupported.
1735                // The planner doesn't allow Range frame mode for now (except for the default
1736                // frame), see https://github.com/MaterializeInc/database-issues/issues/6585
1737                // Note that it would be easy to handle (Range, CurrentRow, UnboundedFollowing):
1738                // it would be similar to (Rows, CurrentRow, UnboundedFollowing), but would call
1739                // groups_between_unbounded_preceding_current_row.
1740                unreachable!()
1741            }
1742            (Groups, _, _) => {
1743                // Unsupported.
1744                // The planner doesn't allow Groups frame mode for now, see
1745                // https://github.com/MaterializeInc/database-issues/issues/6588
1746                unreachable!()
1747            }
1748        }
1749    }
1750
1751    result
1752}
1753
1754/// Computes a bundle of fused window aggregations.
1755/// The input is similar to `window_aggr`, but `InputValue` is not just a single value, but a record
1756/// where each component is the input to one of the aggregations.
1757fn fused_window_aggr<'a, I, A>(
1758    input_datums: I,
1759    callers_temp_storage: &'a RowArena,
1760    wrapped_aggregates: &Vec<AggregateFunc>,
1761    order_by: &Vec<ColumnOrder>,
1762    window_frame: &WindowFrame,
1763) -> Datum<'a>
1764where
1765    I: IntoIterator<Item = Datum<'a>>,
1766    A: OneByOneAggr,
1767{
1768    let temp_storage = RowArena::new();
1769    let iter = fused_window_aggr_no_list::<_, A>(
1770        input_datums,
1771        &temp_storage,
1772        wrapped_aggregates,
1773        order_by,
1774        window_frame,
1775    );
1776    callers_temp_storage.make_datum(|packer| {
1777        packer.push_list(iter);
1778    })
1779}
1780
1781/// Like `fused_window_aggr`, but doesn't perform the final wrapping in a list, returning an
1782/// Iterator instead.
1783fn fused_window_aggr_no_list<'a: 'b, 'b, I, A>(
1784    input_datums: I,
1785    callers_temp_storage: &'b RowArena,
1786    wrapped_aggregates: &Vec<AggregateFunc>,
1787    order_by: &Vec<ColumnOrder>,
1788    window_frame: &WindowFrame,
1789) -> impl Iterator<Item = Datum<'b>>
1790where
1791    I: IntoIterator<Item = Datum<'a>>,
1792    A: OneByOneAggr,
1793{
1794    // Sort the datums according to the ORDER BY expressions and return the ((OriginalRow, InputValue), OrderByRow) record
1795    // The OrderByRow is kept around because it is required to compute the peer groups in RANGE mode
1796    let datums = order_aggregate_datums_with_rank(input_datums, order_by);
1797
1798    let size_hint = datums.size_hint().0;
1799    let mut argss = vec![Vec::with_capacity(size_hint); wrapped_aggregates.len()];
1800    let mut original_rows = Vec::with_capacity(size_hint);
1801    let mut order_by_rows = Vec::with_capacity(size_hint);
1802    for (d, order_by_row) in datums {
1803        let mut iter = d.unwrap_list().iter();
1804        let original_row = iter.next().unwrap();
1805        original_rows.push(original_row);
1806        let args_iter = iter.next().unwrap().unwrap_list().iter();
1807        // Push each argument into the respective list
1808        for (args, arg) in argss.iter_mut().zip_eq(args_iter) {
1809            args.push(arg);
1810        }
1811        order_by_rows.push(order_by_row);
1812    }
1813
1814    let mut results_per_row =
1815        vec![Vec::with_capacity(wrapped_aggregates.len()); original_rows.len()];
1816    for (wrapped_aggr, args) in wrapped_aggregates.iter().zip_eq(argss) {
1817        let results = window_aggr_inner::<A>(
1818            args,
1819            &order_by_rows,
1820            wrapped_aggr,
1821            order_by,
1822            window_frame,
1823            callers_temp_storage,
1824        );
1825        for (results, result) in results_per_row.iter_mut().zip_eq(results) {
1826            results.push(result);
1827        }
1828    }
1829
1830    callers_temp_storage.reserve(2 * original_rows.len());
1831    results_per_row
1832        .into_iter()
1833        .enumerate()
1834        .map(move |(i, results)| {
1835            callers_temp_storage.make_datum(|packer| {
1836                packer.push_list_with(|packer| {
1837                    packer
1838                        .push(callers_temp_storage.make_datum(|packer| packer.push_list(results)));
1839                    packer.push(original_rows[i]);
1840                });
1841            })
1842        })
1843}
1844
1845/// An implementation of an aggregation where we can send in the input elements one-by-one, and
1846/// can also ask the current aggregate at any moment. (This just delegates to other aggregation
1847/// evaluation approaches.)
1848pub trait OneByOneAggr {
1849    /// The `reverse` parameter makes the aggregations process input elements in reverse order.
1850    /// This has an effect only for non-commutative aggregations, e.g. `list_agg`. These are
1851    /// currently only some of the Basic aggregations. (Basic aggregations are handled by
1852    /// `NaiveOneByOneAggr`).
1853    fn new(agg: &AggregateFunc, reverse: bool) -> Self;
1854    /// Pushes one input element into the aggregation.
1855    fn give(&mut self, d: &Datum);
1856    /// Returns the value of the aggregate computed on the given values so far.
1857    fn get_current_aggregate<'a>(&self, temp_storage: &'a RowArena) -> Datum<'a>;
1858}
1859
1860/// Naive implementation of [OneByOneAggr], suitable for stuff like const folding, but too slow for
1861/// rendering. This relies only on infrastructure available in `mz-expr`. It simply saves all the
1862/// given input, and calls the given [AggregateFunc]'s `eval` method when asked about the current
1863/// aggregate. (For Accumulable and Hierarchical aggregations, the rendering has more efficient
1864/// implementations, but for Basic aggregations even the rendering uses this naive implementation.)
1865#[derive(Debug)]
1866pub struct NaiveOneByOneAggr {
1867    agg: AggregateFunc,
1868    input: Vec<Row>,
1869    reverse: bool,
1870}
1871
1872impl OneByOneAggr for NaiveOneByOneAggr {
1873    fn new(agg: &AggregateFunc, reverse: bool) -> Self {
1874        NaiveOneByOneAggr {
1875            agg: agg.clone(),
1876            input: Vec::new(),
1877            reverse,
1878        }
1879    }
1880
1881    fn give(&mut self, d: &Datum) {
1882        let mut row = Row::default();
1883        row.packer().push(d);
1884        self.input.push(row);
1885    }
1886
1887    fn get_current_aggregate<'a>(&self, temp_storage: &'a RowArena) -> Datum<'a> {
1888        temp_storage.make_datum(|packer| {
1889            packer.push(if !self.reverse {
1890                self.agg.eval(
1891                    self.input.iter().map(|r| (r.unpack_first(), Diff::ONE)),
1892                    temp_storage,
1893                )
1894            } else {
1895                self.agg.eval(
1896                    self.input
1897                        .iter()
1898                        .rev()
1899                        .map(|r| (r.unpack_first(), Diff::ONE)),
1900                    temp_storage,
1901                )
1902            });
1903        })
1904    }
1905}
1906
1907/// Identify whether the given aggregate function is Lag or Lead, since they share
1908/// implementations.
1909#[derive(
1910    Clone,
1911    Debug,
1912    Eq,
1913    PartialEq,
1914    Ord,
1915    PartialOrd,
1916    Serialize,
1917    Deserialize,
1918    Hash
1919)]
1920pub enum LagLeadType {
1921    Lag,
1922    Lead,
1923}
1924
1925#[derive(
1926    Clone,
1927    Debug,
1928    Eq,
1929    PartialEq,
1930    Ord,
1931    PartialOrd,
1932    Serialize,
1933    Deserialize,
1934    Hash
1935)]
1936pub enum AggregateFunc {
1937    MaxNumeric,
1938    MaxInt16,
1939    MaxInt32,
1940    MaxInt64,
1941    MaxUInt16,
1942    MaxUInt32,
1943    MaxUInt64,
1944    MaxMzTimestamp,
1945    MaxFloat32,
1946    MaxFloat64,
1947    MaxBool,
1948    MaxString,
1949    MaxDate,
1950    MaxTimestamp,
1951    MaxTimestampTz,
1952    MaxInterval,
1953    MaxTime,
1954    MinNumeric,
1955    MinInt16,
1956    MinInt32,
1957    MinInt64,
1958    MinUInt16,
1959    MinUInt32,
1960    MinUInt64,
1961    MinMzTimestamp,
1962    MinFloat32,
1963    MinFloat64,
1964    MinBool,
1965    MinString,
1966    MinDate,
1967    MinTimestamp,
1968    MinTimestampTz,
1969    MinInterval,
1970    MinTime,
1971    SumInt16,
1972    SumInt32,
1973    SumInt64,
1974    SumUInt16,
1975    SumUInt32,
1976    SumUInt64,
1977    SumFloat32,
1978    SumFloat64,
1979    SumNumeric,
1980    SumInterval,
1981    Count,
1982    Any,
1983    All,
1984    /// Accumulates `Datum::List`s whose first element is a JSON-typed `Datum`s
1985    /// into a JSON list. The other elements are columns used by `order_by`.
1986    ///
1987    /// WARNING: Unlike the `jsonb_agg` function that is exposed by the SQL
1988    /// layer, this function filters out `Datum::Null`, for consistency with
1989    /// the other aggregate functions.
1990    JsonbAgg {
1991        order_by: Vec<ColumnOrder>,
1992    },
1993    /// Zips `Datum::List`s whose first element is a JSON-typed `Datum`s into a
1994    /// JSON map. The other elements are columns used by `order_by`.
1995    ///
1996    /// WARNING: Unlike the `jsonb_object_agg` function that is exposed by the SQL
1997    /// layer, this function filters out `Datum::Null`, for consistency with
1998    /// the other aggregate functions.
1999    JsonbObjectAgg {
2000        order_by: Vec<ColumnOrder>,
2001    },
2002    /// Zips a `Datum::List` whose first element is a `Datum::List` guaranteed
2003    /// to be non-empty and whose len % 2 == 0 into a `Datum::Map`. The other
2004    /// elements are columns used by `order_by`.
2005    MapAgg {
2006        order_by: Vec<ColumnOrder>,
2007        value_type: SqlScalarType,
2008    },
2009    /// Accumulates `Datum::Array`s of `SqlScalarType::Record` whose first element is a `Datum::Array`
2010    /// into a single `Datum::Array` (the remaining fields are used by `order_by`).
2011    ArrayConcat {
2012        order_by: Vec<ColumnOrder>,
2013    },
2014    /// Accumulates `Datum::List`s of `SqlScalarType::Record` whose first field is a `Datum::List`
2015    /// into a single `Datum::List` (the remaining fields are used by `order_by`).
2016    ListConcat {
2017        order_by: Vec<ColumnOrder>,
2018    },
2019    StringAgg {
2020        order_by: Vec<ColumnOrder>,
2021    },
2022    RowNumber {
2023        order_by: Vec<ColumnOrder>,
2024    },
2025    Rank {
2026        order_by: Vec<ColumnOrder>,
2027    },
2028    DenseRank {
2029        order_by: Vec<ColumnOrder>,
2030    },
2031    LagLead {
2032        order_by: Vec<ColumnOrder>,
2033        lag_lead: LagLeadType,
2034        ignore_nulls: bool,
2035    },
2036    FirstValue {
2037        order_by: Vec<ColumnOrder>,
2038        window_frame: WindowFrame,
2039    },
2040    LastValue {
2041        order_by: Vec<ColumnOrder>,
2042        window_frame: WindowFrame,
2043    },
2044    /// Several value window functions fused into one function, to amortize overheads.
2045    FusedValueWindowFunc {
2046        funcs: Vec<AggregateFunc>,
2047        /// Currently, all the fused functions must have the same `order_by`. (We can later
2048        /// eliminate this limitation.)
2049        order_by: Vec<ColumnOrder>,
2050    },
2051    WindowAggregate {
2052        wrapped_aggregate: Box<AggregateFunc>,
2053        order_by: Vec<ColumnOrder>,
2054        window_frame: WindowFrame,
2055    },
2056    FusedWindowAggregate {
2057        wrapped_aggregates: Vec<AggregateFunc>,
2058        order_by: Vec<ColumnOrder>,
2059        window_frame: WindowFrame,
2060    },
2061    /// Accumulates any number of `Datum::Dummy`s into `Datum::Dummy`.
2062    ///
2063    /// Useful for removing an expensive aggregation while maintaining the shape
2064    /// of a reduce operator.
2065    Dummy,
2066}
2067
2068/// Expands an iterator of `(datum, diff)` into one `datum` per unit of `diff`.
2069///
2070/// A non-positive `diff` contributes no copies. This is used by aggregates that
2071/// are sensitive to multiplicity (e.g. `sum`), to recover a flat datum stream
2072/// from the count-aware surface.
2073fn expand_counts<'a, I>(datums: I) -> impl Iterator<Item = Datum<'a>>
2074where
2075    I: IntoIterator<Item = (Datum<'a>, Diff)>,
2076{
2077    datums.into_iter().flat_map(|(datum, diff)| {
2078        let copies = usize::try_from(diff.into_inner()).unwrap_or(0);
2079        std::iter::repeat(datum).take(copies)
2080    })
2081}
2082
2083impl AggregateFunc {
2084    /// Whether this aggregate's result is independent of the multiplicity of its
2085    /// inputs (e.g. `min`/`max`/`any`/`all`).
2086    ///
2087    /// Such aggregates can ignore the `diff` of each input, evaluating over the
2088    /// distinct datums rather than expanding by count. This keeps idempotent
2089    /// reductions linear in the number of distinct inputs.
2090    fn ignores_multiplicity(&self) -> bool {
2091        use AggregateFunc::*;
2092        matches!(
2093            self,
2094            MaxNumeric
2095                | MaxInt16
2096                | MaxInt32
2097                | MaxInt64
2098                | MaxUInt16
2099                | MaxUInt32
2100                | MaxUInt64
2101                | MaxMzTimestamp
2102                | MaxFloat32
2103                | MaxFloat64
2104                | MaxBool
2105                | MaxString
2106                | MaxDate
2107                | MaxTimestamp
2108                | MaxTimestampTz
2109                | MaxInterval
2110                | MaxTime
2111                | MinNumeric
2112                | MinInt16
2113                | MinInt32
2114                | MinInt64
2115                | MinUInt16
2116                | MinUInt32
2117                | MinUInt64
2118                | MinMzTimestamp
2119                | MinFloat32
2120                | MinFloat64
2121                | MinBool
2122                | MinString
2123                | MinDate
2124                | MinTimestamp
2125                | MinTimestampTz
2126                | MinInterval
2127                | MinTime
2128                | Any
2129                | All
2130        )
2131    }
2132
2133    /// Evaluates the aggregate over an iterator of `(datum, diff)` pairs.
2134    ///
2135    /// Each aggregate consumes the multiplicity (`diff`) in whatever way is most
2136    /// efficient: `count` sums the diffs, multiplicity-insensitive aggregates
2137    /// (see `AggregateFunc::ignores_multiplicity`) ignore them, and everything
2138    /// else expands each datum into `diff` copies (see `expand_counts`).
2139    pub fn eval<'a, I>(&self, datums: I, temp_storage: &'a RowArena) -> Datum<'a>
2140    where
2141        I: IntoIterator<Item = (Datum<'a>, Diff)>,
2142    {
2143        // Accumulable aggregates consume multiplicity directly rather than
2144        // expanding each `(datum, diff)` into `diff` copies. The cases handled
2145        // here mirror the dataflow's accumulable reduction (`build_accumulable`
2146        // in `mz_compute::render::reduce`) so that constant folding produces the
2147        // same result the dataflow would. Signed integer and interval sums are
2148        // folded here; unsigned sums are not, because their
2149        // negative-accumulation case is a query error in the dataflow that this
2150        // `Datum`-returning path cannot signal. Floats and numerics use bespoke
2151        // fixed-point/wide-decimal accumulators in the dataflow that
2152        // `expand_counts` does not reproduce.
2153        match self {
2154            AggregateFunc::Count => count(datums),
2155            AggregateFunc::SumInt16 | AggregateFunc::SumInt32 => {
2156                // `finalize_accum` narrows these to `i64` with wrapping.
2157                sum_signed_int_counted(datums, |accum| {
2158                    #[allow(clippy::as_conversions)]
2159                    let narrowed = accum as i64;
2160                    Datum::Int64(narrowed)
2161                })
2162            }
2163            AggregateFunc::SumInt64 => sum_signed_int_counted(datums, Datum::from),
2164            AggregateFunc::SumInterval => sum_interval_counted(datums),
2165            _ if self.ignores_multiplicity() => {
2166                self.eval_datums(datums.into_iter().map(|(datum, _diff)| datum), temp_storage)
2167            }
2168            _ => self.eval_datums(expand_counts(datums), temp_storage),
2169        }
2170    }
2171
2172    /// Evaluates the aggregate over a flat iterator of datums, ignoring multiplicity.
2173    fn eval_datums<'a, I>(&self, datums: I, temp_storage: &'a RowArena) -> Datum<'a>
2174    where
2175        I: IntoIterator<Item = Datum<'a>>,
2176    {
2177        match self {
2178            AggregateFunc::MaxNumeric => {
2179                max_datum::<'a, I, OrderedDecimal<numeric::Numeric>>(datums)
2180            }
2181            AggregateFunc::MaxInt16 => max_datum::<'a, I, i16>(datums),
2182            AggregateFunc::MaxInt32 => max_datum::<'a, I, i32>(datums),
2183            AggregateFunc::MaxInt64 => max_datum::<'a, I, i64>(datums),
2184            AggregateFunc::MaxUInt16 => max_datum::<'a, I, u16>(datums),
2185            AggregateFunc::MaxUInt32 => max_datum::<'a, I, u32>(datums),
2186            AggregateFunc::MaxUInt64 => max_datum::<'a, I, u64>(datums),
2187            AggregateFunc::MaxMzTimestamp => max_datum::<'a, I, mz_repr::Timestamp>(datums),
2188            AggregateFunc::MaxFloat32 => max_datum::<'a, I, OrderedFloat<f32>>(datums),
2189            AggregateFunc::MaxFloat64 => max_datum::<'a, I, OrderedFloat<f64>>(datums),
2190            AggregateFunc::MaxBool => max_datum::<'a, I, bool>(datums),
2191            AggregateFunc::MaxString => max_string(datums),
2192            AggregateFunc::MaxDate => max_datum::<'a, I, Date>(datums),
2193            AggregateFunc::MaxTimestamp => {
2194                max_datum::<'a, I, CheckedTimestamp<NaiveDateTime>>(datums)
2195            }
2196            AggregateFunc::MaxTimestampTz => {
2197                max_datum::<'a, I, CheckedTimestamp<DateTime<Utc>>>(datums)
2198            }
2199            AggregateFunc::MaxInterval => max_datum::<'a, I, Interval>(datums),
2200            AggregateFunc::MaxTime => max_datum::<'a, I, NaiveTime>(datums),
2201            AggregateFunc::MinNumeric => {
2202                min_datum::<'a, I, OrderedDecimal<numeric::Numeric>>(datums)
2203            }
2204            AggregateFunc::MinInt16 => min_datum::<'a, I, i16>(datums),
2205            AggregateFunc::MinInt32 => min_datum::<'a, I, i32>(datums),
2206            AggregateFunc::MinInt64 => min_datum::<'a, I, i64>(datums),
2207            AggregateFunc::MinUInt16 => min_datum::<'a, I, u16>(datums),
2208            AggregateFunc::MinUInt32 => min_datum::<'a, I, u32>(datums),
2209            AggregateFunc::MinUInt64 => min_datum::<'a, I, u64>(datums),
2210            AggregateFunc::MinMzTimestamp => min_datum::<'a, I, mz_repr::Timestamp>(datums),
2211            AggregateFunc::MinFloat32 => min_datum::<'a, I, OrderedFloat<f32>>(datums),
2212            AggregateFunc::MinFloat64 => min_datum::<'a, I, OrderedFloat<f64>>(datums),
2213            AggregateFunc::MinBool => min_datum::<'a, I, bool>(datums),
2214            AggregateFunc::MinString => min_string(datums),
2215            AggregateFunc::MinDate => min_datum::<'a, I, Date>(datums),
2216            AggregateFunc::MinTimestamp => {
2217                min_datum::<'a, I, CheckedTimestamp<NaiveDateTime>>(datums)
2218            }
2219            AggregateFunc::MinTimestampTz => {
2220                min_datum::<'a, I, CheckedTimestamp<DateTime<Utc>>>(datums)
2221            }
2222            AggregateFunc::MinInterval => min_datum::<'a, I, Interval>(datums),
2223            AggregateFunc::MinTime => min_datum::<'a, I, NaiveTime>(datums),
2224            AggregateFunc::SumInt16 => sum_datum::<'a, I, i16, i64>(datums),
2225            AggregateFunc::SumInt32 => sum_datum::<'a, I, i32, i64>(datums),
2226            AggregateFunc::SumInt64 => sum_datum::<'a, I, i64, i128>(datums),
2227            AggregateFunc::SumUInt16 => sum_datum::<'a, I, u16, u64>(datums),
2228            AggregateFunc::SumUInt32 => sum_datum::<'a, I, u32, u64>(datums),
2229            AggregateFunc::SumUInt64 => sum_datum::<'a, I, u64, u128>(datums),
2230            AggregateFunc::SumFloat32 => sum_datum::<'a, I, f32, f32>(datums),
2231            AggregateFunc::SumFloat64 => sum_datum::<'a, I, f64, f64>(datums),
2232            AggregateFunc::SumNumeric => sum_numeric(datums),
2233            AggregateFunc::SumInterval => sum_interval(datums),
2234            AggregateFunc::Count => unreachable!("Count is handled in `eval`"),
2235            AggregateFunc::Any => any(datums),
2236            AggregateFunc::All => all(datums),
2237            AggregateFunc::JsonbAgg { order_by } => jsonb_agg(datums, temp_storage, order_by),
2238            AggregateFunc::MapAgg { order_by, .. } | AggregateFunc::JsonbObjectAgg { order_by } => {
2239                dict_agg(datums, temp_storage, order_by)
2240            }
2241            AggregateFunc::ArrayConcat { order_by } => array_concat(datums, temp_storage, order_by),
2242            AggregateFunc::ListConcat { order_by } => list_concat(datums, temp_storage, order_by),
2243            AggregateFunc::StringAgg { order_by } => string_agg(datums, temp_storage, order_by),
2244            AggregateFunc::RowNumber { order_by } => row_number(datums, temp_storage, order_by),
2245            AggregateFunc::Rank { order_by } => rank(datums, temp_storage, order_by),
2246            AggregateFunc::DenseRank { order_by } => dense_rank(datums, temp_storage, order_by),
2247            AggregateFunc::LagLead {
2248                order_by,
2249                lag_lead: lag_lead_type,
2250                ignore_nulls,
2251            } => lag_lead(datums, temp_storage, order_by, lag_lead_type, ignore_nulls),
2252            AggregateFunc::FirstValue {
2253                order_by,
2254                window_frame,
2255            } => first_value(datums, temp_storage, order_by, window_frame),
2256            AggregateFunc::LastValue {
2257                order_by,
2258                window_frame,
2259            } => last_value(datums, temp_storage, order_by, window_frame),
2260            AggregateFunc::WindowAggregate {
2261                wrapped_aggregate,
2262                order_by,
2263                window_frame,
2264            } => window_aggr::<_, NaiveOneByOneAggr>(
2265                datums,
2266                temp_storage,
2267                wrapped_aggregate,
2268                order_by,
2269                window_frame,
2270            ),
2271            AggregateFunc::FusedValueWindowFunc { funcs, order_by } => {
2272                fused_value_window_func(datums, temp_storage, funcs, order_by)
2273            }
2274            AggregateFunc::FusedWindowAggregate {
2275                wrapped_aggregates,
2276                order_by,
2277                window_frame,
2278            } => fused_window_aggr::<_, NaiveOneByOneAggr>(
2279                datums,
2280                temp_storage,
2281                wrapped_aggregates,
2282                order_by,
2283                window_frame,
2284            ),
2285            AggregateFunc::Dummy => Datum::Dummy,
2286        }
2287    }
2288
2289    /// Like `eval`, but it's given a [OneByOneAggr]. If `self` is a `WindowAggregate`, then
2290    /// the given [OneByOneAggr] will be used to evaluate the wrapped aggregate inside the
2291    /// `WindowAggregate`. If `self` is not a `WindowAggregate`, then it simply calls `eval`.
2292    pub fn eval_with_fast_window_agg<'a, I, W>(
2293        &self,
2294        datums: I,
2295        temp_storage: &'a RowArena,
2296    ) -> Datum<'a>
2297    where
2298        I: IntoIterator<Item = (Datum<'a>, Diff)>,
2299        W: OneByOneAggr,
2300    {
2301        match self {
2302            AggregateFunc::WindowAggregate {
2303                wrapped_aggregate,
2304                order_by,
2305                window_frame,
2306            } => window_aggr::<_, W>(
2307                expand_counts(datums),
2308                temp_storage,
2309                wrapped_aggregate,
2310                order_by,
2311                window_frame,
2312            ),
2313            AggregateFunc::FusedWindowAggregate {
2314                wrapped_aggregates,
2315                order_by,
2316                window_frame,
2317            } => fused_window_aggr::<_, W>(
2318                expand_counts(datums),
2319                temp_storage,
2320                wrapped_aggregates,
2321                order_by,
2322                window_frame,
2323            ),
2324            _ => self.eval(datums, temp_storage),
2325        }
2326    }
2327
2328    pub fn eval_with_unnest_list<'a, I, W>(
2329        &self,
2330        datums: I,
2331        temp_storage: &'a RowArena,
2332    ) -> impl Iterator<Item = Datum<'a>>
2333    where
2334        I: IntoIterator<Item = (Datum<'a>, Diff)>,
2335        W: OneByOneAggr,
2336    {
2337        // TODO: Use `enum_dispatch` to construct a unified iterator instead of `collect_vec`.
2338        assert!(self.can_fuse_with_unnest_list());
2339        // Window functions are sensitive to multiplicity, so expand counts.
2340        let datums = expand_counts(datums);
2341        match self {
2342            AggregateFunc::RowNumber { order_by } => {
2343                row_number_no_list(datums, temp_storage, order_by).collect_vec()
2344            }
2345            AggregateFunc::Rank { order_by } => {
2346                rank_no_list(datums, temp_storage, order_by).collect_vec()
2347            }
2348            AggregateFunc::DenseRank { order_by } => {
2349                dense_rank_no_list(datums, temp_storage, order_by).collect_vec()
2350            }
2351            AggregateFunc::LagLead {
2352                order_by,
2353                lag_lead: lag_lead_type,
2354                ignore_nulls,
2355            } => lag_lead_no_list(datums, temp_storage, order_by, lag_lead_type, ignore_nulls)
2356                .collect_vec(),
2357            AggregateFunc::FirstValue {
2358                order_by,
2359                window_frame,
2360            } => first_value_no_list(datums, temp_storage, order_by, window_frame).collect_vec(),
2361            AggregateFunc::LastValue {
2362                order_by,
2363                window_frame,
2364            } => last_value_no_list(datums, temp_storage, order_by, window_frame).collect_vec(),
2365            AggregateFunc::FusedValueWindowFunc { funcs, order_by } => {
2366                fused_value_window_func_no_list(datums, temp_storage, funcs, order_by).collect_vec()
2367            }
2368            AggregateFunc::WindowAggregate {
2369                wrapped_aggregate,
2370                order_by,
2371                window_frame,
2372            } => window_aggr_no_list::<_, W>(
2373                datums,
2374                temp_storage,
2375                wrapped_aggregate,
2376                order_by,
2377                window_frame,
2378            )
2379            .collect_vec(),
2380            AggregateFunc::FusedWindowAggregate {
2381                wrapped_aggregates,
2382                order_by,
2383                window_frame,
2384            } => fused_window_aggr_no_list::<_, W>(
2385                datums,
2386                temp_storage,
2387                wrapped_aggregates,
2388                order_by,
2389                window_frame,
2390            )
2391            .collect_vec(),
2392            _ => unreachable!("asserted above that `can_fuse_with_unnest_list`"),
2393        }
2394        .into_iter()
2395    }
2396
2397    /// Returns the output of the aggregation function when applied on an empty
2398    /// input relation.
2399    pub fn default(&self) -> Datum<'static> {
2400        match self {
2401            AggregateFunc::Count => Datum::Int64(0),
2402            AggregateFunc::Any => Datum::False,
2403            AggregateFunc::All => Datum::True,
2404            AggregateFunc::Dummy => Datum::Dummy,
2405            _ => Datum::Null,
2406        }
2407    }
2408
2409    /// Returns a datum whose inclusion in the aggregation will not change its
2410    /// result.
2411    pub fn identity_datum(&self) -> Datum<'static> {
2412        match self {
2413            AggregateFunc::Any => Datum::False,
2414            AggregateFunc::All => Datum::True,
2415            AggregateFunc::Dummy => Datum::Dummy,
2416            AggregateFunc::ArrayConcat { .. } => Datum::empty_array(),
2417            AggregateFunc::ListConcat { .. } => Datum::empty_list(),
2418            AggregateFunc::RowNumber { .. }
2419            | AggregateFunc::Rank { .. }
2420            | AggregateFunc::DenseRank { .. }
2421            | AggregateFunc::LagLead { .. }
2422            | AggregateFunc::FirstValue { .. }
2423            | AggregateFunc::LastValue { .. }
2424            | AggregateFunc::WindowAggregate { .. }
2425            | AggregateFunc::FusedValueWindowFunc { .. }
2426            | AggregateFunc::FusedWindowAggregate { .. } => Datum::empty_list(),
2427            AggregateFunc::MaxNumeric
2428            | AggregateFunc::MaxInt16
2429            | AggregateFunc::MaxInt32
2430            | AggregateFunc::MaxInt64
2431            | AggregateFunc::MaxUInt16
2432            | AggregateFunc::MaxUInt32
2433            | AggregateFunc::MaxUInt64
2434            | AggregateFunc::MaxMzTimestamp
2435            | AggregateFunc::MaxFloat32
2436            | AggregateFunc::MaxFloat64
2437            | AggregateFunc::MaxBool
2438            | AggregateFunc::MaxString
2439            | AggregateFunc::MaxDate
2440            | AggregateFunc::MaxTimestamp
2441            | AggregateFunc::MaxTimestampTz
2442            | AggregateFunc::MaxInterval
2443            | AggregateFunc::MaxTime
2444            | AggregateFunc::MinNumeric
2445            | AggregateFunc::MinInt16
2446            | AggregateFunc::MinInt32
2447            | AggregateFunc::MinInt64
2448            | AggregateFunc::MinUInt16
2449            | AggregateFunc::MinUInt32
2450            | AggregateFunc::MinUInt64
2451            | AggregateFunc::MinMzTimestamp
2452            | AggregateFunc::MinFloat32
2453            | AggregateFunc::MinFloat64
2454            | AggregateFunc::MinBool
2455            | AggregateFunc::MinString
2456            | AggregateFunc::MinDate
2457            | AggregateFunc::MinTimestamp
2458            | AggregateFunc::MinTimestampTz
2459            | AggregateFunc::MinInterval
2460            | AggregateFunc::MinTime
2461            | AggregateFunc::SumInt16
2462            | AggregateFunc::SumInt32
2463            | AggregateFunc::SumInt64
2464            | AggregateFunc::SumUInt16
2465            | AggregateFunc::SumUInt32
2466            | AggregateFunc::SumUInt64
2467            | AggregateFunc::SumFloat32
2468            | AggregateFunc::SumFloat64
2469            | AggregateFunc::SumNumeric
2470            | AggregateFunc::SumInterval
2471            | AggregateFunc::Count
2472            | AggregateFunc::JsonbAgg { .. }
2473            | AggregateFunc::JsonbObjectAgg { .. }
2474            | AggregateFunc::MapAgg { .. }
2475            | AggregateFunc::StringAgg { .. } => Datum::Null,
2476        }
2477    }
2478
2479    pub fn can_fuse_with_unnest_list(&self) -> bool {
2480        match self {
2481            AggregateFunc::RowNumber { .. }
2482            | AggregateFunc::Rank { .. }
2483            | AggregateFunc::DenseRank { .. }
2484            | AggregateFunc::LagLead { .. }
2485            | AggregateFunc::FirstValue { .. }
2486            | AggregateFunc::LastValue { .. }
2487            | AggregateFunc::WindowAggregate { .. }
2488            | AggregateFunc::FusedValueWindowFunc { .. }
2489            | AggregateFunc::FusedWindowAggregate { .. } => true,
2490            AggregateFunc::ArrayConcat { .. }
2491            | AggregateFunc::ListConcat { .. }
2492            | AggregateFunc::Any
2493            | AggregateFunc::All
2494            | AggregateFunc::Dummy
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::SumInt16
2530            | AggregateFunc::SumInt32
2531            | AggregateFunc::SumInt64
2532            | AggregateFunc::SumUInt16
2533            | AggregateFunc::SumUInt32
2534            | AggregateFunc::SumUInt64
2535            | AggregateFunc::SumFloat32
2536            | AggregateFunc::SumFloat64
2537            | AggregateFunc::SumNumeric
2538            | AggregateFunc::SumInterval
2539            | AggregateFunc::Count
2540            | AggregateFunc::JsonbAgg { .. }
2541            | AggregateFunc::JsonbObjectAgg { .. }
2542            | AggregateFunc::MapAgg { .. }
2543            | AggregateFunc::StringAgg { .. } => false,
2544        }
2545    }
2546
2547    /// The output column type for the result of an aggregation.
2548    ///
2549    /// The output column type also contains nullability information, which
2550    /// is (without further information) true for aggregations that are not
2551    /// counts.
2552    pub fn output_sql_type(&self, input_type: SqlColumnType) -> SqlColumnType {
2553        let scalar_type = match self {
2554            AggregateFunc::Count => SqlScalarType::Int64,
2555            AggregateFunc::Any => SqlScalarType::Bool,
2556            AggregateFunc::All => SqlScalarType::Bool,
2557            AggregateFunc::JsonbAgg { .. } => SqlScalarType::Jsonb,
2558            AggregateFunc::JsonbObjectAgg { .. } => SqlScalarType::Jsonb,
2559            AggregateFunc::SumInt16 => SqlScalarType::Int64,
2560            AggregateFunc::SumInt32 => SqlScalarType::Int64,
2561            AggregateFunc::SumInt64 => SqlScalarType::Numeric {
2562                max_scale: Some(NumericMaxScale::ZERO),
2563            },
2564            AggregateFunc::SumUInt16 => SqlScalarType::UInt64,
2565            AggregateFunc::SumUInt32 => SqlScalarType::UInt64,
2566            AggregateFunc::SumUInt64 => SqlScalarType::Numeric {
2567                max_scale: Some(NumericMaxScale::ZERO),
2568            },
2569            AggregateFunc::MapAgg { value_type, .. } => SqlScalarType::Map {
2570                value_type: Box::new(value_type.clone()),
2571                custom_id: None,
2572            },
2573            AggregateFunc::ArrayConcat { .. } | AggregateFunc::ListConcat { .. } => {
2574                match input_type.scalar_type {
2575                    // The input is wrapped in a Record if there's an ORDER BY, so extract it out.
2576                    SqlScalarType::Record { ref fields, .. } => fields[0].1.scalar_type.clone(),
2577                    _ => unreachable!(),
2578                }
2579            }
2580            AggregateFunc::StringAgg { .. } => SqlScalarType::String,
2581            AggregateFunc::RowNumber { .. } => {
2582                AggregateFunc::output_type_ranking_window_funcs(&input_type, "?row_number?")
2583            }
2584            AggregateFunc::Rank { .. } => {
2585                AggregateFunc::output_type_ranking_window_funcs(&input_type, "?rank?")
2586            }
2587            AggregateFunc::DenseRank { .. } => {
2588                AggregateFunc::output_type_ranking_window_funcs(&input_type, "?dense_rank?")
2589            }
2590            AggregateFunc::LagLead { lag_lead: lag_lead_type, .. } => {
2591                // The input type for Lag is ((OriginalRow, EncodedArgs), OrderByExprs...)
2592                let fields = input_type.scalar_type.unwrap_record_element_type();
2593                let original_row_type = fields[0].unwrap_record_element_type()[0]
2594                    .clone()
2595                    .nullable(false);
2596                let encoded_args = fields[0].unwrap_record_element_type()[1];
2597                let output_type_inner =
2598                    Self::lag_lead_output_type_inner_from_encoded_args(encoded_args);
2599                let column_name = Self::lag_lead_result_column_name(lag_lead_type);
2600
2601                SqlScalarType::List {
2602                    element_type: Box::new(SqlScalarType::Record {
2603                        fields: [
2604                            (column_name, output_type_inner),
2605                            (ColumnName::from("?orig_row?"), original_row_type),
2606                        ].into(),
2607                        custom_id: None,
2608                    }),
2609                    custom_id: None,
2610                }
2611            }
2612            AggregateFunc::FirstValue { .. } => {
2613                // The input type for FirstValue is ((OriginalRow, Arg), OrderByExprs...)
2614                let fields = input_type.scalar_type.unwrap_record_element_type();
2615                let original_row_type = fields[0].unwrap_record_element_type()[0]
2616                    .clone()
2617                    .nullable(false);
2618                let value_type = fields[0].unwrap_record_element_type()[1]
2619                    .clone()
2620                    .nullable(true); // null when the partition is empty
2621
2622                SqlScalarType::List {
2623                    element_type: Box::new(SqlScalarType::Record {
2624                        fields: [
2625                            (ColumnName::from("?first_value?"), value_type),
2626                            (ColumnName::from("?orig_row?"), original_row_type),
2627                        ].into(),
2628                        custom_id: None,
2629                    }),
2630                    custom_id: None,
2631                }
2632            }
2633            AggregateFunc::LastValue { .. } => {
2634                // The input type for LastValue is ((OriginalRow, Arg), OrderByExprs...)
2635                let fields = input_type.scalar_type.unwrap_record_element_type();
2636                let original_row_type = fields[0].unwrap_record_element_type()[0]
2637                    .clone()
2638                    .nullable(false);
2639                let value_type = fields[0].unwrap_record_element_type()[1]
2640                    .clone()
2641                    .nullable(true); // null when the partition is empty
2642
2643                SqlScalarType::List {
2644                    element_type: Box::new(SqlScalarType::Record {
2645                        fields: [
2646                            (ColumnName::from("?last_value?"), value_type),
2647                            (ColumnName::from("?orig_row?"), original_row_type),
2648                        ].into(),
2649                        custom_id: None,
2650                    }),
2651                    custom_id: None,
2652                }
2653            }
2654            AggregateFunc::WindowAggregate {
2655                wrapped_aggregate, ..
2656            } => {
2657                // The input type for a window aggregate is ((OriginalRow, Arg), OrderByExprs...)
2658                let fields = input_type.scalar_type.unwrap_record_element_type();
2659                let original_row_type = fields[0].unwrap_record_element_type()[0]
2660                    .clone()
2661                    .nullable(false);
2662                let arg_type = fields[0].unwrap_record_element_type()[1]
2663                    .clone()
2664                    .nullable(true);
2665                let wrapped_aggr_out_type = wrapped_aggregate.output_sql_type(arg_type);
2666
2667                SqlScalarType::List {
2668                    element_type: Box::new(SqlScalarType::Record {
2669                        fields: [
2670                            (ColumnName::from("?window_agg?"), wrapped_aggr_out_type),
2671                            (ColumnName::from("?orig_row?"), original_row_type),
2672                        ].into(),
2673                        custom_id: None,
2674                    }),
2675                    custom_id: None,
2676                }
2677            }
2678            AggregateFunc::FusedWindowAggregate {
2679                wrapped_aggregates, ..
2680            } => {
2681                // The input type for a fused window aggregate is ((OriginalRow, Args), OrderByExprs...)
2682                // where `Args` is a record.
2683                let fields = input_type.scalar_type.unwrap_record_element_type();
2684                let original_row_type = fields[0].unwrap_record_element_type()[0]
2685                    .clone()
2686                    .nullable(false);
2687                let args_type = fields[0].unwrap_record_element_type()[1];
2688                let arg_types = args_type.unwrap_record_element_type();
2689                let out_fields = arg_types.iter().zip_eq(wrapped_aggregates).map(
2690                    |(arg_type, wrapped_agg)| {
2691                    (
2692                        ColumnName::from(wrapped_agg.name()),
2693                        wrapped_agg.output_sql_type((**arg_type).clone().nullable(true)),
2694                    )
2695                }).collect_vec();
2696
2697                SqlScalarType::List {
2698                    element_type: Box::new(SqlScalarType::Record {
2699                        fields: [
2700                            (ColumnName::from("?fused_window_agg?"), SqlScalarType::Record {
2701                                fields: out_fields.into(),
2702                                custom_id: None,
2703                            }.nullable(false)),
2704                            (ColumnName::from("?orig_row?"), original_row_type),
2705                        ].into(),
2706                        custom_id: None,
2707                    }),
2708                    custom_id: None,
2709                }
2710            }
2711            AggregateFunc::FusedValueWindowFunc { funcs, order_by: _ } => {
2712                // The input type is ((OriginalRow, EncodedArgs), OrderByExprs...)
2713                // where EncodedArgs is a record, where each element is the argument to one of the
2714                // function calls that got fused. This is a record for lag/lead, and a simple type
2715                // for first_value/last_value.
2716                let fields = input_type.scalar_type.unwrap_record_element_type();
2717                let original_row_type = fields[0].unwrap_record_element_type()[0]
2718                    .clone()
2719                    .nullable(false);
2720                let encoded_args_type = fields[0]
2721                    .unwrap_record_element_type()[1]
2722                    .unwrap_record_element_type();
2723
2724                SqlScalarType::List {
2725                    element_type: Box::new(SqlScalarType::Record {
2726                        fields: [
2727                            (
2728                                ColumnName::from("?fused_value_window_func?"),
2729                                SqlScalarType::Record {
2730                                fields: encoded_args_type.into_iter().zip_eq(funcs).map(
2731                                    |(arg_type, func)| {
2732                                    match func {
2733                                        AggregateFunc::LagLead {
2734                                            lag_lead: lag_lead_type, ..
2735                                        } => {
2736                                            let name = Self::lag_lead_result_column_name(
2737                                                lag_lead_type,
2738                                            );
2739                                            let ty = Self
2740                                                ::lag_lead_output_type_inner_from_encoded_args(
2741                                                    arg_type,
2742                                                );
2743                                            (name, ty)
2744                                        },
2745                                        AggregateFunc::FirstValue { .. } => {
2746                                            (
2747                                                ColumnName::from("?first_value?"),
2748                                                arg_type.clone().nullable(true),
2749                                            )
2750                                        }
2751                                        AggregateFunc::LastValue { .. } => {
2752                                            (
2753                                                ColumnName::from("?last_value?"),
2754                                                arg_type.clone().nullable(true),
2755                                            )
2756                                        }
2757                                        _ => panic!("FusedValueWindowFunc has an unknown function"),
2758                                    }
2759                                }).collect(),
2760                                custom_id: None,
2761                            }.nullable(false)),
2762                            (ColumnName::from("?orig_row?"), original_row_type),
2763                        ].into(),
2764                        custom_id: None,
2765                    }),
2766                    custom_id: None,
2767                }
2768            }
2769            AggregateFunc::Dummy
2770            | AggregateFunc::MaxNumeric
2771            | AggregateFunc::MaxInt16
2772            | AggregateFunc::MaxInt32
2773            | AggregateFunc::MaxInt64
2774            | AggregateFunc::MaxUInt16
2775            | AggregateFunc::MaxUInt32
2776            | AggregateFunc::MaxUInt64
2777            | AggregateFunc::MaxMzTimestamp
2778            | AggregateFunc::MaxFloat32
2779            | AggregateFunc::MaxFloat64
2780            | AggregateFunc::MaxBool
2781            // Note AggregateFunc::MaxString, MinString rely on returning input
2782            // type as output type to support the proper return type for
2783            // character input.
2784            | AggregateFunc::MaxString
2785            | AggregateFunc::MaxDate
2786            | AggregateFunc::MaxTimestamp
2787            | AggregateFunc::MaxTimestampTz
2788            | AggregateFunc::MaxInterval
2789            | AggregateFunc::MaxTime
2790            | AggregateFunc::MinNumeric
2791            | AggregateFunc::MinInt16
2792            | AggregateFunc::MinInt32
2793            | AggregateFunc::MinInt64
2794            | AggregateFunc::MinUInt16
2795            | AggregateFunc::MinUInt32
2796            | AggregateFunc::MinUInt64
2797            | AggregateFunc::MinMzTimestamp
2798            | AggregateFunc::MinFloat32
2799            | AggregateFunc::MinFloat64
2800            | AggregateFunc::MinBool
2801            | AggregateFunc::MinString
2802            | AggregateFunc::MinDate
2803            | AggregateFunc::MinTimestamp
2804            | AggregateFunc::MinTimestampTz
2805            | AggregateFunc::MinInterval
2806            | AggregateFunc::MinTime
2807            | AggregateFunc::SumFloat32
2808            | AggregateFunc::SumFloat64
2809            | AggregateFunc::SumNumeric
2810            | AggregateFunc::SumInterval => input_type.scalar_type.clone(),
2811        };
2812        // Count never produces null, and other aggregations only produce
2813        // null in the presence of null inputs.
2814        let nullable = match self {
2815            AggregateFunc::Count => false,
2816            // Use the nullability of the underlying column being aggregated, not the Records wrapping it
2817            AggregateFunc::StringAgg { .. } => match input_type.scalar_type {
2818                // The outer Record wraps the input in the first position, and any ORDER BY expressions afterwards
2819                SqlScalarType::Record { fields, .. } => match &fields[0].1.scalar_type {
2820                    // The inner Record is a (value, separator) tuple
2821                    SqlScalarType::Record { fields, .. } => fields[0].1.nullable,
2822                    _ => unreachable!(),
2823                },
2824                _ => unreachable!(),
2825            },
2826            _ => input_type.nullable,
2827        };
2828        scalar_type.nullable(nullable)
2829    }
2830
2831    /// Computes the representation type of this aggregate function.
2832    ///
2833    /// This is a wrapper around [`Self::output_sql_type`] that converts the result to a representation type.
2834    pub fn output_type(&self, input_type: ReprColumnType) -> ReprColumnType {
2835        ReprColumnType::from(&self.output_sql_type(SqlColumnType::from_repr(&input_type)))
2836    }
2837
2838    /// Compute output type for ROW_NUMBER, RANK, DENSE_RANK
2839    fn output_type_ranking_window_funcs(
2840        input_type: &SqlColumnType,
2841        col_name: &str,
2842    ) -> SqlScalarType {
2843        match input_type.scalar_type {
2844            SqlScalarType::Record { ref fields, .. } => SqlScalarType::List {
2845                element_type: Box::new(SqlScalarType::Record {
2846                    fields: [
2847                        (
2848                            ColumnName::from(col_name),
2849                            SqlScalarType::Int64.nullable(false),
2850                        ),
2851                        (ColumnName::from("?orig_row?"), {
2852                            let inner = match &fields[0].1.scalar_type {
2853                                SqlScalarType::List { element_type, .. } => element_type.clone(),
2854                                _ => unreachable!(),
2855                            };
2856                            inner.nullable(false)
2857                        }),
2858                    ]
2859                    .into(),
2860                    custom_id: None,
2861                }),
2862                custom_id: None,
2863            },
2864            _ => unreachable!(),
2865        }
2866    }
2867
2868    /// Given the `EncodedArgs` part of `((OriginalRow, EncodedArgs), OrderByExprs...)`,
2869    /// this computes the type of the first field of the output type. (The first field is the
2870    /// real result, the rest is the original row.)
2871    fn lag_lead_output_type_inner_from_encoded_args(
2872        encoded_args_type: &SqlScalarType,
2873    ) -> SqlColumnType {
2874        // lag/lead have 3 arguments, and the output type is
2875        // the same as the first of these, but always nullable. (It's null when the
2876        // lag/lead computation reaches over the bounds of the window partition.)
2877        encoded_args_type.unwrap_record_element_type()[0]
2878            .clone()
2879            .nullable(true)
2880    }
2881
2882    fn lag_lead_result_column_name(lag_lead_type: &LagLeadType) -> ColumnName {
2883        ColumnName::from(match lag_lead_type {
2884            LagLeadType::Lag => "?lag?",
2885            LagLeadType::Lead => "?lead?",
2886        })
2887    }
2888
2889    /// Returns true if the non-null constraint on the aggregation can be
2890    /// converted into a non-null constraint on its parameter expression, ie.
2891    /// whether the result of the aggregation is null if all the input values
2892    /// are null.
2893    pub fn propagates_nonnull_constraint(&self) -> bool {
2894        match self {
2895            AggregateFunc::MaxNumeric
2896            | AggregateFunc::MaxInt16
2897            | AggregateFunc::MaxInt32
2898            | AggregateFunc::MaxInt64
2899            | AggregateFunc::MaxUInt16
2900            | AggregateFunc::MaxUInt32
2901            | AggregateFunc::MaxUInt64
2902            | AggregateFunc::MaxMzTimestamp
2903            | AggregateFunc::MaxFloat32
2904            | AggregateFunc::MaxFloat64
2905            | AggregateFunc::MaxBool
2906            | AggregateFunc::MaxString
2907            | AggregateFunc::MaxDate
2908            | AggregateFunc::MaxTimestamp
2909            | AggregateFunc::MaxTimestampTz
2910            | AggregateFunc::MaxInterval
2911            | AggregateFunc::MaxTime
2912            | AggregateFunc::MinNumeric
2913            | AggregateFunc::MinInt16
2914            | AggregateFunc::MinInt32
2915            | AggregateFunc::MinInt64
2916            | AggregateFunc::MinUInt16
2917            | AggregateFunc::MinUInt32
2918            | AggregateFunc::MinUInt64
2919            | AggregateFunc::MinMzTimestamp
2920            | AggregateFunc::MinFloat32
2921            | AggregateFunc::MinFloat64
2922            | AggregateFunc::MinBool
2923            | AggregateFunc::MinString
2924            | AggregateFunc::MinDate
2925            | AggregateFunc::MinTimestamp
2926            | AggregateFunc::MinTimestampTz
2927            | AggregateFunc::MinInterval
2928            | AggregateFunc::MinTime
2929            | AggregateFunc::SumInt16
2930            | AggregateFunc::SumInt32
2931            | AggregateFunc::SumInt64
2932            | AggregateFunc::SumUInt16
2933            | AggregateFunc::SumUInt32
2934            | AggregateFunc::SumUInt64
2935            | AggregateFunc::SumFloat32
2936            | AggregateFunc::SumFloat64
2937            | AggregateFunc::SumNumeric
2938            | AggregateFunc::SumInterval
2939            | AggregateFunc::StringAgg { .. } => true,
2940            // Count is never null
2941            AggregateFunc::Count
2942            | AggregateFunc::Any
2943            | AggregateFunc::All
2944            | AggregateFunc::JsonbAgg { .. }
2945            | AggregateFunc::JsonbObjectAgg { .. }
2946            | AggregateFunc::MapAgg { .. }
2947            | AggregateFunc::ArrayConcat { .. }
2948            | AggregateFunc::ListConcat { .. }
2949            | AggregateFunc::RowNumber { .. }
2950            | AggregateFunc::Rank { .. }
2951            | AggregateFunc::DenseRank { .. }
2952            | AggregateFunc::LagLead { .. }
2953            | AggregateFunc::FirstValue { .. }
2954            | AggregateFunc::LastValue { .. }
2955            | AggregateFunc::FusedValueWindowFunc { .. }
2956            | AggregateFunc::WindowAggregate { .. }
2957            | AggregateFunc::FusedWindowAggregate { .. }
2958            | AggregateFunc::Dummy => false,
2959        }
2960    }
2961}
2962
2963fn jsonb_each<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
2964    // First produce a map, so that a common iterator can be returned.
2965    let map = match a {
2966        Datum::Map(dict) => dict,
2967        _ => mz_repr::DatumMap::empty(),
2968    };
2969
2970    map.iter()
2971        .map(move |(k, v)| (Row::pack_slice(&[Datum::String(k), v]), Diff::ONE))
2972}
2973
2974fn jsonb_each_stringify<'a>(
2975    a: Datum<'a>,
2976    temp_storage: &'a RowArena,
2977) -> impl Iterator<Item = (Row, Diff)> + 'a {
2978    // First produce a map, so that a common iterator can be returned.
2979    let map = match a {
2980        Datum::Map(dict) => dict,
2981        _ => mz_repr::DatumMap::empty(),
2982    };
2983
2984    map.iter().map(move |(k, mut v)| {
2985        v = jsonb_stringify(v, temp_storage)
2986            .map(Datum::String)
2987            .unwrap_or(Datum::Null);
2988        (Row::pack_slice(&[Datum::String(k), v]), Diff::ONE)
2989    })
2990}
2991
2992fn jsonb_object_keys<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
2993    let map = match a {
2994        Datum::Map(dict) => dict,
2995        _ => mz_repr::DatumMap::empty(),
2996    };
2997
2998    map.iter()
2999        .map(move |(k, _)| (Row::pack_slice(&[Datum::String(k)]), Diff::ONE))
3000}
3001
3002fn jsonb_array_elements<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
3003    let list = match a {
3004        Datum::List(list) => list,
3005        _ => mz_repr::DatumList::empty(),
3006    };
3007    list.iter().map(move |e| (Row::pack_slice(&[e]), Diff::ONE))
3008}
3009
3010fn jsonb_array_elements_stringify<'a>(
3011    a: Datum<'a>,
3012    temp_storage: &'a RowArena,
3013) -> impl Iterator<Item = (Row, Diff)> + 'a {
3014    let list = match a {
3015        Datum::List(list) => list,
3016        _ => mz_repr::DatumList::empty(),
3017    };
3018    list.iter().map(move |mut e| {
3019        e = jsonb_stringify(e, temp_storage)
3020            .map(Datum::String)
3021            .unwrap_or(Datum::Null);
3022        (Row::pack_slice(&[e]), Diff::ONE)
3023    })
3024}
3025
3026fn regexp_extract(a: Datum, r: &AnalyzedRegex) -> Option<(Row, Diff)> {
3027    let r = r.inner();
3028    let a = a.unwrap_str();
3029    let captures = r.captures(a)?;
3030    let datums = captures
3031        .iter()
3032        .skip(1)
3033        .map(|m| Datum::from(m.map(|m| m.as_str())));
3034    Some((Row::pack(datums), Diff::ONE))
3035}
3036
3037fn regexp_matches<'a>(
3038    exprs: &[Datum<'a>],
3039) -> Result<impl Iterator<Item = (Row, Diff)> + 'a, EvalError> {
3040    // There are only two acceptable ways to call this function:
3041    // 1. regexp_matches(string, regex)
3042    // 2. regexp_matches(string, regex, flag)
3043    assert!(exprs.len() == 2 || exprs.len() == 3);
3044    let a = exprs[0].unwrap_str();
3045    let r = exprs[1].unwrap_str();
3046
3047    let (regex, opts) = if exprs.len() == 3 {
3048        let flag = exprs[2].unwrap_str();
3049        let opts = AnalyzedRegexOpts::from_str(flag)?;
3050        (AnalyzedRegex::new(r, opts)?, opts)
3051    } else {
3052        let opts = AnalyzedRegexOpts::default();
3053        (AnalyzedRegex::new(r, opts)?, opts)
3054    };
3055
3056    let regex = regex.inner().clone();
3057
3058    let iter = regex.captures_iter(a).map(move |captures| {
3059        let matches = captures
3060            .iter()
3061            // The first match is the *entire* match, we want the capture groups by themselves.
3062            .skip(1)
3063            .map(|m| Datum::from(m.map(|m| m.as_str())))
3064            .collect::<Vec<_>>();
3065
3066        let mut binding = SharedRow::get();
3067        let mut packer = binding.packer();
3068
3069        let dimension = ArrayDimension {
3070            lower_bound: 1,
3071            length: matches.len(),
3072        };
3073        packer
3074            .try_push_array(&[dimension], matches)
3075            .expect("generated dimensions above");
3076
3077        (binding.clone(), Diff::ONE)
3078    });
3079
3080    // This is slightly unfortunate, but we need to collect the captures into a
3081    // Vec before we can yield them, because we can't return a iter with a
3082    // reference to the local `regex` variable.
3083    // We attempt to minimize the cost of this by using a SmallVec.
3084    let out = iter.collect::<SmallVec<[_; 3]>>();
3085
3086    if opts.global {
3087        Ok(Either::Left(out.into_iter()))
3088    } else {
3089        Ok(Either::Right(out.into_iter().take(1)))
3090    }
3091}
3092
3093fn generate_series<N>(
3094    start: N,
3095    stop: N,
3096    step: N,
3097) -> Result<impl Iterator<Item = (Row, Diff)>, EvalError>
3098where
3099    N: Integer + Signed + CheckedAdd + Clone,
3100    Datum<'static>: From<N>,
3101{
3102    if step == N::zero() {
3103        return Err(EvalError::InvalidParameterValue(
3104            "step size cannot equal zero".into(),
3105        ));
3106    }
3107    Ok(num::range_step_inclusive(start, stop, step)
3108        .map(move |i| (Row::pack_slice(&[Datum::from(i)]), Diff::ONE)))
3109}
3110
3111/// Like
3112/// [`num::range_step_inclusive`](https://github.com/rust-num/num-iter/blob/ddb14c1e796d401014c6c7a727de61d8109ad986/src/lib.rs#L279),
3113/// but for our timestamp types using [`Interval`] for `step`.xwxw
3114#[derive(Clone)]
3115pub struct TimestampRangeStepInclusive<T> {
3116    state: CheckedTimestamp<T>,
3117    stop: CheckedTimestamp<T>,
3118    step: Interval,
3119    rev: bool,
3120    done: bool,
3121}
3122
3123impl<T: TimestampLike> Iterator for TimestampRangeStepInclusive<T> {
3124    type Item = CheckedTimestamp<T>;
3125
3126    #[inline]
3127    fn next(&mut self) -> Option<CheckedTimestamp<T>> {
3128        if !self.done
3129            && ((self.rev && self.state >= self.stop) || (!self.rev && self.state <= self.stop))
3130        {
3131            let result = self.state.clone();
3132            match add_timestamp_months(self.state.deref(), self.step.months) {
3133                Ok(state) => match state.checked_add_signed(self.step.duration_as_chrono()) {
3134                    Some(v) => match CheckedTimestamp::from_timestamplike(v) {
3135                        Ok(v) => {
3136                            // Advance only if the step makes progress toward `stop`. A mixed
3137                            // month/day step can reach an in-bounds fixed point (month addition
3138                            // saturates a short month back onto the start day), which would
3139                            // otherwise loop forever.
3140                            let progressed = if self.rev {
3141                                v < self.state
3142                            } else {
3143                                v > self.state
3144                            };
3145                            if progressed {
3146                                self.state = v
3147                            } else {
3148                                self.done = true
3149                            }
3150                        }
3151                        Err(_) => self.done = true,
3152                    },
3153                    None => self.done = true,
3154                },
3155                Err(..) => {
3156                    self.done = true;
3157                }
3158            }
3159
3160            Some(result)
3161        } else {
3162            None
3163        }
3164    }
3165}
3166
3167fn generate_series_ts<T: TimestampLike>(
3168    start: CheckedTimestamp<T>,
3169    stop: CheckedTimestamp<T>,
3170    step: Interval,
3171    conv: fn(CheckedTimestamp<T>) -> Datum<'static>,
3172) -> Result<impl Iterator<Item = (Row, Diff)>, EvalError> {
3173    let normalized_step = step.as_microseconds();
3174    if normalized_step == 0 {
3175        return Err(EvalError::InvalidParameterValue(
3176            "step size cannot equal zero".into(),
3177        ));
3178    }
3179    let rev = normalized_step < 0;
3180
3181    let trsi = TimestampRangeStepInclusive {
3182        state: start,
3183        stop,
3184        step,
3185        rev,
3186        done: false,
3187    };
3188
3189    Ok(trsi.map(move |i| (Row::pack_slice(&[conv(i)]), Diff::ONE)))
3190}
3191
3192fn generate_subscripts_array(
3193    a: Datum,
3194    dim: i32,
3195) -> Result<Box<dyn Iterator<Item = (Row, Diff)>>, EvalError> {
3196    if dim <= 0 {
3197        return Ok(Box::new(iter::empty()));
3198    }
3199
3200    match a.unwrap_array().dims().into_iter().nth(
3201        (dim - 1)
3202            .try_into()
3203            .map_err(|_| EvalError::Int32OutOfRange((dim - 1).to_string().into()))?,
3204    ) {
3205        Some(requested_dim) => {
3206            let lower_bound: i32 = requested_dim.lower_bound.try_into().map_err(|_| {
3207                EvalError::Int32OutOfRange(requested_dim.lower_bound.to_string().into())
3208            })?;
3209            // The subscripts run from the lower bound to the upper bound,
3210            // inclusive. The upper bound is `lower_bound + length - 1`.
3211            let length: i32 = requested_dim
3212                .length
3213                .try_into()
3214                .map_err(|_| EvalError::Int32OutOfRange(requested_dim.length.to_string().into()))?;
3215            let upper_bound = lower_bound.checked_add(length - 1).ok_or_else(|| {
3216                EvalError::Int32OutOfRange(requested_dim.length.to_string().into())
3217            })?;
3218            Ok(Box::new(generate_series::<i32>(
3219                lower_bound,
3220                upper_bound,
3221                1,
3222            )?))
3223        }
3224        None => Ok(Box::new(iter::empty())),
3225    }
3226}
3227
3228fn unnest_array<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
3229    a.unwrap_array()
3230        .elements()
3231        .iter()
3232        .map(move |e| (Row::pack_slice(&[e]), Diff::ONE))
3233}
3234
3235fn unnest_list<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
3236    a.unwrap_list()
3237        .iter()
3238        .map(move |e| (Row::pack_slice(&[e]), Diff::ONE))
3239}
3240
3241fn unnest_map<'a>(a: Datum<'a>) -> impl Iterator<Item = (Row, Diff)> + 'a {
3242    a.unwrap_map()
3243        .iter()
3244        .map(move |(k, v)| (Row::pack_slice(&[Datum::from(k), v]), Diff::ONE))
3245}
3246
3247impl AggregateFunc {
3248    /// The base function name without the `~[...]` suffix used when rendering
3249    /// variants that represent a parameterized function family.
3250    pub fn name(&self) -> &'static str {
3251        match self {
3252            Self::MaxNumeric => "max",
3253            Self::MaxInt16 => "max",
3254            Self::MaxInt32 => "max",
3255            Self::MaxInt64 => "max",
3256            Self::MaxUInt16 => "max",
3257            Self::MaxUInt32 => "max",
3258            Self::MaxUInt64 => "max",
3259            Self::MaxMzTimestamp => "max",
3260            Self::MaxFloat32 => "max",
3261            Self::MaxFloat64 => "max",
3262            Self::MaxBool => "max",
3263            Self::MaxString => "max",
3264            Self::MaxDate => "max",
3265            Self::MaxTimestamp => "max",
3266            Self::MaxTimestampTz => "max",
3267            Self::MaxInterval => "max",
3268            Self::MaxTime => "max",
3269            Self::MinNumeric => "min",
3270            Self::MinInt16 => "min",
3271            Self::MinInt32 => "min",
3272            Self::MinInt64 => "min",
3273            Self::MinUInt16 => "min",
3274            Self::MinUInt32 => "min",
3275            Self::MinUInt64 => "min",
3276            Self::MinMzTimestamp => "min",
3277            Self::MinFloat32 => "min",
3278            Self::MinFloat64 => "min",
3279            Self::MinBool => "min",
3280            Self::MinString => "min",
3281            Self::MinDate => "min",
3282            Self::MinTimestamp => "min",
3283            Self::MinTimestampTz => "min",
3284            Self::MinInterval => "min",
3285            Self::MinTime => "min",
3286            Self::SumInt16 => "sum",
3287            Self::SumInt32 => "sum",
3288            Self::SumInt64 => "sum",
3289            Self::SumUInt16 => "sum",
3290            Self::SumUInt32 => "sum",
3291            Self::SumUInt64 => "sum",
3292            Self::SumFloat32 => "sum",
3293            Self::SumFloat64 => "sum",
3294            Self::SumNumeric => "sum",
3295            Self::SumInterval => "sum",
3296            Self::Count => "count",
3297            Self::Any => "any",
3298            Self::All => "all",
3299            Self::JsonbAgg { .. } => "jsonb_agg",
3300            Self::JsonbObjectAgg { .. } => "jsonb_object_agg",
3301            Self::MapAgg { .. } => "map_agg",
3302            Self::ArrayConcat { .. } => "array_agg",
3303            Self::ListConcat { .. } => "list_agg",
3304            Self::StringAgg { .. } => "string_agg",
3305            Self::RowNumber { .. } => "row_number",
3306            Self::Rank { .. } => "rank",
3307            Self::DenseRank { .. } => "dense_rank",
3308            Self::LagLead {
3309                lag_lead: LagLeadType::Lag,
3310                ..
3311            } => "lag",
3312            Self::LagLead {
3313                lag_lead: LagLeadType::Lead,
3314                ..
3315            } => "lead",
3316            Self::FirstValue { .. } => "first_value",
3317            Self::LastValue { .. } => "last_value",
3318            Self::WindowAggregate { .. } => "window_agg",
3319            Self::FusedValueWindowFunc { .. } => "fused_value_window_func",
3320            Self::FusedWindowAggregate { .. } => "fused_window_agg",
3321            Self::Dummy => "dummy",
3322        }
3323    }
3324}
3325
3326impl<'a, M> fmt::Display for HumanizedExpr<'a, AggregateFunc, M>
3327where
3328    M: HumanizerMode,
3329{
3330    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
3331        use AggregateFunc::*;
3332        let name = self.expr.name();
3333        match self.expr {
3334            JsonbAgg { order_by }
3335            | JsonbObjectAgg { order_by }
3336            | MapAgg { order_by, .. }
3337            | ArrayConcat { order_by }
3338            | ListConcat { order_by }
3339            | StringAgg { order_by }
3340            | RowNumber { order_by }
3341            | Rank { order_by }
3342            | DenseRank { order_by } => {
3343                let order_by = order_by.iter().map(|col| self.child(col));
3344                write!(f, "{}[order_by=[{}]]", name, separated(", ", order_by))
3345            }
3346            LagLead {
3347                lag_lead: _,
3348                ignore_nulls,
3349                order_by,
3350            } => {
3351                let order_by = order_by.iter().map(|col| self.child(col));
3352                f.write_str(name)?;
3353                f.write_str("[")?;
3354                if *ignore_nulls {
3355                    f.write_str("ignore_nulls=true, ")?;
3356                }
3357                write!(f, "order_by=[{}]", separated(", ", order_by))?;
3358                f.write_str("]")
3359            }
3360            FirstValue {
3361                order_by,
3362                window_frame,
3363            } => {
3364                let order_by = order_by.iter().map(|col| self.child(col));
3365                f.write_str(name)?;
3366                f.write_str("[")?;
3367                write!(f, "order_by=[{}]", separated(", ", order_by))?;
3368                if *window_frame != WindowFrame::default() {
3369                    write!(f, " {}", window_frame)?;
3370                }
3371                f.write_str("]")
3372            }
3373            LastValue {
3374                order_by,
3375                window_frame,
3376            } => {
3377                let order_by = order_by.iter().map(|col| self.child(col));
3378                f.write_str(name)?;
3379                f.write_str("[")?;
3380                write!(f, "order_by=[{}]", separated(", ", order_by))?;
3381                if *window_frame != WindowFrame::default() {
3382                    write!(f, " {}", window_frame)?;
3383                }
3384                f.write_str("]")
3385            }
3386            WindowAggregate {
3387                wrapped_aggregate,
3388                order_by,
3389                window_frame,
3390            } => {
3391                let order_by = order_by.iter().map(|col| self.child(col));
3392                let wrapped_aggregate = self.child(wrapped_aggregate.deref());
3393                f.write_str(name)?;
3394                f.write_str("[")?;
3395                write!(f, "{} ", wrapped_aggregate)?;
3396                write!(f, "order_by=[{}]", separated(", ", order_by))?;
3397                if *window_frame != WindowFrame::default() {
3398                    write!(f, " {}", window_frame)?;
3399                }
3400                f.write_str("]")
3401            }
3402            FusedValueWindowFunc { funcs, order_by } => {
3403                let order_by = order_by.iter().map(|col| self.child(col));
3404                let funcs = separated(", ", funcs.iter().map(|func| self.child(func)));
3405                f.write_str(name)?;
3406                f.write_str("[")?;
3407                write!(f, "{} ", funcs)?;
3408                write!(f, "order_by=[{}]", separated(", ", order_by))?;
3409                f.write_str("]")
3410            }
3411            _ => f.write_str(name),
3412        }
3413    }
3414}
3415
3416#[derive(
3417    Clone,
3418    Debug,
3419    Eq,
3420    PartialEq,
3421    Ord,
3422    PartialOrd,
3423    Serialize,
3424    Deserialize,
3425    Hash
3426)]
3427pub struct CaptureGroupDesc {
3428    pub index: u32,
3429    pub name: Option<String>,
3430    pub nullable: bool,
3431}
3432
3433#[derive(
3434    Clone,
3435    Copy,
3436    Debug,
3437    Eq,
3438    PartialEq,
3439    Ord,
3440    PartialOrd,
3441    Serialize,
3442    Deserialize,
3443    Hash,
3444    Default
3445)]
3446pub struct AnalyzedRegexOpts {
3447    pub case_insensitive: bool,
3448    pub global: bool,
3449}
3450
3451impl FromStr for AnalyzedRegexOpts {
3452    type Err = EvalError;
3453
3454    fn from_str(s: &str) -> Result<Self, Self::Err> {
3455        let mut opts = AnalyzedRegexOpts::default();
3456        for c in s.chars() {
3457            match c {
3458                'i' => opts.case_insensitive = true,
3459                'g' => opts.global = true,
3460                _ => return Err(EvalError::InvalidRegexFlag(c)),
3461            }
3462        }
3463        Ok(opts)
3464    }
3465}
3466
3467#[derive(
3468    Clone,
3469    Debug,
3470    Eq,
3471    PartialEq,
3472    Ord,
3473    PartialOrd,
3474    Serialize,
3475    Deserialize,
3476    Hash
3477)]
3478pub struct AnalyzedRegex(ReprRegex, Vec<CaptureGroupDesc>, AnalyzedRegexOpts);
3479
3480impl AnalyzedRegex {
3481    pub fn new(s: &str, opts: AnalyzedRegexOpts) -> Result<Self, RegexCompilationError> {
3482        let r = ReprRegex::new(s, opts.case_insensitive)?;
3483        // TODO(benesch): remove potentially dangerous usage of `as`.
3484        #[allow(clippy::as_conversions)]
3485        let descs: Vec<_> = r
3486            .capture_names()
3487            .enumerate()
3488            // The first capture is the entire matched string.
3489            // This will often not be useful, so skip it.
3490            // If people want it they can just surround their
3491            // entire regex in an explicit capture group.
3492            .skip(1)
3493            .map(|(i, name)| CaptureGroupDesc {
3494                index: i as u32,
3495                name: name.map(String::from),
3496                // TODO -- we can do better.
3497                // https://github.com/MaterializeInc/database-issues/issues/612
3498                nullable: true,
3499            })
3500            .collect();
3501        Ok(Self(r, descs, opts))
3502    }
3503    pub fn capture_groups_len(&self) -> usize {
3504        self.1.len()
3505    }
3506    pub fn capture_groups_iter(&self) -> impl Iterator<Item = &CaptureGroupDesc> {
3507        self.1.iter()
3508    }
3509    pub fn inner(&self) -> &Regex {
3510        &(self.0).regex
3511    }
3512    pub fn opts(&self) -> &AnalyzedRegexOpts {
3513        &self.2
3514    }
3515}
3516
3517pub fn csv_extract(a: Datum<'_>, n_cols: usize) -> impl Iterator<Item = (Row, Diff)> + '_ {
3518    let bytes = a.unwrap_str().as_bytes();
3519    let mut row = Row::default();
3520    let csv_reader = csv::ReaderBuilder::new()
3521        .has_headers(false)
3522        .from_reader(bytes);
3523    csv_reader.into_records().filter_map(move |res| match res {
3524        Ok(sr) if sr.len() == n_cols => {
3525            row.packer().extend(sr.iter().map(Datum::String));
3526            Some((row.clone(), Diff::ONE))
3527        }
3528        _ => None,
3529    })
3530}
3531
3532pub fn repeat_row(a: Datum) -> Option<(Row, Diff)> {
3533    let n = a.unwrap_int64();
3534    if n != 0 {
3535        Some((Row::default(), n.into()))
3536    } else {
3537        None
3538    }
3539}
3540
3541pub fn repeat_row_non_negative<'a>(
3542    a: Datum,
3543) -> Result<Box<dyn Iterator<Item = (Row, Diff)> + 'a>, EvalError> {
3544    let n = a.unwrap_int64();
3545    if n < 0 {
3546        Err(EvalError::InvalidParameterValue(
3547            format!("repeat_row_non_negative got {}", n).into(),
3548        ))
3549    } else if n == 0 {
3550        Ok(Box::new(iter::empty()))
3551    } else {
3552        // iterator with 1 element; n goes into the diff
3553        Ok(Box::new(iter::once((Row::default(), n.into()))))
3554    }
3555}
3556
3557fn wrap<'a>(datums: &'a [Datum<'a>], width: usize) -> impl Iterator<Item = (Row, Diff)> + 'a {
3558    datums
3559        .chunks(width)
3560        .map(|chunk| (Row::pack(chunk), Diff::ONE))
3561}
3562
3563fn acl_explode<'a>(
3564    acl_items: Datum<'a>,
3565    temp_storage: &'a RowArena,
3566) -> Result<impl Iterator<Item = (Row, Diff)> + 'a, EvalError> {
3567    let acl_items = acl_items.unwrap_array();
3568    let mut res = Vec::new();
3569    for acl_item in acl_items.elements().iter() {
3570        if acl_item.is_null() {
3571            return Err(EvalError::AclArrayNullElement);
3572        }
3573        let acl_item = acl_item.unwrap_acl_item();
3574        for privilege in acl_item.acl_mode.explode() {
3575            let row = [
3576                Datum::UInt32(acl_item.grantor.0),
3577                Datum::UInt32(acl_item.grantee.0),
3578                Datum::String(temp_storage.push_string(privilege.to_string())),
3579                // GRANT OPTION is not implemented, so we hardcode false.
3580                Datum::False,
3581            ];
3582            res.push((Row::pack_slice(&row), Diff::ONE));
3583        }
3584    }
3585    Ok(res.into_iter())
3586}
3587
3588fn mz_acl_explode<'a>(
3589    mz_acl_items: Datum<'a>,
3590    temp_storage: &'a RowArena,
3591) -> Result<impl Iterator<Item = (Row, Diff)> + 'a, EvalError> {
3592    let mz_acl_items = mz_acl_items.unwrap_array();
3593    let mut res = Vec::new();
3594    for mz_acl_item in mz_acl_items.elements().iter() {
3595        if mz_acl_item.is_null() {
3596            return Err(EvalError::MzAclArrayNullElement);
3597        }
3598        let mz_acl_item = mz_acl_item.unwrap_mz_acl_item();
3599        for privilege in mz_acl_item.acl_mode.explode() {
3600            let row = [
3601                Datum::String(temp_storage.push_string(mz_acl_item.grantor.to_string())),
3602                Datum::String(temp_storage.push_string(mz_acl_item.grantee.to_string())),
3603                Datum::String(temp_storage.push_string(privilege.to_string())),
3604                // GRANT OPTION is not implemented, so we hardcode false.
3605                Datum::False,
3606            ];
3607            res.push((Row::pack_slice(&row), Diff::ONE));
3608        }
3609    }
3610    Ok(res.into_iter())
3611}
3612
3613/// When adding a new `TableFunc` variant, please consider adding it to
3614/// `TableFunc::with_ordinality`!
3615#[derive(
3616    Clone,
3617    Debug,
3618    Eq,
3619    PartialEq,
3620    Ord,
3621    PartialOrd,
3622    Serialize,
3623    Deserialize,
3624    Hash
3625)]
3626pub enum TableFunc {
3627    AclExplode,
3628    MzAclExplode,
3629    JsonbEach,
3630    JsonbEachStringify,
3631    JsonbObjectKeys,
3632    JsonbArrayElements,
3633    JsonbArrayElementsStringify,
3634    RegexpExtract(AnalyzedRegex),
3635    CsvExtract(usize),
3636    GenerateSeriesInt32,
3637    GenerateSeriesInt64,
3638    /// An int64 `generate_series` that the optimizer promises to leave as an
3639    /// enumeration: no transform may match on this variant to replace its
3640    /// evaluation with a cardinality shortcut (compare the collapse of an
3641    /// unused `GenerateSeriesInt64` into `RepeatRowNonNegative`). Its
3642    /// *argument* expressions are still simplified like any other scalar.
3643    ///
3644    /// Exposed as `mz_unsafe.generate_series_unoptimized` for tests that rely
3645    /// on the work of enumeration actually happening (e.g. stress tests whose
3646    /// load would otherwise be optimized away). As with everything in
3647    /// `mz_unsafe`, it is not a supported surface: bug reports must not
3648    /// depend on it.
3649    GenerateSeriesUnoptimized,
3650    GenerateSeriesTimestamp,
3651    GenerateSeriesTimestampTz,
3652    /// Supplied with an input count,
3653    ///   1. Adds a column as if a typed subquery result,
3654    ///   2. Filters the row away if the count is only one,
3655    ///   3. Errors if the count is not exactly one.
3656    /// The intent is that this presents as if a subquery result with too many
3657    /// records contributing. The error column has the same type as the result
3658    /// should have, but we only produce it if the count exceeds one.
3659    ///
3660    /// This logic could nearly be achieved with map, filter, project logic,
3661    /// but has been challenging to do in a way that respects the vagaries of
3662    /// SQL and our semantics. If we reveal a constant value in the column we
3663    /// risk the optimizer pruning the branch; if we reveal that this will not
3664    /// produce rows we risk the optimizer pruning the branch; if we reveal that
3665    /// the only possible value is an error we risk the optimizer propagating that
3666    /// error without guards.
3667    ///
3668    /// Before replacing this by an `MirScalarExpr`, quadruple check that it
3669    /// would not result in misoptimizations due to expression evaluation order
3670    /// being utterly undefined, and predicate pushdown trimming any fragments
3671    /// that might produce columns that will not be needed.
3672    GuardSubquerySize {
3673        column_type: SqlScalarType,
3674    },
3675    /// Repeats the input row the given number of times. Can even repeat a negative number of times,
3676    /// which has some important consequences:
3677    /// - can lead to negative accumulations downstream;
3678    /// - can't be used in `WITH ORDINALITY` and other constructs that are implemented by
3679    ///   `TableFunc::WithOrdinality`, e.g., `ROWS FROM`;
3680    /// - output is non-monotonic.
3681    RepeatRow,
3682    /// Same as `RepeatRow`, but errors on a negative count, and thereby avoids the above
3683    /// peculiarities.
3684    RepeatRowNonNegative,
3685    UnnestArray {
3686        el_typ: SqlScalarType,
3687    },
3688    UnnestList {
3689        el_typ: SqlScalarType,
3690    },
3691    UnnestMap {
3692        value_type: SqlScalarType,
3693    },
3694    /// Given `n` input expressions, wraps them into `n / width` rows, each of
3695    /// `width` columns.
3696    ///
3697    /// This function is not intended to be called directly by end users, but
3698    /// is useful in the planning of e.g. VALUES clauses.
3699    Wrap {
3700        types: Vec<SqlColumnType>,
3701        width: usize,
3702    },
3703    GenerateSubscriptsArray,
3704    /// Execute some arbitrary scalar function as a table function.
3705    TabletizedScalar {
3706        name: String,
3707        relation: SqlRelationType,
3708    },
3709    RegexpMatches,
3710    /// Implements the WITH ORDINALITY clause.
3711    ///
3712    /// Don't construct `TableFunc::WithOrdinality` manually! Use the `with_ordinality` constructor
3713    /// function instead, which checks whether the given table function supports `WithOrdinality`.
3714    #[allow(private_interfaces)]
3715    WithOrdinality(WithOrdinality),
3716}
3717
3718/// Evaluates the inner table function, expands its results into unary (repeating each row as
3719/// many times as the diff indicates), and appends an integer corresponding to the ordinal
3720/// position (starting from 1). For example, it numbers the elements of a list when calling
3721/// `unnest_list`.
3722///
3723/// Private enum variant of `TableFunc`. Don't construct this directly, but use
3724/// `TableFunc::with_ordinality` instead.
3725#[derive(
3726    Clone,
3727    Debug,
3728    Eq,
3729    PartialEq,
3730    Ord,
3731    PartialOrd,
3732    Serialize,
3733    Deserialize,
3734    Hash
3735)]
3736struct WithOrdinality {
3737    inner: Box<TableFunc>,
3738}
3739
3740impl TableFunc {
3741    /// Adds `WITH ORDINALITY` to a table function if it's allowed on the given table function.
3742    pub fn with_ordinality(inner: TableFunc) -> Option<TableFunc> {
3743        match inner {
3744            TableFunc::AclExplode
3745            | TableFunc::MzAclExplode
3746            | TableFunc::JsonbEach
3747            | TableFunc::JsonbEachStringify
3748            | TableFunc::JsonbObjectKeys
3749            | TableFunc::JsonbArrayElements
3750            | TableFunc::JsonbArrayElementsStringify
3751            | TableFunc::RegexpExtract(_)
3752            | TableFunc::CsvExtract(_)
3753            | TableFunc::GenerateSeriesInt32
3754            | TableFunc::GenerateSeriesInt64
3755            | TableFunc::GenerateSeriesUnoptimized
3756            | TableFunc::GenerateSeriesTimestamp
3757            | TableFunc::GenerateSeriesTimestampTz
3758            | TableFunc::GuardSubquerySize { .. }
3759            | TableFunc::RepeatRowNonNegative
3760            | TableFunc::UnnestArray { .. }
3761            | TableFunc::UnnestList { .. }
3762            | TableFunc::UnnestMap { .. }
3763            | TableFunc::Wrap { .. }
3764            | TableFunc::GenerateSubscriptsArray
3765            | TableFunc::TabletizedScalar { .. }
3766            | TableFunc::RegexpMatches => Some(TableFunc::WithOrdinality(WithOrdinality {
3767                inner: Box::new(inner),
3768            })),
3769            // IMPORTANT: Before adding a new table function above, consider negative diffs:
3770            // `WithOrdinality::eval` will panic if the inner table function emits a negative diff.
3771            // (Note that negative diffs in the table function's _input_ don't matter. The table
3772            // function implementation doesn't see the input diffs, so the thing that matters here
3773            // is whether the table function itself can emit a negative diff.)
3774            TableFunc::RepeatRow // can produce negative diffs
3775            | TableFunc::WithOrdinality(_) => None, // no nesting of `WITH ORDINALITY` allowed
3776        }
3777    }
3778}
3779
3780impl TableFunc {
3781    /// Executes `self` on the given input row (`datums`).
3782    pub fn eval<'a>(
3783        &'a self,
3784        datums: &'a [Datum<'a>],
3785        temp_storage: &'a RowArena,
3786    ) -> Result<Box<dyn Iterator<Item = (Row, Diff)> + 'a>, EvalError> {
3787        if self.empty_on_null_input() && datums.iter().any(|d| d.is_null()) {
3788            return Ok(Box::new(vec![].into_iter()));
3789        }
3790        match self {
3791            TableFunc::AclExplode => Ok(Box::new(acl_explode(datums[0], temp_storage)?)),
3792            TableFunc::MzAclExplode => Ok(Box::new(mz_acl_explode(datums[0], temp_storage)?)),
3793            TableFunc::JsonbEach => Ok(Box::new(jsonb_each(datums[0]))),
3794            TableFunc::JsonbEachStringify => {
3795                Ok(Box::new(jsonb_each_stringify(datums[0], temp_storage)))
3796            }
3797            TableFunc::JsonbObjectKeys => Ok(Box::new(jsonb_object_keys(datums[0]))),
3798            TableFunc::JsonbArrayElements => Ok(Box::new(jsonb_array_elements(datums[0]))),
3799            TableFunc::JsonbArrayElementsStringify => Ok(Box::new(jsonb_array_elements_stringify(
3800                datums[0],
3801                temp_storage,
3802            ))),
3803            TableFunc::RegexpExtract(a) => Ok(Box::new(regexp_extract(datums[0], a).into_iter())),
3804            TableFunc::CsvExtract(n_cols) => Ok(Box::new(csv_extract(datums[0], *n_cols))),
3805            TableFunc::GenerateSeriesInt32 => {
3806                let res = generate_series(
3807                    datums[0].unwrap_int32(),
3808                    datums[1].unwrap_int32(),
3809                    datums[2].unwrap_int32(),
3810                )?;
3811                Ok(Box::new(res))
3812            }
3813            TableFunc::GenerateSeriesInt64 | TableFunc::GenerateSeriesUnoptimized => {
3814                let res = generate_series(
3815                    datums[0].unwrap_int64(),
3816                    datums[1].unwrap_int64(),
3817                    datums[2].unwrap_int64(),
3818                )?;
3819                Ok(Box::new(res))
3820            }
3821            TableFunc::GenerateSeriesTimestamp => {
3822                fn pass_through<'a>(d: CheckedTimestamp<NaiveDateTime>) -> Datum<'a> {
3823                    Datum::from(d)
3824                }
3825                let res = generate_series_ts(
3826                    datums[0].unwrap_timestamp(),
3827                    datums[1].unwrap_timestamp(),
3828                    datums[2].unwrap_interval(),
3829                    pass_through,
3830                )?;
3831                Ok(Box::new(res))
3832            }
3833            TableFunc::GenerateSeriesTimestampTz => {
3834                fn gen_ts_tz<'a>(d: CheckedTimestamp<DateTime<Utc>>) -> Datum<'a> {
3835                    Datum::from(d)
3836                }
3837                let res = generate_series_ts(
3838                    datums[0].unwrap_timestamptz(),
3839                    datums[1].unwrap_timestamptz(),
3840                    datums[2].unwrap_interval(),
3841                    gen_ts_tz,
3842                )?;
3843                Ok(Box::new(res))
3844            }
3845            TableFunc::GenerateSubscriptsArray => {
3846                generate_subscripts_array(datums[0], datums[1].unwrap_int32())
3847            }
3848            TableFunc::GuardSubquerySize { column_type: _ } => {
3849                // A subquery used as an expression may return at most one row.
3850                // For 0 or 1 we emit no rows and let the subquery's own output
3851                // flow through. Zero can't come directly from the count that
3852                // lowering plants (an MIR `count(true)`, at least 1 per group),
3853                // but over a provably empty subquery body the optimizer may
3854                // vacuously rewrite the counted expression to `null`, and a
3855                // count of `null` over a group is 0. Later simplifications can
3856                // surface that 0 as a literal argument that is evaluated
3857                // during optimization. Emitting no rows is also the correct
3858                // semantics: the empty subquery decorrelates to NULL via the
3859                // outer lookup.
3860                let count = datums[0].unwrap_int64();
3861                if count > 1 {
3862                    Err(EvalError::MultipleRowsFromSubquery)
3863                } else if count < 0 {
3864                    // Would require negative multiplicities to reach the guard.
3865                    Err(EvalError::NegativeRowsFromSubquery)
3866                } else {
3867                    Ok(Box::new([].into_iter()))
3868                }
3869            }
3870            TableFunc::RepeatRow => Ok(Box::new(repeat_row(datums[0]).into_iter())),
3871            TableFunc::RepeatRowNonNegative => repeat_row_non_negative(datums[0]),
3872            TableFunc::UnnestArray { .. } => Ok(Box::new(unnest_array(datums[0]))),
3873            TableFunc::UnnestList { .. } => Ok(Box::new(unnest_list(datums[0]))),
3874            TableFunc::UnnestMap { .. } => Ok(Box::new(unnest_map(datums[0]))),
3875            TableFunc::Wrap { width, .. } => Ok(Box::new(wrap(datums, *width))),
3876            TableFunc::TabletizedScalar { .. } => {
3877                let r = Row::pack_slice(datums);
3878                Ok(Box::new(std::iter::once((r, Diff::ONE))))
3879            }
3880            TableFunc::RegexpMatches => Ok(Box::new(regexp_matches(datums)?)),
3881            TableFunc::WithOrdinality(func_with_ordinality) => {
3882                func_with_ordinality.eval(datums, temp_storage)
3883            }
3884        }
3885    }
3886
3887    pub fn output_sql_type(&self) -> SqlRelationType {
3888        let (column_types, keys) = match self {
3889            TableFunc::AclExplode => {
3890                let column_types = vec![
3891                    SqlScalarType::Oid.nullable(false),
3892                    SqlScalarType::Oid.nullable(false),
3893                    SqlScalarType::String.nullable(false),
3894                    SqlScalarType::Bool.nullable(false),
3895                ];
3896                let keys = vec![];
3897                (column_types, keys)
3898            }
3899            TableFunc::MzAclExplode => {
3900                let column_types = vec![
3901                    SqlScalarType::String.nullable(false),
3902                    SqlScalarType::String.nullable(false),
3903                    SqlScalarType::String.nullable(false),
3904                    SqlScalarType::Bool.nullable(false),
3905                ];
3906                let keys = vec![];
3907                (column_types, keys)
3908            }
3909            TableFunc::JsonbEach => {
3910                let column_types = vec![
3911                    SqlScalarType::String.nullable(false),
3912                    SqlScalarType::Jsonb.nullable(false),
3913                ];
3914                let keys = vec![];
3915                (column_types, keys)
3916            }
3917            TableFunc::JsonbEachStringify => {
3918                let column_types = vec![
3919                    SqlScalarType::String.nullable(false),
3920                    SqlScalarType::String.nullable(true),
3921                ];
3922                let keys = vec![];
3923                (column_types, keys)
3924            }
3925            TableFunc::JsonbObjectKeys => {
3926                let column_types = vec![SqlScalarType::String.nullable(false)];
3927                let keys = vec![];
3928                (column_types, keys)
3929            }
3930            TableFunc::JsonbArrayElements => {
3931                let column_types = vec![SqlScalarType::Jsonb.nullable(false)];
3932                let keys = vec![];
3933                (column_types, keys)
3934            }
3935            TableFunc::JsonbArrayElementsStringify => {
3936                let column_types = vec![SqlScalarType::String.nullable(true)];
3937                let keys = vec![];
3938                (column_types, keys)
3939            }
3940            TableFunc::RegexpExtract(a) => {
3941                let column_types = a
3942                    .capture_groups_iter()
3943                    .map(|cg| SqlScalarType::String.nullable(cg.nullable))
3944                    .collect();
3945                let keys = vec![];
3946                (column_types, keys)
3947            }
3948            TableFunc::CsvExtract(n_cols) => {
3949                let column_types = iter::repeat(SqlScalarType::String.nullable(false))
3950                    .take(*n_cols)
3951                    .collect();
3952                let keys = vec![];
3953                (column_types, keys)
3954            }
3955            TableFunc::GenerateSeriesInt32 => {
3956                let column_types = vec![SqlScalarType::Int32.nullable(false)];
3957                let keys = vec![vec![0]];
3958                (column_types, keys)
3959            }
3960            TableFunc::GenerateSeriesInt64 | TableFunc::GenerateSeriesUnoptimized => {
3961                let column_types = vec![SqlScalarType::Int64.nullable(false)];
3962                let keys = vec![vec![0]];
3963                (column_types, keys)
3964            }
3965            TableFunc::GenerateSeriesTimestamp => {
3966                let column_types =
3967                    vec![SqlScalarType::Timestamp { precision: None }.nullable(false)];
3968                let keys = vec![vec![0]];
3969                (column_types, keys)
3970            }
3971            TableFunc::GenerateSeriesTimestampTz => {
3972                let column_types =
3973                    vec![SqlScalarType::TimestampTz { precision: None }.nullable(false)];
3974                let keys = vec![vec![0]];
3975                (column_types, keys)
3976            }
3977            TableFunc::GenerateSubscriptsArray => {
3978                let column_types = vec![SqlScalarType::Int32.nullable(false)];
3979                let keys = vec![vec![0]];
3980                (column_types, keys)
3981            }
3982            TableFunc::GuardSubquerySize { column_type } => {
3983                let column_types = vec![column_type.clone().nullable(false)];
3984                let keys = vec![];
3985                (column_types, keys)
3986            }
3987            TableFunc::RepeatRow | TableFunc::RepeatRowNonNegative => {
3988                let column_types = vec![];
3989                let keys = vec![];
3990                (column_types, keys)
3991            }
3992            TableFunc::UnnestArray { el_typ } => {
3993                let column_types = vec![el_typ.clone().nullable(true)];
3994                let keys = vec![];
3995                (column_types, keys)
3996            }
3997            TableFunc::UnnestList { el_typ } => {
3998                let column_types = vec![el_typ.clone().nullable(true)];
3999                let keys = vec![];
4000                (column_types, keys)
4001            }
4002            TableFunc::UnnestMap { value_type } => {
4003                let column_types = vec![
4004                    SqlScalarType::String.nullable(false),
4005                    value_type.clone().nullable(true),
4006                ];
4007                let keys = vec![vec![0]];
4008                (column_types, keys)
4009            }
4010            TableFunc::Wrap { types, .. } => {
4011                let column_types = types.clone();
4012                let keys = vec![];
4013                (column_types, keys)
4014            }
4015            TableFunc::TabletizedScalar { relation, .. } => {
4016                return relation.clone();
4017            }
4018            TableFunc::RegexpMatches => {
4019                let column_types =
4020                    vec![SqlScalarType::Array(Box::new(SqlScalarType::String)).nullable(false)];
4021                let keys = vec![];
4022
4023                (column_types, keys)
4024            }
4025            TableFunc::WithOrdinality(WithOrdinality { inner }) => {
4026                let mut typ = inner.output_sql_type();
4027                // Add the ordinality column.
4028                typ.column_types.push(SqlScalarType::Int64.nullable(false));
4029                // The ordinality column is always a key.
4030                typ.keys.push(vec![typ.column_types.len() - 1]);
4031                (typ.column_types, typ.keys)
4032            }
4033        };
4034
4035        soft_assert_eq_no_log!(column_types.len(), self.output_arity());
4036
4037        if !keys.is_empty() {
4038            SqlRelationType::new(column_types).with_keys(keys)
4039        } else {
4040            SqlRelationType::new(column_types)
4041        }
4042    }
4043
4044    /// Computes the representation type of this table function.
4045    ///
4046    /// This is a wrapper around [`Self::output_sql_type`] that converts the result to a representation type.
4047    pub fn output_type(&self) -> ReprRelationType {
4048        ReprRelationType::from(&self.output_sql_type())
4049    }
4050
4051    pub fn output_arity(&self) -> usize {
4052        match self {
4053            TableFunc::AclExplode => 4,
4054            TableFunc::MzAclExplode => 4,
4055            TableFunc::JsonbEach => 2,
4056            TableFunc::JsonbEachStringify => 2,
4057            TableFunc::JsonbObjectKeys => 1,
4058            TableFunc::JsonbArrayElements => 1,
4059            TableFunc::JsonbArrayElementsStringify => 1,
4060            TableFunc::RegexpExtract(a) => a.capture_groups_len(),
4061            TableFunc::CsvExtract(n_cols) => *n_cols,
4062            TableFunc::GenerateSeriesInt32 => 1,
4063            TableFunc::GenerateSeriesInt64 => 1,
4064            TableFunc::GenerateSeriesUnoptimized => 1,
4065            TableFunc::GenerateSeriesTimestamp => 1,
4066            TableFunc::GenerateSeriesTimestampTz => 1,
4067            TableFunc::GenerateSubscriptsArray => 1,
4068            TableFunc::GuardSubquerySize { .. } => 1,
4069            TableFunc::RepeatRow => 0,
4070            TableFunc::RepeatRowNonNegative => 0,
4071            TableFunc::UnnestArray { .. } => 1,
4072            TableFunc::UnnestList { .. } => 1,
4073            TableFunc::UnnestMap { .. } => 2,
4074            TableFunc::Wrap { width, .. } => *width,
4075            TableFunc::TabletizedScalar { relation, .. } => relation.column_types.len(),
4076            TableFunc::RegexpMatches => 1,
4077            TableFunc::WithOrdinality(WithOrdinality { inner }) => inner.output_arity() + 1,
4078        }
4079    }
4080
4081    pub fn empty_on_null_input(&self) -> bool {
4082        match self {
4083            TableFunc::AclExplode
4084            | TableFunc::MzAclExplode
4085            | TableFunc::JsonbEach
4086            | TableFunc::JsonbEachStringify
4087            | TableFunc::JsonbObjectKeys
4088            | TableFunc::JsonbArrayElements
4089            | TableFunc::JsonbArrayElementsStringify
4090            | TableFunc::GenerateSeriesInt32
4091            | TableFunc::GenerateSeriesInt64
4092            | TableFunc::GenerateSeriesUnoptimized
4093            | TableFunc::GenerateSeriesTimestamp
4094            | TableFunc::GenerateSeriesTimestampTz
4095            | TableFunc::GenerateSubscriptsArray
4096            | TableFunc::RegexpExtract(_)
4097            | TableFunc::CsvExtract(_)
4098            | TableFunc::RepeatRow
4099            | TableFunc::RepeatRowNonNegative
4100            | TableFunc::UnnestArray { .. }
4101            | TableFunc::UnnestList { .. }
4102            | TableFunc::UnnestMap { .. }
4103            | TableFunc::RegexpMatches => true,
4104            TableFunc::GuardSubquerySize { .. } => false,
4105            TableFunc::Wrap { .. } => false,
4106            TableFunc::TabletizedScalar { .. } => false,
4107            TableFunc::WithOrdinality(WithOrdinality { inner }) => inner.empty_on_null_input(),
4108        }
4109    }
4110
4111    /// True iff the table function preserves the append-only property of its input.
4112    pub fn preserves_monotonicity(&self) -> bool {
4113        // Most variants preserve monotonicity, but all variants are enumerated to
4114        // ensure that added variants at least check that this is the case.
4115        match self {
4116            TableFunc::AclExplode => false,
4117            TableFunc::MzAclExplode => false,
4118            TableFunc::JsonbEach => true,
4119            TableFunc::JsonbEachStringify => true,
4120            TableFunc::JsonbObjectKeys => true,
4121            TableFunc::JsonbArrayElements => true,
4122            TableFunc::JsonbArrayElementsStringify => true,
4123            TableFunc::RegexpExtract(_) => true,
4124            TableFunc::CsvExtract(_) => true,
4125            TableFunc::GenerateSeriesInt32 => true,
4126            TableFunc::GenerateSeriesInt64 => true,
4127            TableFunc::GenerateSeriesUnoptimized => true,
4128            TableFunc::GenerateSeriesTimestamp => true,
4129            TableFunc::GenerateSeriesTimestampTz => true,
4130            TableFunc::GenerateSubscriptsArray => true,
4131            TableFunc::RepeatRow => false,
4132            TableFunc::RepeatRowNonNegative => true,
4133            TableFunc::UnnestArray { .. } => true,
4134            TableFunc::UnnestList { .. } => true,
4135            TableFunc::UnnestMap { .. } => true,
4136            TableFunc::Wrap { .. } => true,
4137            TableFunc::TabletizedScalar { .. } => true,
4138            TableFunc::RegexpMatches => true,
4139            TableFunc::GuardSubquerySize { .. } => false,
4140            TableFunc::WithOrdinality(WithOrdinality { inner }) => inner.preserves_monotonicity(),
4141        }
4142    }
4143}
4144
4145impl fmt::Display for TableFunc {
4146    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
4147        match self {
4148            TableFunc::AclExplode => f.write_str("aclexplode"),
4149            TableFunc::MzAclExplode => f.write_str("mz_aclexplode"),
4150            TableFunc::JsonbEach => f.write_str("jsonb_each"),
4151            TableFunc::JsonbEachStringify => f.write_str("jsonb_each_text"),
4152            TableFunc::JsonbObjectKeys => f.write_str("jsonb_object_keys"),
4153            TableFunc::JsonbArrayElements => f.write_str("jsonb_array_elements"),
4154            TableFunc::JsonbArrayElementsStringify => f.write_str("jsonb_array_elements_text"),
4155            TableFunc::RegexpExtract(a) => write!(f, "regexp_extract({:?}, _)", a.0),
4156            TableFunc::CsvExtract(n_cols) => write!(f, "csv_extract({}, _)", n_cols),
4157            TableFunc::GenerateSeriesInt32 => f.write_str("generate_series"),
4158            TableFunc::GenerateSeriesInt64 => f.write_str("generate_series"),
4159            TableFunc::GenerateSeriesUnoptimized => f.write_str("generate_series_unoptimized"),
4160            TableFunc::GenerateSeriesTimestamp => f.write_str("generate_series"),
4161            TableFunc::GenerateSeriesTimestampTz => f.write_str("generate_series"),
4162            TableFunc::GenerateSubscriptsArray => f.write_str("generate_subscripts"),
4163            TableFunc::GuardSubquerySize { .. } => f.write_str("guard_subquery_size"),
4164            TableFunc::RepeatRow => f.write_str(REPEAT_ROW_NAME),
4165            TableFunc::RepeatRowNonNegative => f.write_str("repeat_row_non_negative"),
4166            TableFunc::UnnestArray { .. } => f.write_str("unnest_array"),
4167            TableFunc::UnnestList { .. } => f.write_str("unnest_list"),
4168            TableFunc::UnnestMap { .. } => f.write_str("unnest_map"),
4169            TableFunc::Wrap { width, .. } => write!(f, "wrap{}", width),
4170            TableFunc::TabletizedScalar { name, .. } => f.write_str(name),
4171            TableFunc::RegexpMatches => write!(f, "regexp_matches(_, _, _)"),
4172            TableFunc::WithOrdinality(WithOrdinality { inner }) => {
4173                write!(f, "{}[with_ordinality]", inner)
4174            }
4175        }
4176    }
4177}
4178
4179impl WithOrdinality {
4180    /// Executes the `self.inner` table function on the given input row (`datums`), and zips
4181    /// 1, 2, 3, ... to the result as a new column. We need to expand rows with non-1 diffs into the
4182    /// corresponding number of rows with unit diffs, because the ordinality column will have
4183    /// different values for each copy.
4184    ///
4185    /// # Panics
4186    ///
4187    /// Panics if the `inner` table function emits a negative diff.
4188    fn eval<'a>(
4189        &'a self,
4190        datums: &'a [Datum<'a>],
4191        temp_storage: &'a RowArena,
4192    ) -> Result<Box<dyn Iterator<Item = (Row, Diff)> + 'a>, EvalError> {
4193        let mut next_ordinal: i64 = 1;
4194        let it = self
4195            .inner
4196            .eval(datums, temp_storage)?
4197            .flat_map(move |(mut row, diff)| {
4198                let diff = diff.into_inner();
4199                // WITH ORDINALITY is not well-defined for negative diffs. This is ok, and
4200                // `TableFunc::with_ordinality` refuses to wrap such table functions in
4201                // `WithOrdinality` that can emit negative diffs, e.g., `repeat_row`.
4202                //
4203                // (Note that we don't need to worry about negative diffs in FlatMap's input,
4204                // because the diff of the input of the FlatMap is factored in after we return from
4205                // here.)
4206                assert!(diff >= 0);
4207                // The ordinals that will be associated with this row.
4208                let mut ordinals = next_ordinal..(next_ordinal + diff);
4209                next_ordinal += diff;
4210                // The maximum byte capacity we need for the original row and its ordinal.
4211                let cap = row.data_len() + datum_size(&Datum::Int64(next_ordinal));
4212                iter::from_fn(move || {
4213                    let ordinal = ordinals.next()?;
4214                    let mut row = if ordinals.is_empty() {
4215                        // This is the last row, so no need to clone. (Most table functions emit
4216                        // only 1 diffs, so this completely avoids cloning in most cases.)
4217                        std::mem::take(&mut row)
4218                    } else {
4219                        let mut new_row = Row::with_capacity(cap);
4220                        new_row.clone_from(&row);
4221                        new_row
4222                    };
4223                    RowPacker::for_existing_row(&mut row).push(Datum::Int64(ordinal));
4224                    Some((row, Diff::ONE))
4225                })
4226            });
4227        Ok(Box::new(it))
4228    }
4229}
4230
4231pub const REPEAT_ROW_NAME: &str = "repeat_row";
4232
4233#[cfg(test)]
4234mod tests {
4235    use mz_repr::{Datum, RowArena, SqlScalarType};
4236
4237    use super::TableFunc;
4238    use crate::EvalError;
4239
4240    /// 0 and 1 are valid (no guard rows), >1 errors with
4241    /// `MultipleRowsFromSubquery`, <0 with `NegativeRowsFromSubquery`. Zero is
4242    /// legitimate, not "can't happen": the optimizer can turn an empty
4243    /// subquery's count into a literal `0` that is evaluated during
4244    /// optimization (see the comment in `eval`), so it must not panic (exposed
4245    /// by #37049).
4246    #[mz_ore::test]
4247    fn guard_subquery_size_accepts_zero_and_one() {
4248        let func = TableFunc::GuardSubquerySize {
4249            column_type: SqlScalarType::Int64,
4250        };
4251        let temp_storage = RowArena::new();
4252
4253        for count in [0_i64, 1] {
4254            let rows = func
4255                .eval(&[Datum::Int64(count)], &temp_storage)
4256                .unwrap_or_else(|e| panic!("count {count} should be accepted, got {e:?}"))
4257                .count();
4258            assert_eq!(rows, 0, "count {count} should emit no guard rows");
4259        }
4260
4261        assert_eq!(
4262            func.eval(&[Datum::Int64(2)], &temp_storage).err(),
4263            Some(EvalError::MultipleRowsFromSubquery),
4264        );
4265        assert_eq!(
4266            func.eval(&[Datum::Int64(-1)], &temp_storage).err(),
4267            Some(EvalError::NegativeRowsFromSubquery),
4268        );
4269    }
4270}