Skip to main content

mz_expr/scalar/func/
registry.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//! Registry of the declared properties of every scalar function variant.
11//!
12//! The stable LIR schema (see `lir_schema.rs` in `mz-compute-types`) pins the
13//! *shape* of [`UnaryFunc`], [`BinaryFunc`] and [`VariadicFunc`]: their
14//! variants and payloads. It cannot see the properties the optimizer and the
15//! renderer rely on, such as null propagation, error behavior, monotonicity
16//! and output typing, nor the function bodies themselves. A change to any of
17//! those changes the meaning of a stored LIR plan without changing its
18//! serialized form.
19//!
20//! [`FuncRegistry::build`] records those properties for one representative
21//! instance of every variant, and `tests/func_registry.rs` in
22//! `mz-compute-types` compares the result against checked-in snapshots, one
23//! per half of [`Record`]. The module exists only under the `func-registry` feature,
24//! which that test enables. Production builds carry none of it. For `#[sqlfunc]`
25//! functions the record also carries the declaration text, the types-only
26//! signature and a fingerprint of the function body, see [`SqlFuncSource`],
27//! and the output type is probed at the column types the parameter types
28//! map to, see [`ColumnTypeProbe`].
29//!
30//! Variants whose payload cannot be constructed without data need a
31//! representative [`Sample`] in this module. Building the registry panics
32//! and names the variant otherwise. Samples may also be given for a
33//! constructible variant, to probe its output type at chosen input types or
34//! to record payloads whose properties differ, see [`Sample`].
35
36use std::collections::{BTreeMap, BTreeSet};
37use std::fmt;
38
39use mz_repr::adt::char::CharLength;
40use mz_repr::adt::datetime::DateTimeUnits;
41use mz_repr::adt::numeric::NumericMaxScale;
42use mz_repr::adt::regex::Regex;
43use mz_repr::adt::timestamp::TimestampPrecision;
44use mz_repr::adt::varchar::VarCharMaxLength;
45use mz_repr::{ColumnName, Datum, Row, SqlColumnType, SqlScalarType, StableRow};
46use serde::Serialize;
47
48use crate::func::format::DateTimeFormat;
49use crate::func::variadic::{
50    And, ArrayCreate, ArrayFill, ArrayIndex, ArrayToString, Coalesce, ErrorIfNull, Greatest, Least,
51    ListCreate, MapBuild, Or, RangeCreate, RecordCreate,
52};
53use crate::func::*;
54use crate::{BinaryFunc, MirScalarExpr, UnaryFunc, VariadicFunc, like_pattern};
55
56/// A representative instance of a function variant.
57///
58/// `input_types` are the column types the output type is probed at. Leave
59/// them empty to skip the probe, for functions whose output typing does not
60/// depend on a specific input shape.
61///
62/// A variant's properties can depend on its payload, for example a numeric
63/// cast errors only when it has a scale. One sample sees one point of each
64/// such property function, so a variant may have several samples, told apart
65/// by `label`. The unlabeled sample is the variant's primary record, keyed by
66/// its canonical name. Labeled samples are keyed `name[label]`.
67#[derive(Debug)]
68pub struct Sample<F> {
69    pub func: F,
70    pub input_types: Vec<SqlColumnType>,
71    pub label: &'static str,
72}
73
74impl<F> Sample<F> {
75    /// Marks this as an additional sample of its variant, recorded alongside
76    /// the primary one under `name[label]`.
77    pub fn labeled(self, label: &'static str) -> Self {
78        assert!(!label.is_empty(), "sample labels must be non-empty");
79        Sample { label, ..self }
80    }
81}
82
83/// The registry of every variant of the three scalar function enums, keyed
84/// by canonical variant name (see [`FuncName`]) as described on [`Sample`].
85#[derive(Debug)]
86pub struct FuncRegistry {
87    pub unary: BTreeMap<String, Record<UnaryFuncProperties>>,
88    pub binary: BTreeMap<String, Record<BinaryFuncProperties>>,
89    pub variadic: BTreeMap<String, Record<VariadicFuncProperties>>,
90}
91
92/// One registry entry: the properties that decide what a stored plan using
93/// the function computes, and the source they were read from.
94///
95/// The two halves are snapshotted separately by the test in
96/// `mz-compute-types`. A change to `properties` of a shipped LIR version
97/// requires a version bump. A change to `source` alone is informational.
98#[derive(Debug)]
99pub struct Record<P> {
100    pub properties: P,
101    pub source: FuncSource,
102}
103
104/// Where a record's properties came from.
105///
106/// `display` only feeds EXPLAIN output, since LIR stores variant names.
107/// `sqlfunc_decl` is the declaration text, whose semantic content the
108/// properties already carry as their own fields. `body_fingerprint` tracks
109/// the implementation, whose semantics only a reader can judge.
110#[derive(Debug, Serialize)]
111pub struct FuncSource {
112    pub display: String,
113    pub sqlfunc_decl: Option<&'static str>,
114    pub body_fingerprint: Option<String>,
115}
116
117impl FuncSource {
118    fn of(display: String, source: Option<SqlFuncSource>) -> Self {
119        FuncSource {
120            display,
121            sqlfunc_decl: source.map(|s| s.decl),
122            body_fingerprint: source.map(|s| format!("{:016x}", s.body_fingerprint)),
123        }
124    }
125}
126
127/// Resolves a Rust parameter type to its column type by autoref
128/// specialization: `(&ColumnTypeProbe::<T>(PhantomData)).column_type()` picks
129/// [`ProbeColumnType`] when `T: AsColumnType` and [`ProbeColumnTypeFallback`]
130/// otherwise. `#[sqlfunc]` emits that expression for each parameter, which is
131/// how it can ask for a column type without knowing whether one exists.
132#[derive(Debug)]
133pub struct ColumnTypeProbe<T>(pub std::marker::PhantomData<T>);
134
135/// The specialized arm of [`ColumnTypeProbe`].
136pub trait ProbeColumnType {
137    fn column_type(&self) -> Option<SqlColumnType>;
138}
139
140impl<T: mz_repr::AsColumnType> ProbeColumnType for ColumnTypeProbe<T> {
141    fn column_type(&self) -> Option<SqlColumnType> {
142        Some(T::as_column_type())
143    }
144}
145
146/// The fallback arm of [`ColumnTypeProbe`], reached through one more autoref
147/// than [`ProbeColumnType`] so it only applies when that one does not.
148pub trait ProbeColumnTypeFallback {
149    fn column_type(&self) -> Option<SqlColumnType>;
150}
151
152impl<T> ProbeColumnTypeFallback for &ColumnTypeProbe<T> {
153    fn column_type(&self) -> Option<SqlColumnType> {
154        None
155    }
156}
157
158#[derive(Debug, Serialize)]
159pub struct UnaryFuncProperties {
160    pub variant: String,
161    pub propagates_nulls: bool,
162    pub introduces_nulls: bool,
163    pub could_error: bool,
164    pub preserves_uniqueness: bool,
165    pub is_monotone: bool,
166    pub is_eliminable_cast: bool,
167    /// Canonical name of the inverse function, if one is declared.
168    pub inverse: Option<&'static str>,
169    pub input_types: Vec<String>,
170    pub output_type: Option<String>,
171    /// The `#[sqlfunc]` parameter and return types as written, see
172    /// [`SqlFuncSource::signature`].
173    pub sqlfunc_signature: Option<&'static str>,
174}
175
176#[derive(Debug, Serialize)]
177pub struct BinaryFuncProperties {
178    pub variant: String,
179    pub propagates_nulls: bool,
180    pub introduces_nulls: bool,
181    pub could_error: bool,
182    pub is_monotone: (bool, bool),
183    pub is_infinity_monotone: bool,
184    pub is_infix_op: bool,
185    /// Canonical name of the negated function, if one is declared.
186    pub negate: Option<&'static str>,
187    pub input_types: Vec<String>,
188    pub output_type: Option<String>,
189    /// The `#[sqlfunc]` parameter and return types as written, see
190    /// [`SqlFuncSource::signature`].
191    pub sqlfunc_signature: Option<&'static str>,
192}
193
194#[derive(Debug, Serialize)]
195pub struct VariadicFuncProperties {
196    pub variant: String,
197    pub propagates_nulls: bool,
198    pub introduces_nulls: bool,
199    pub could_error: bool,
200    pub is_monotone: bool,
201    pub is_associative: bool,
202    pub is_infix_op: bool,
203    pub input_types: Vec<String>,
204    pub output_type: Option<String>,
205    /// The `#[sqlfunc]` parameter and return types as written, see
206    /// [`SqlFuncSource::signature`].
207    pub sqlfunc_signature: Option<&'static str>,
208}
209
210impl FuncRegistry {
211    /// Records the properties of every variant of the three function enums.
212    ///
213    /// Panics if a variant can neither be constructed from its canonical name
214    /// nor has a [`Sample`] in this module, or if a sample's input count does
215    /// not fit its function.
216    pub fn build() -> FuncRegistry {
217        FuncRegistry {
218            unary: collect(
219                "UnaryFunc",
220                UnaryFunc::variant_names(),
221                UnaryFunc::from_variant_name,
222                UnaryFunc::sqlfunc_input_types,
223                unary_samples(),
224                UnaryFunc::variant_name,
225                UnaryFuncProperties::of,
226            ),
227            binary: collect(
228                "BinaryFunc",
229                BinaryFunc::variant_names(),
230                BinaryFunc::from_variant_name,
231                BinaryFunc::sqlfunc_input_types,
232                binary_samples(),
233                BinaryFunc::variant_name,
234                BinaryFuncProperties::of,
235            ),
236            variadic: collect(
237                "VariadicFunc",
238                VariadicFunc::variant_names(),
239                VariadicFunc::from_variant_name,
240                VariadicFunc::sqlfunc_input_types,
241                variadic_samples(),
242                VariadicFunc::variant_name,
243                VariadicFuncProperties::of,
244            ),
245        }
246    }
247}
248
249/// Records the properties of every variant of one function enum, keyed as
250/// described on [`Sample`].
251///
252/// A variant's primary record comes from its unlabeled hand-written sample if
253/// there is one, otherwise from a payload-free instance built by `construct`
254/// (see `from_variant_name` on the enums), probed at the column types
255/// `natural_inputs` reports for it, if any. Hand-written primaries take
256/// precedence so a constructible variant can still be probed at chosen input
257/// types.
258///
259/// A variant whose payload has fields must have at least one hand-written
260/// sample, even when the payload is constructible with every field defaulted,
261/// because a defaulted payload sees only one branch of any property that
262/// depends on it.
263///
264/// Panics if a variant has no primary or no required sample, if two samples
265/// share a name and label, or if probing a sample panics, in which case the
266/// panic names the sample.
267fn collect<F: fmt::Debug, P>(
268    enum_name: &str,
269    names: impl Iterator<Item = &'static str>,
270    construct: fn(&str) -> Option<F>,
271    natural_inputs: fn(&F) -> Option<Vec<SqlColumnType>>,
272    samples: Vec<Sample<F>>,
273    variant_name: fn(&F) -> &'static str,
274    properties: fn(&Sample<F>) -> P,
275) -> BTreeMap<String, P> {
276    let samples_fn = format!(
277        "{}_samples() in src/expr/src/scalar/func/registry.rs",
278        enum_name.trim_end_matches("Func").to_lowercase()
279    );
280    let mut by_name: BTreeMap<(&'static str, &'static str), Sample<F>> = BTreeMap::new();
281    let mut hand_written = BTreeSet::new();
282    for sample in samples {
283        let name = variant_name(&sample.func);
284        let label = sample.label;
285        hand_written.insert(name);
286        let duplicate = by_name.insert((name, label), sample).is_some();
287        assert!(
288            !duplicate,
289            "duplicate {enum_name} sample for `{name}` with label `{label}`"
290        );
291    }
292
293    let mut records = BTreeMap::new();
294    for name in names {
295        let primary = by_name.remove(&(name, "")).or_else(|| {
296            construct(name).map(|func| Sample {
297                input_types: natural_inputs(&func).unwrap_or_default(),
298                func,
299                label: "",
300            })
301        });
302        let primary = primary.unwrap_or_else(|| {
303            panic!(
304                "{enum_name} variant `{name}` cannot be constructed from its name because \
305                 its payload needs data. Add a representative Sample for it to {samples_fn}."
306            )
307        });
308        assert!(
309            hand_written.contains(name) || !payload_has_fields(&primary.func),
310            "{enum_name} variant `{name}` has a payload with fields but only its defaulted \
311             instance is recorded. Add a Sample with a non-default payload to {samples_fn}, \
312             labeled if the defaulted instance should stay the primary record."
313        );
314        records.insert(
315            name.to_string(),
316            probe(enum_name, name, &primary, properties),
317        );
318    }
319    for ((name, label), sample) in by_name {
320        assert!(
321            records.contains_key(name),
322            "{enum_name} sample `{name}[{label}]` names an unknown variant"
323        );
324        records.insert(
325            format!("{name}[{label}]"),
326            probe(enum_name, &format!("{name}[{label}]"), &sample, properties),
327        );
328    }
329    records
330}
331
332/// Whether a variant's payload struct has any fields.
333///
334/// Decided from the `Debug` rendering, which every payload derives: a unit
335/// payload renders as `Variant(Payload)`, one with fields opens a second
336/// parenthesis or a brace.
337fn payload_has_fields<F: fmt::Debug>(func: &F) -> bool {
338    let rendered = format!("{func:?}");
339    rendered.matches('(').count() > 1 || rendered.contains('{')
340}
341
342/// Runs `properties` on a sample, attributing any panic (typically an
343/// `output_sql_type` impl rejecting the sample's input types) to the sample.
344fn probe<F, P>(
345    enum_name: &str,
346    key: &str,
347    sample: &Sample<F>,
348    properties: fn(&Sample<F>) -> P,
349) -> P {
350    mz_ore::panic::catch_unwind_str(std::panic::AssertUnwindSafe(|| properties(sample)))
351        .unwrap_or_else(|message| {
352            panic!(
353                "probing {enum_name} sample `{key}` panicked: {message}\n\
354                 Check the sample's input types in src/expr/src/scalar/func/registry.rs."
355            )
356        })
357}
358
359/// The serde variant name, which is what the stable LIR format stores.
360fn variant_ident<F: Serialize>(func: &F) -> String {
361    let value = serde_json::to_value(func).expect("function variants serialize");
362    match value {
363        serde_json::Value::Object(map) if map.len() == 1 => {
364            map.into_iter().next().expect("one entry").0
365        }
366        serde_json::Value::String(name) => name,
367        other => panic!("unexpected function variant encoding: {other}"),
368    }
369}
370
371fn type_strings(types: &[SqlColumnType]) -> Vec<String> {
372    types.iter().map(ToString::to_string).collect()
373}
374
375impl UnaryFuncProperties {
376    fn of(sample: &Sample<UnaryFunc>) -> Record<Self> {
377        let func = &sample.func;
378        let output_type = match sample.input_types.as_slice() {
379            [] => None,
380            [input] => Some(func.output_sql_type(input.clone()).to_string()),
381            _ => panic!("unary sample `{}` needs exactly one input type", func),
382        };
383        let source = func.sqlfunc_source();
384        let properties = UnaryFuncProperties {
385            variant: variant_ident(func),
386            propagates_nulls: func.propagates_nulls(),
387            introduces_nulls: func.introduces_nulls(),
388            could_error: func.could_error(),
389            preserves_uniqueness: func.preserves_uniqueness(),
390            is_monotone: func.is_monotone(),
391            is_eliminable_cast: func.is_eliminable_cast(),
392            inverse: func.inverse().map(|f| f.variant_name()),
393            input_types: type_strings(&sample.input_types),
394            output_type,
395            sqlfunc_signature: source.map(|s| s.signature),
396        };
397        Record {
398            properties,
399            source: FuncSource::of(func.to_string(), source),
400        }
401    }
402}
403
404impl BinaryFuncProperties {
405    fn of(sample: &Sample<BinaryFunc>) -> Record<Self> {
406        let func = &sample.func;
407        let output_type = match sample.input_types.as_slice() {
408            [] => None,
409            [_, _] => Some(func.output_sql_type(&sample.input_types).to_string()),
410            _ => panic!("binary sample `{}` needs exactly two input types", func),
411        };
412        let source = func.sqlfunc_source();
413        let properties = BinaryFuncProperties {
414            variant: variant_ident(func),
415            propagates_nulls: func.propagates_nulls(),
416            introduces_nulls: func.introduces_nulls(),
417            could_error: func.could_error(),
418            is_monotone: func.is_monotone(),
419            is_infinity_monotone: func.is_infinity_monotone(),
420            is_infix_op: func.is_infix_op(),
421            negate: func.negate().map(|f| f.variant_name()),
422            input_types: type_strings(&sample.input_types),
423            output_type,
424            sqlfunc_signature: source.map(|s| s.signature),
425        };
426        Record {
427            properties,
428            source: FuncSource::of(func.to_string(), source),
429        }
430    }
431}
432
433impl VariadicFuncProperties {
434    fn of(sample: &Sample<VariadicFunc>) -> Record<Self> {
435        let func = &sample.func;
436        let output_type = (!sample.input_types.is_empty())
437            .then(|| func.output_sql_type(sample.input_types.clone()).to_string());
438        let source = func.sqlfunc_source();
439        let properties = VariadicFuncProperties {
440            variant: variant_ident(func),
441            propagates_nulls: func.propagates_nulls(),
442            introduces_nulls: func.introduces_nulls(),
443            could_error: func.could_error(),
444            is_monotone: func.is_monotone(),
445            is_associative: func.is_associative(),
446            is_infix_op: func.is_infix_op(),
447            input_types: type_strings(&sample.input_types),
448            output_type,
449            sqlfunc_signature: source.map(|s| s.signature),
450        };
451        Record {
452            properties,
453            source: FuncSource::of(func.to_string(), source),
454        }
455    }
456}
457
458fn column(i: usize) -> Box<MirScalarExpr> {
459    Box::new(MirScalarExpr::column(i))
460}
461
462fn record_type() -> SqlScalarType {
463    SqlScalarType::Record {
464        fields: [
465            ("a".into(), SqlScalarType::Int32.nullable(false)),
466            ("b".into(), SqlScalarType::String.nullable(true)),
467        ]
468        .into(),
469        custom_id: None,
470    }
471}
472
473fn list_type(element: SqlScalarType) -> SqlScalarType {
474    SqlScalarType::List {
475        element_type: Box::new(element),
476        custom_id: None,
477    }
478}
479
480fn map_type(value: SqlScalarType) -> SqlScalarType {
481    SqlScalarType::Map {
482        value_type: Box::new(value),
483        custom_id: None,
484    }
485}
486
487fn range_type(element: SqlScalarType) -> SqlScalarType {
488    SqlScalarType::Range {
489        element_type: Box::new(element),
490    }
491}
492
493fn array_type(element: SqlScalarType) -> SqlScalarType {
494    SqlScalarType::Array(Box::new(element))
495}
496
497fn regex() -> Regex {
498    Regex::new("a+", false).expect("valid regex")
499}
500
501fn unary<F: Into<UnaryFunc>>(func: F, input: SqlScalarType) -> Sample<UnaryFunc> {
502    Sample {
503        func: func.into(),
504        input_types: vec![input.nullable(false)],
505        label: "",
506    }
507}
508
509fn unary_samples() -> Vec<Sample<UnaryFunc>> {
510    let timestamp_precision = Some(TimestampPrecision::try_from(3i64).expect("valid precision"));
511    let numeric_scale = NumericMaxScale::try_from(2i64).expect("valid scale");
512    let char_length = Some(CharLength::try_from(5i64).expect("valid length"));
513    let varchar_length = Some(VarCharMaxLength::try_from(5i64).expect("valid length"));
514    let tz = mz_pgtz::timezone::Timezone::Tz(chrono_tz::Tz::UTC);
515    vec![
516        // Casts between container types, which carry their target type.
517        unary(
518            CastArrayToString {
519                ty: array_type(SqlScalarType::Int32),
520            },
521            array_type(SqlScalarType::Int32),
522        ),
523        unary(
524            CastArrayToJsonb {
525                cast_element: column(0),
526            },
527            array_type(SqlScalarType::Int32),
528        ),
529        unary(
530            CastArrayToArray {
531                return_ty: array_type(SqlScalarType::Int64),
532                cast_expr: column(0),
533            },
534            array_type(SqlScalarType::Int32),
535        ),
536        unary(
537            CastListToString {
538                ty: list_type(SqlScalarType::Int32),
539            },
540            list_type(SqlScalarType::Int32),
541        ),
542        unary(
543            CastListToJsonb {
544                cast_element: column(0),
545            },
546            list_type(SqlScalarType::Int32),
547        ),
548        unary(
549            CastList1ToList2 {
550                return_ty: list_type(SqlScalarType::Int64),
551                cast_expr: column(0),
552            },
553            list_type(SqlScalarType::Int32),
554        ),
555        unary(
556            CastMapToString {
557                ty: map_type(SqlScalarType::Int32),
558            },
559            map_type(SqlScalarType::Int32),
560        ),
561        unary(
562            MapBuildFromRecordList {
563                value_type: SqlScalarType::Int32,
564            },
565            list_type(record_type()),
566        ),
567        unary(
568            CastRangeToString {
569                ty: range_type(SqlScalarType::Int32),
570            },
571            range_type(SqlScalarType::Int32),
572        ),
573        unary(CastRecordToString { ty: record_type() }, record_type()),
574        unary(
575            CastRecord1ToRecord2 {
576                return_ty: record_type(),
577                cast_exprs: vec![*column(0), *column(1)].into(),
578            },
579            record_type(),
580        ),
581        unary(RecordGet(1), record_type()),
582        unary(
583            CastStringToArray {
584                return_ty: array_type(SqlScalarType::Int32),
585                cast_expr: column(0),
586            },
587            SqlScalarType::String,
588        ),
589        unary(
590            CastStringToList {
591                return_ty: list_type(SqlScalarType::Int32),
592                cast_expr: column(0),
593            },
594            SqlScalarType::String,
595        ),
596        unary(
597            CastStringToMap {
598                return_ty: map_type(SqlScalarType::Int32),
599                cast_expr: column(0),
600            },
601            SqlScalarType::String,
602        ),
603        unary(
604            CastStringToRange {
605                return_ty: range_type(SqlScalarType::Int32),
606                cast_expr: column(0),
607            },
608            SqlScalarType::String,
609        ),
610        // Length-parameterized string casts.
611        unary(
612            CastStringToChar {
613                length: char_length,
614                fail_on_len: true,
615            },
616            SqlScalarType::String,
617        ),
618        unary(
619            PadChar {
620                length: char_length,
621            },
622            SqlScalarType::String,
623        ),
624        unary(
625            CastStringToVarChar {
626                length: varchar_length,
627                fail_on_len: true,
628            },
629            SqlScalarType::String,
630        ),
631        // Precision and scale parameterized casts.
632        unary(
633            CastTimestampToTimestampTz {
634                from: None,
635                to: timestamp_precision,
636            },
637            SqlScalarType::Timestamp { precision: None },
638        ),
639        unary(
640            CastTimestampTzToTimestamp {
641                from: None,
642                to: timestamp_precision,
643            },
644            SqlScalarType::TimestampTz { precision: None },
645        ),
646        unary(
647            AdjustTimestampPrecision {
648                from: None,
649                to: timestamp_precision,
650            },
651            SqlScalarType::Timestamp { precision: None },
652        ),
653        unary(
654            AdjustTimestampTzPrecision {
655                from: None,
656                to: timestamp_precision,
657            },
658            SqlScalarType::TimestampTz { precision: None },
659        ),
660        unary(
661            AdjustNumericScale(numeric_scale),
662            SqlScalarType::Numeric { max_scale: None },
663        ),
664        // Pattern matching.
665        unary(
666            IsLikeMatch(like_pattern::compile("%a%", false).expect("valid pattern")),
667            SqlScalarType::String,
668        ),
669        unary(IsRegexpMatch(regex()), SqlScalarType::String),
670        unary(RegexpMatch(regex()), SqlScalarType::String),
671        unary(RegexpSplitToArray(regex()), SqlScalarType::String),
672        // Date and time functions carrying units, time zones and formats.
673        unary(
674            ExtractInterval(DateTimeUnits::Epoch),
675            SqlScalarType::Interval,
676        ),
677        unary(ExtractTime(DateTimeUnits::Epoch), SqlScalarType::Time),
678        unary(
679            ExtractTimestamp(DateTimeUnits::Epoch),
680            SqlScalarType::Timestamp { precision: None },
681        ),
682        unary(
683            ExtractTimestampTz(DateTimeUnits::Epoch),
684            SqlScalarType::TimestampTz { precision: None },
685        ),
686        unary(ExtractDate(DateTimeUnits::Epoch), SqlScalarType::Date),
687        unary(
688            DatePartInterval(DateTimeUnits::Epoch),
689            SqlScalarType::Interval,
690        ),
691        unary(DatePartTime(DateTimeUnits::Epoch), SqlScalarType::Time),
692        unary(
693            DatePartTimestamp(DateTimeUnits::Epoch),
694            SqlScalarType::Timestamp { precision: None },
695        ),
696        unary(
697            DatePartTimestampTz(DateTimeUnits::Epoch),
698            SqlScalarType::TimestampTz { precision: None },
699        ),
700        unary(
701            DateTruncTimestamp(DateTimeUnits::Day),
702            SqlScalarType::Timestamp { precision: None },
703        ),
704        unary(
705            DateTruncTimestampTz(DateTimeUnits::Day),
706            SqlScalarType::TimestampTz { precision: None },
707        ),
708        unary(
709            TimezoneTimestamp(tz.clone()),
710            SqlScalarType::Timestamp { precision: None },
711        ),
712        unary(
713            TimezoneTimestampTz(tz.clone()),
714            SqlScalarType::TimestampTz { precision: None },
715        ),
716        unary(
717            TimezoneTime {
718                tz,
719                wall_time: chrono::NaiveDateTime::default(),
720            },
721            SqlScalarType::Time,
722        ),
723        unary(
724            ToCharTimestamp {
725                format_string: "YYYY".into(),
726                format: DateTimeFormat::compile("YYYY"),
727            },
728            SqlScalarType::Timestamp { precision: None },
729        ),
730        unary(
731            ToCharTimestampTz {
732                format_string: "YYYY".into(),
733                format: DateTimeFormat::compile("YYYY"),
734            },
735            SqlScalarType::TimestampTz { precision: None },
736        ),
737        // Payload-free hand-written casts, probed at their natural input.
738        unary(CastStringToInt2Vector, SqlScalarType::String),
739        unary(CastDateToTimestamp(None), SqlScalarType::Date),
740        unary(CastDateToTimestampTz(None), SqlScalarType::Date),
741        unary(CastStringToTimestamp(None), SqlScalarType::String),
742        unary(CastStringToTimestampTz(None), SqlScalarType::String),
743        unary(
744            CastDateToTimestamp(timestamp_precision),
745            SqlScalarType::Date,
746        )
747        .labeled("precision"),
748        unary(
749            CastDateToTimestampTz(timestamp_precision),
750            SqlScalarType::Date,
751        )
752        .labeled("precision"),
753        unary(
754            CastStringToTimestamp(timestamp_precision),
755            SqlScalarType::String,
756        )
757        .labeled("precision"),
758        unary(
759            CastStringToTimestampTz(timestamp_precision),
760            SqlScalarType::String,
761        )
762        .labeled("precision"),
763        // Without a length to enforce, these casts cannot error.
764        unary(
765            CastStringToChar {
766                length: None,
767                fail_on_len: false,
768            },
769            SqlScalarType::String,
770        )
771        .labeled("unbounded"),
772        unary(PadChar { length: None }, SqlScalarType::String).labeled("unbounded"),
773        unary(
774            CastStringToVarChar {
775                length: None,
776                fail_on_len: false,
777            },
778            SqlScalarType::String,
779        )
780        .labeled("unbounded"),
781        // Widening a precision preserves uniqueness, narrowing does not.
782        unary(
783            CastTimestampToTimestampTz {
784                from: timestamp_precision,
785                to: None,
786            },
787            SqlScalarType::Timestamp {
788                precision: timestamp_precision,
789            },
790        )
791        .labeled("widening"),
792        unary(
793            CastTimestampTzToTimestamp {
794                from: timestamp_precision,
795                to: None,
796            },
797            SqlScalarType::TimestampTz {
798                precision: timestamp_precision,
799            },
800        )
801        .labeled("widening"),
802        unary(
803            AdjustTimestampPrecision {
804                from: timestamp_precision,
805                to: None,
806            },
807            SqlScalarType::Timestamp {
808                precision: timestamp_precision,
809            },
810        )
811        .labeled("widening"),
812        unary(
813            AdjustTimestampTzPrecision {
814                from: timestamp_precision,
815                to: None,
816            },
817            SqlScalarType::TimestampTz {
818                precision: timestamp_precision,
819            },
820        )
821        .labeled("widening"),
822        // Units below the most significant ones are not monotone.
823        unary(
824            ExtractInterval(DateTimeUnits::Month),
825            SqlScalarType::Interval,
826        )
827        .labeled("month"),
828        unary(ExtractTime(DateTimeUnits::Minute), SqlScalarType::Time).labeled("minute"),
829        unary(
830            ExtractTimestamp(DateTimeUnits::Month),
831            SqlScalarType::Timestamp { precision: None },
832        )
833        .labeled("month"),
834        unary(
835            ExtractTimestampTz(DateTimeUnits::Month),
836            SqlScalarType::TimestampTz { precision: None },
837        )
838        .labeled("month"),
839        unary(ExtractDate(DateTimeUnits::Month), SqlScalarType::Date).labeled("month"),
840        unary(
841            DatePartInterval(DateTimeUnits::Month),
842            SqlScalarType::Interval,
843        )
844        .labeled("month"),
845        unary(DatePartTime(DateTimeUnits::Minute), SqlScalarType::Time).labeled("minute"),
846        unary(
847            DatePartTimestamp(DateTimeUnits::Month),
848            SqlScalarType::Timestamp { precision: None },
849        )
850        .labeled("month"),
851        unary(
852            DatePartTimestampTz(DateTimeUnits::Month),
853            SqlScalarType::TimestampTz { precision: None },
854        )
855        .labeled("month"),
856    ]
857    .into_iter()
858    .chain(numeric_cast_samples(numeric_scale))
859    .collect()
860}
861
862/// The numeric casts, each without a scale (the primary record) and with
863/// one. Only the scaled cast can error, because it rounds.
864fn numeric_cast_samples(scale: NumericMaxScale) -> Vec<Sample<UnaryFunc>> {
865    let casts: [(fn(Option<NumericMaxScale>) -> UnaryFunc, SqlScalarType); 10] = [
866        (|s| CastInt16ToNumeric(s).into(), SqlScalarType::Int16),
867        (|s| CastInt32ToNumeric(s).into(), SqlScalarType::Int32),
868        (|s| CastInt64ToNumeric(s).into(), SqlScalarType::Int64),
869        (|s| CastUint16ToNumeric(s).into(), SqlScalarType::UInt16),
870        (|s| CastUint32ToNumeric(s).into(), SqlScalarType::UInt32),
871        (|s| CastUint64ToNumeric(s).into(), SqlScalarType::UInt64),
872        (|s| CastFloat32ToNumeric(s).into(), SqlScalarType::Float32),
873        (|s| CastFloat64ToNumeric(s).into(), SqlScalarType::Float64),
874        (|s| CastStringToNumeric(s).into(), SqlScalarType::String),
875        (|s| CastJsonbToNumeric(s).into(), SqlScalarType::Jsonb),
876    ];
877    casts
878        .into_iter()
879        .flat_map(|(cast, input)| {
880            [
881                unary(cast(None), input.clone()),
882                unary(cast(Some(scale)), input).labeled("scale"),
883            ]
884        })
885        .collect()
886}
887
888fn binary<F: Into<BinaryFunc>>(
889    func: F,
890    left: SqlScalarType,
891    right: SqlScalarType,
892) -> Sample<BinaryFunc> {
893    Sample {
894        func: func.into(),
895        input_types: vec![left.nullable(false), right.nullable(false)],
896        label: "",
897    }
898}
899
900fn binary_samples() -> Vec<Sample<BinaryFunc>> {
901    vec![
902        binary(
903            ListLengthMax { max_layer: 1 },
904            list_type(SqlScalarType::Int32),
905            SqlScalarType::Int64,
906        ),
907        binary(
908            RegexpReplace {
909                regex: regex(),
910                limit: 1,
911            },
912            SqlScalarType::String,
913            SqlScalarType::String,
914        ),
915    ]
916}
917
918fn variadic<F: Into<VariadicFunc>>(func: F, inputs: Vec<SqlScalarType>) -> Sample<VariadicFunc> {
919    Sample {
920        func: func.into(),
921        input_types: inputs.into_iter().map(|ty| ty.nullable(false)).collect(),
922        label: "",
923    }
924}
925
926fn variadic_samples() -> Vec<Sample<VariadicFunc>> {
927    vec![
928        variadic(
929            ArrayCreate {
930                elem_type: SqlScalarType::Int32,
931            },
932            vec![SqlScalarType::Int32, SqlScalarType::Int32],
933        ),
934        variadic(
935            ArrayFill {
936                elem_type: SqlScalarType::Int32,
937            },
938            vec![SqlScalarType::Int32, array_type(SqlScalarType::Int32)],
939        ),
940        variadic(
941            ArrayIndex { offset: 1 },
942            vec![array_type(SqlScalarType::Int32), SqlScalarType::Int64],
943        ),
944        variadic(
945            ArrayToString {
946                elem_type: SqlScalarType::Int32,
947            },
948            vec![array_type(SqlScalarType::Int32), SqlScalarType::String],
949        ),
950        variadic(
951            ListCreate {
952                elem_type: SqlScalarType::Int32,
953            },
954            vec![SqlScalarType::Int32, SqlScalarType::Int32],
955        ),
956        variadic(
957            RangeCreate {
958                elem_type: SqlScalarType::Int32,
959            },
960            vec![
961                SqlScalarType::Int32,
962                SqlScalarType::Int32,
963                SqlScalarType::String,
964            ],
965        ),
966        variadic(
967            RecordCreate {
968                field_names: vec![ColumnName::from("a".to_string())],
969            },
970            vec![SqlScalarType::Int32],
971        ),
972        variadic(
973            MapBuild {
974                value_type: SqlScalarType::Int32,
975            },
976            vec![SqlScalarType::String, SqlScalarType::Int32],
977        ),
978        variadic(
979            CaseLiteral {
980                lookup: vec![CaseLiteralEntry {
981                    literal: StableRow::from(Row::pack_slice(&[Datum::Int32(1)])),
982                    expr_index: 0,
983                }],
984                return_type: SqlScalarType::String.nullable(true),
985            },
986            vec![SqlScalarType::Int32, SqlScalarType::String],
987        ),
988        // Payload-free hand-written functions, probed at their natural inputs.
989        variadic(And, vec![SqlScalarType::Bool, SqlScalarType::Bool]),
990        variadic(Or, vec![SqlScalarType::Bool, SqlScalarType::Bool]),
991        variadic(Coalesce, vec![SqlScalarType::Int32, SqlScalarType::Int32]),
992        variadic(Greatest, vec![SqlScalarType::Int32, SqlScalarType::Int32]),
993        variadic(Least, vec![SqlScalarType::Int32, SqlScalarType::Int32]),
994        variadic(
995            ErrorIfNull,
996            vec![SqlScalarType::Int32, SqlScalarType::String],
997        ),
998    ]
999}
1000
1001#[cfg(test)]
1002mod tests {
1003    use super::*;
1004
1005    #[mz_ore::test]
1006    fn payload_field_detection() {
1007        fn unary(func: impl Into<UnaryFunc>) -> UnaryFunc {
1008            func.into()
1009        }
1010        assert!(!payload_has_fields(&unary(Not)));
1011        assert!(payload_has_fields(&unary(CastInt32ToNumeric(None))));
1012        assert!(payload_has_fields(&unary(PadChar { length: None })));
1013        assert!(!payload_has_fields(&VariadicFunc::from(And)));
1014        assert!(payload_has_fields(&VariadicFunc::from(ArrayIndex {
1015            offset: 0
1016        })));
1017    }
1018}