Skip to main content

mz_expr/scalar/func/
macros.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// Convenience macro for generating `inverse` values.
11macro_rules! to_unary {
12    ($f:expr) => {
13        Some(crate::UnaryFunc::from($f))
14    };
15}
16
17#[cfg(test)]
18mod test {
19    use mz_expr_derive::sqlfunc;
20    use mz_repr::SqlScalarType;
21
22    use crate::EvalError;
23    use crate::scalar::func::LazyUnaryFunc;
24
25    #[sqlfunc(sqlname = "INFALLIBLE")]
26    fn infallible1(a: f32) -> f32 {
27        a
28    }
29
30    #[sqlfunc]
31    fn infallible2(a: Option<f32>) -> f32 {
32        a.unwrap_or_default()
33    }
34
35    #[sqlfunc]
36    fn infallible3(a: f32) -> Option<f32> {
37        Some(a)
38    }
39
40    #[mz_ore::test]
41    fn elision_rules_infallible() {
42        assert_eq!(format!("{}", Infallible1), "INFALLIBLE");
43        assert!(Infallible1.propagates_nulls());
44        assert!(!Infallible1.introduces_nulls());
45
46        assert!(!Infallible2.propagates_nulls());
47        assert!(!Infallible2.introduces_nulls());
48
49        assert!(Infallible3.propagates_nulls());
50        assert!(Infallible3.introduces_nulls());
51    }
52
53    #[mz_ore::test]
54    fn output_types_infallible() {
55        assert_eq!(
56            Infallible1.output_sql_type(SqlScalarType::Float32.nullable(true)),
57            SqlScalarType::Float32.nullable(true)
58        );
59        assert_eq!(
60            Infallible1.output_sql_type(SqlScalarType::Float32.nullable(false)),
61            SqlScalarType::Float32.nullable(false)
62        );
63
64        assert_eq!(
65            Infallible2.output_sql_type(SqlScalarType::Float32.nullable(true)),
66            SqlScalarType::Float32.nullable(false)
67        );
68        assert_eq!(
69            Infallible2.output_sql_type(SqlScalarType::Float32.nullable(false)),
70            SqlScalarType::Float32.nullable(false)
71        );
72
73        assert_eq!(
74            Infallible3.output_sql_type(SqlScalarType::Float32.nullable(true)),
75            SqlScalarType::Float32.nullable(true)
76        );
77        assert_eq!(
78            Infallible3.output_sql_type(SqlScalarType::Float32.nullable(false)),
79            SqlScalarType::Float32.nullable(true)
80        );
81    }
82
83    #[sqlfunc]
84    fn fallible1(a: f32) -> Result<f32, EvalError> {
85        Ok(a)
86    }
87
88    #[sqlfunc]
89    fn fallible2(a: Option<f32>) -> Result<f32, EvalError> {
90        Ok(a.unwrap_or_default())
91    }
92
93    #[sqlfunc]
94    fn fallible3(a: f32) -> Result<Option<f32>, EvalError> {
95        Ok(Some(a))
96    }
97
98    #[mz_ore::test]
99    fn elision_rules_fallible() {
100        assert!(Fallible1.propagates_nulls());
101        assert!(!Fallible1.introduces_nulls());
102
103        assert!(!Fallible2.propagates_nulls());
104        assert!(!Fallible2.introduces_nulls());
105
106        assert!(Fallible3.propagates_nulls());
107        assert!(Fallible3.introduces_nulls());
108    }
109
110    #[mz_ore::test]
111    fn output_types_fallible() {
112        assert_eq!(
113            Fallible1.output_sql_type(SqlScalarType::Float32.nullable(true)),
114            SqlScalarType::Float32.nullable(true)
115        );
116        assert_eq!(
117            Fallible1.output_sql_type(SqlScalarType::Float32.nullable(false)),
118            SqlScalarType::Float32.nullable(false)
119        );
120
121        assert_eq!(
122            Fallible2.output_sql_type(SqlScalarType::Float32.nullable(true)),
123            SqlScalarType::Float32.nullable(false)
124        );
125        assert_eq!(
126            Fallible2.output_sql_type(SqlScalarType::Float32.nullable(false)),
127            SqlScalarType::Float32.nullable(false)
128        );
129
130        assert_eq!(
131            Fallible3.output_sql_type(SqlScalarType::Float32.nullable(true)),
132            SqlScalarType::Float32.nullable(true)
133        );
134        assert_eq!(
135            Fallible3.output_sql_type(SqlScalarType::Float32.nullable(false)),
136            SqlScalarType::Float32.nullable(true)
137        );
138    }
139}
140
141/// Temporary macro that generates the equivalent of what enum_dispatch will do in the future. We
142/// need this manual macro implementation to delegate to the previous manual implementation for
143/// variants that use the old definitions.
144///
145/// Once everything is handled by this macro we can remove it and replace it with `enum_dispatch`
146///
147/// Variants marked `(E)` have payload structs generic over the stored expression type. The
148/// marker must be the literal ident `E`, matching the generic parameter the macro declares.
149macro_rules! derive_unary {
150    // Internal rules that rebuild one variant payload during expression
151    // conversion. Payloads without a marker are expression free and clone.
152    (@try_map_payload $f:ident) => { $f.clone() };
153    (@try_map_payload $f:ident $marker:ident) => { $f.try_map_expr()? };
154    (@map_payload $f:ident) => { $f.clone() };
155    (@map_payload $f:ident $marker:ident) => { $f.map_expr() };
156    // Internal rules that name a payload type at the MIR instantiation, for
157    // use in the impl block where `E` is not in scope.
158    (@mir_ty $name:ident) => { $name };
159    (@mir_ty $name:ident $marker:ident) => { $name<crate::MirScalarExpr> };
160    ($($name:ident $(($marker:ident))?),* $(,)?) => {
161        #[derive(
162            Ord, PartialOrd, Clone, Debug, Eq, PartialEq,
163            serde::Serialize, serde::Deserialize, Hash,
164                )]
165        pub enum UnaryFunc<E = crate::MirScalarExpr> {
166            $($name($name $(<$marker>)?),)*
167        }
168
169        // The `Eval` bound is required by the expression bearing payloads'
170        // `LazyUnaryFunc` impls, even for methods that never evaluate.
171        impl<E: Eval> UnaryFunc<E> {
172            pub fn eval<'a>(
173                &'a self,
174                datums: &[Datum<'a>],
175                temp_storage: &'a RowArena,
176                a: &'a impl Eval,
177            ) -> Result<Datum<'a>, EvalError> {
178                match self {
179                    $(Self::$name(f) => f.eval(datums, temp_storage, a),)*
180                }
181            }
182
183            pub fn output_sql_type(&self, input_type: SqlColumnType) -> SqlColumnType {
184                match self {
185                    $(Self::$name(f) => LazyUnaryFunc::output_sql_type(f, input_type),)*
186                }
187            }
188            pub fn output_type(&self, input_type: ReprColumnType) -> ReprColumnType {
189                match self {
190                    $(Self::$name(f) => LazyUnaryFunc::output_type(f, input_type),)*
191                }
192            }
193            pub fn propagates_nulls(&self) -> bool {
194                match self {
195                    $(Self::$name(f) => LazyUnaryFunc::propagates_nulls(f),)*
196                }
197            }
198            pub fn introduces_nulls(&self) -> bool {
199                match self {
200                    $(Self::$name(f) => LazyUnaryFunc::introduces_nulls(f),)*
201                }
202            }
203            pub fn preserves_uniqueness(&self) -> bool {
204                match self {
205                    $(Self::$name(f) => LazyUnaryFunc::preserves_uniqueness(f),)*
206                }
207            }
208            pub fn is_monotone(&self) -> bool {
209                match self {
210                    $(Self::$name(f) => LazyUnaryFunc::is_monotone(f),)*
211                }
212            }
213            pub fn could_error(&self) -> bool {
214                match self {
215                    $(Self::$name(f) => LazyUnaryFunc::could_error(f),)*
216                }
217            }
218            pub fn is_eliminable_cast(&self) -> bool {
219                match self {
220                    $(Self::$name(f) => LazyUnaryFunc::is_eliminable_cast(f),)*
221                }
222            }
223        }
224
225        impl<E> UnaryFunc<E> {
226            /// The canonical name of this variant, as declared by its
227            /// [`FuncName`](crate::func::FuncName) impl.
228            pub fn variant_name(&self) -> &'static str {
229                match self {
230                    $(Self::$name(_) =>
231                        <$name $(<$marker>)? as crate::func::FuncName>::NAME,)*
232                }
233            }
234
235            /// The `#[sqlfunc]` source of this variant's function, `None` for
236            /// hand-written functions. See [`FuncName::SQLFUNC`](crate::func::FuncName::SQLFUNC).
237            #[cfg(feature = "func-registry")]
238            pub fn sqlfunc_source(&self) -> Option<crate::func::SqlFuncSource> {
239                match self {
240                    $(Self::$name(_) =>
241                        <$name $(<$marker>)? as crate::func::FuncName>::SQLFUNC,)*
242                }
243            }
244
245            /// See [`FuncName::sqlfunc_input_types`](crate::func::FuncName::sqlfunc_input_types).
246            #[cfg(feature = "func-registry")]
247            pub fn sqlfunc_input_types(&self) -> Option<Vec<SqlColumnType>> {
248                match self {
249                    $(Self::$name(_) =>
250                        <$name $(<$marker>)? as crate::func::FuncName>::sqlfunc_input_types(),)*
251                }
252            }
253
254            /// Rebuilds this function with any stored expressions converted to
255            /// `E2`. Fails if any expression conversion fails, reporting the
256            /// first failure.
257            pub fn try_map_expr<'a, E2: TryFrom<&'a E>>(
258                &'a self,
259            ) -> Result<UnaryFunc<E2>, E2::Error> {
260                Ok(match self {
261                    $(Self::$name(f) => UnaryFunc::$name(
262                        derive_unary!(@try_map_payload f $($marker)?),
263                    ),)*
264                })
265            }
266
267            /// Rebuilds this function with any stored expressions converted to
268            /// `E2`.
269            pub fn map_expr<'a, E2: From<&'a E>>(&'a self) -> UnaryFunc<E2> {
270                match self {
271                    $(Self::$name(f) => UnaryFunc::$name(
272                        derive_unary!(@map_payload f $($marker)?),
273                    ),)*
274                }
275            }
276        }
277
278        // Methods that only exist at the MIR instantiation. `inverse` is an
279        // optimizer concept and the name lookup helpers serve MIR test
280        // tooling.
281        impl UnaryFunc {
282            pub fn inverse(&self) -> Option<UnaryFunc> {
283                match self {
284                    $(Self::$name(f) => LazyUnaryFunc::inverse(f),)*
285                }
286            }
287
288            /// Attempts to construct a `UnaryFunc` from the canonical name of
289            /// one of its variants, as declared by the variant's
290            /// [`FuncName`](crate::func::FuncName) impl (the name of the
291            /// underlying Rust function, e.g. `cast_int32_to_numeric`).
292            ///
293            /// Intended for test tooling that needs to name exact function
294            /// variants. Only variants whose inner function deserializes from
295            /// no data (unit functions and functions all of whose parameters
296            /// may default to `None`) are constructible. Returns `None` for
297            /// other variants and for unknown names. Returns the MIR
298            /// instantiation.
299            pub fn from_variant_name(name: &str) -> Option<Self> {
300                use crate::func::FuncName;
301                $(
302                    if name == <derive_unary!(@mir_ty $name $($marker)?) as FuncName>::NAME {
303                        return serde_json::from_value(
304                            serde_json::json!({ stringify!($name): null }),
305                        )
306                        .ok();
307                    }
308                )*
309                None
310            }
311
312            /// The canonical names of all variants, in declaration order.
313            pub fn variant_names() -> impl Iterator<Item = &'static str> {
314                use crate::func::FuncName;
315                [$(
316                    <derive_unary!(@mir_ty $name $($marker)?) as FuncName>::NAME,
317                )*]
318                .into_iter()
319            }
320        }
321
322        impl<E> fmt::Display for UnaryFunc<E> {
323            fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
324                match self {
325                    $(Self::$name(func) => func.fmt(f),)*
326                }
327            }
328        }
329
330        $(
331            impl<E> From<$name $(<$marker>)?> for crate::UnaryFunc<E> {
332                fn from(variant: $name $(<$marker>)?) -> Self {
333                    Self::$name(variant)
334                }
335            }
336        )*
337    }
338}
339
340/// Generates the `VariadicFunc` enum, its `impl` block,
341/// `Display` impl, and `From<InnerType>` impls for each variant.
342///
343/// All variants must use explicit `Name(Type)` syntax. When the variant name equals
344/// the inner type name, write e.g. `And(And)`.
345macro_rules! derive_variadic {
346    ($($name:ident ( $variant:ident )),* $(,)?) => {
347        #[derive(
348            Ord, PartialOrd, Clone, Debug, Eq, PartialEq,
349            serde::Serialize, serde::Deserialize, Hash,
350        )]
351        pub enum VariadicFunc {
352            $($name($variant),)*
353        }
354
355        impl VariadicFunc {
356            pub fn eval<'a>(
357                &'a self,
358                datums: &[Datum<'a>],
359                temp_storage: &'a RowArena,
360                exprs: &'a [impl Eval],
361            ) -> Result<Datum<'a>, EvalError> {
362                match self {
363                    $(Self::$name(f) => f.eval(datums, temp_storage, exprs),)*
364                }
365            }
366
367            pub fn output_sql_type(&self, input_types: Vec<SqlColumnType>) -> SqlColumnType {
368                match self {
369                    $(Self::$name(f) => LazyVariadicFunc::output_type(f, &input_types),)*
370                }
371            }
372
373            /// Computes the representation type of this variadic function.
374            ///
375            /// Wrapper around [`Self::output_sql_type`] that converts to representation types.
376            pub fn output_type(&self, input_types: Vec<ReprColumnType>) -> ReprColumnType {
377                let sql_types = input_types.iter().map(SqlColumnType::from_repr).collect();
378                ReprColumnType::from(&self.output_sql_type(sql_types))
379            }
380
381            pub fn propagates_nulls(&self) -> bool {
382                match self {
383                    $(Self::$name(f) => LazyVariadicFunc::propagates_nulls(f),)*
384                }
385            }
386
387            pub fn introduces_nulls(&self) -> bool {
388                match self {
389                    $(Self::$name(f) => LazyVariadicFunc::introduces_nulls(f),)*
390                }
391            }
392
393            pub fn could_error(&self) -> bool {
394                match self {
395                    $(Self::$name(f) => LazyVariadicFunc::could_error(f),)*
396                }
397            }
398
399            pub fn is_monotone(&self) -> bool {
400                match self {
401                    $(Self::$name(f) => LazyVariadicFunc::is_monotone(f),)*
402                }
403            }
404
405            pub fn is_associative(&self) -> bool {
406                match self {
407                    $(Self::$name(f) => LazyVariadicFunc::is_associative(f),)*
408                }
409            }
410
411            pub fn is_infix_op(&self) -> bool {
412                match self {
413                    $(Self::$name(f) => LazyVariadicFunc::is_infix_op(f),)*
414                }
415            }
416
417            /// The canonical name of this variant, as declared by its
418            /// [`FuncName`](crate::func::FuncName) impl.
419            pub fn variant_name(&self) -> &'static str {
420                match self {
421                    $(Self::$name(_) => <$variant as crate::func::FuncName>::NAME,)*
422                }
423            }
424
425            /// The `#[sqlfunc]` source of this variant's function, `None` for
426            /// hand-written functions. See [`FuncName::SQLFUNC`](crate::func::FuncName::SQLFUNC).
427            #[cfg(feature = "func-registry")]
428            pub fn sqlfunc_source(&self) -> Option<crate::func::SqlFuncSource> {
429                match self {
430                    $(Self::$name(_) => <$variant as crate::func::FuncName>::SQLFUNC,)*
431                }
432            }
433
434            /// See [`FuncName::sqlfunc_input_types`](crate::func::FuncName::sqlfunc_input_types).
435            #[cfg(feature = "func-registry")]
436            pub fn sqlfunc_input_types(&self) -> Option<Vec<SqlColumnType>> {
437                match self {
438                    $(Self::$name(_) =>
439                        <$variant as crate::func::FuncName>::sqlfunc_input_types(),)*
440                }
441            }
442
443            /// Attempts to construct a `VariadicFunc` from the canonical name
444            /// of one of its variants, as declared by the variant's
445            /// [`FuncName`](crate::func::FuncName) impl (the name of the
446            /// underlying Rust function, e.g. `record_create`).
447            ///
448            /// Intended for test tooling that needs to name exact function
449            /// variants. Only variants whose inner function deserializes from
450            /// no data (unit functions and functions all of whose parameters
451            /// may default to `None`) are constructible. Returns `None` for
452            /// other variants and for unknown names.
453            pub fn from_variant_name(name: &str) -> Option<Self> {
454                $(
455                    if name == <$variant as crate::func::FuncName>::NAME {
456                        return serde_json::from_value(
457                            serde_json::json!({ stringify!($name): null }),
458                        )
459                        .ok();
460                    }
461                )*
462                None
463            }
464
465            /// The canonical names of all variants, in declaration order.
466            pub fn variant_names() -> impl Iterator<Item = &'static str> {
467                [$(<$variant as crate::func::FuncName>::NAME,)*].into_iter()
468            }
469        }
470
471        impl fmt::Display for VariadicFunc {
472            fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
473                match self {
474                    $(Self::$name(func) => func.fmt(f),)*
475                }
476            }
477        }
478
479        $(
480            impl From<$variant> for crate::VariadicFunc {
481                fn from(variant: $variant) -> Self {
482                    Self::$name(variant)
483                }
484            }
485        )*
486    }
487}
488
489/// Generates the `BinaryFunc` enum, its `impl` block (with 8 delegating methods),
490/// `Display` impl, and `From<InnerType>` impls for each variant.
491///
492/// All variants must use explicit `Name(Type)` syntax. When the variant name equals
493/// the inner type name, write e.g. `AddInt16(AddInt16)`.
494macro_rules! derive_binary {
495    ($($name:ident ( $variant:ident )),* $(,)?) => {
496        #[derive(
497            Ord, PartialOrd, Clone, Debug, Eq, PartialEq,
498            serde::Serialize, serde::Deserialize, Hash,
499                )]
500        pub enum BinaryFunc {
501            $($name($variant),)*
502        }
503
504        impl BinaryFunc {
505            pub fn eval<'a>(
506                &'a self,
507                datums: &[Datum<'a>],
508                temp_storage: &'a RowArena,
509                exprs: &[&'a impl Eval],
510            ) -> Result<Datum<'a>, EvalError> {
511                match self {
512                    $(Self::$name(f) => f.eval(datums, temp_storage, exprs),)*
513                }
514            }
515
516            pub fn output_sql_type(&self, input_types: &[SqlColumnType]) -> SqlColumnType {
517                match self {
518                    $(Self::$name(f) => {
519                        LazyBinaryFunc::output_sql_type(f, input_types)
520                    },)*
521                }
522            }
523
524            pub fn output_type(&self, input_types: &[ReprColumnType]) -> ReprColumnType {
525                match self {
526                    $(Self::$name(f) => LazyBinaryFunc::output_type(f, input_types),)*
527                }
528            }
529
530            pub fn propagates_nulls(&self) -> bool {
531                match self {
532                    $(Self::$name(f) => LazyBinaryFunc::propagates_nulls(f),)*
533                }
534            }
535
536            pub fn introduces_nulls(&self) -> bool {
537                match self {
538                    $(Self::$name(f) => LazyBinaryFunc::introduces_nulls(f),)*
539                }
540            }
541
542            pub fn is_infix_op(&self) -> bool {
543                match self {
544                    $(Self::$name(f) => LazyBinaryFunc::is_infix_op(f),)*
545                }
546            }
547
548            pub fn negate(&self) -> Option<BinaryFunc> {
549                match self {
550                    $(Self::$name(f) => LazyBinaryFunc::negate(f),)*
551                }
552            }
553
554            pub fn could_error(&self) -> bool {
555                match self {
556                    $(Self::$name(f) => LazyBinaryFunc::could_error(f),)*
557                }
558            }
559
560            pub fn is_monotone(&self) -> (bool, bool) {
561                match self {
562                    $(Self::$name(f) => LazyBinaryFunc::is_monotone(f),)*
563                }
564            }
565
566            pub fn is_infinity_monotone(&self) -> bool {
567                match self {
568                    $(Self::$name(f) => LazyBinaryFunc::is_infinity_monotone(f),)*
569                }
570            }
571
572            /// The canonical name of this variant, as declared by its
573            /// [`FuncName`](crate::func::FuncName) impl.
574            pub fn variant_name(&self) -> &'static str {
575                match self {
576                    $(Self::$name(_) => <$variant as crate::func::FuncName>::NAME,)*
577                }
578            }
579
580            /// The `#[sqlfunc]` source of this variant's function, `None` for
581            /// hand-written functions. See [`FuncName::SQLFUNC`](crate::func::FuncName::SQLFUNC).
582            #[cfg(feature = "func-registry")]
583            pub fn sqlfunc_source(&self) -> Option<crate::func::SqlFuncSource> {
584                match self {
585                    $(Self::$name(_) => <$variant as crate::func::FuncName>::SQLFUNC,)*
586                }
587            }
588
589            /// See [`FuncName::sqlfunc_input_types`](crate::func::FuncName::sqlfunc_input_types).
590            #[cfg(feature = "func-registry")]
591            pub fn sqlfunc_input_types(&self) -> Option<Vec<SqlColumnType>> {
592                match self {
593                    $(Self::$name(_) =>
594                        <$variant as crate::func::FuncName>::sqlfunc_input_types(),)*
595                }
596            }
597
598            /// Attempts to construct a `BinaryFunc` from the canonical name of
599            /// one of its variants, as declared by the variant's
600            /// [`FuncName`](crate::func::FuncName) impl (the name of the
601            /// underlying Rust function, e.g. `add_int32`).
602            ///
603            /// Intended for test tooling that needs to name exact function
604            /// variants. Only variants whose inner function deserializes from
605            /// no data (unit functions and functions all of whose parameters
606            /// may default to `None`) are constructible. Returns `None` for
607            /// other variants and for unknown names.
608            pub fn from_variant_name(name: &str) -> Option<Self> {
609                $(
610                    if name == <$variant as crate::func::FuncName>::NAME {
611                        return serde_json::from_value(
612                            serde_json::json!({ stringify!($name): null }),
613                        )
614                        .ok();
615                    }
616                )*
617                None
618            }
619
620            /// The canonical names of all variants, in declaration order.
621            pub fn variant_names() -> impl Iterator<Item = &'static str> {
622                [$(<$variant as crate::func::FuncName>::NAME,)*].into_iter()
623            }
624        }
625
626        impl fmt::Display for BinaryFunc {
627            fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
628                match self {
629                    $(Self::$name(func) => func.fmt(f),)*
630                }
631            }
632        }
633
634        $(
635            impl From<$variant> for crate::BinaryFunc {
636                fn from(variant: $variant) -> Self {
637                    Self::$name(variant)
638                }
639            }
640        )*
641    }
642}