1use 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#[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 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#[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#[derive(Debug)]
99pub struct Record<P> {
100 pub properties: P,
101 pub source: FuncSource,
102}
103
104#[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#[derive(Debug)]
133pub struct ColumnTypeProbe<T>(pub std::marker::PhantomData<T>);
134
135pub 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
146pub 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 pub inverse: Option<&'static str>,
169 pub input_types: Vec<String>,
170 pub output_type: Option<String>,
171 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 pub negate: Option<&'static str>,
187 pub input_types: Vec<String>,
188 pub output_type: Option<String>,
189 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 pub sqlfunc_signature: Option<&'static str>,
208}
209
210impl FuncRegistry {
211 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
249fn 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
332fn payload_has_fields<F: fmt::Debug>(func: &F) -> bool {
338 let rendered = format!("{func:?}");
339 rendered.matches('(').count() > 1 || rendered.contains('{')
340}
341
342fn 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
359fn 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 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 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 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 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 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 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 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 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 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
862fn 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 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}