Skip to main content

mz_expr/scalar/
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// Portions of this file are derived from the PostgreSQL project. The original
11// source code is subject to the terms of the PostgreSQL license, a copy of
12// which can be found in the LICENSE file at the root of this repository.
13
14use std::borrow::Cow;
15use std::cmp::Ordering;
16use std::convert::{TryFrom, TryInto};
17use std::str::FromStr;
18use std::{iter, str};
19
20use ::encoding::DecoderTrap;
21use ::encoding::label::encoding_from_whatwg_label;
22use aws_lc_rs::constant_time::verify_slices_are_equal;
23use aws_lc_rs::digest;
24use chrono::{DateTime, Duration, NaiveDate, NaiveDateTime, TimeZone, Timelike, Utc};
25use chrono_tz::{OffsetComponents, OffsetName, Tz};
26use dec::OrderedDecimal;
27use itertools::Itertools;
28use md5::{Digest, Md5};
29use mz_expr_derive::sqlfunc;
30use mz_ore::cast::{self, CastFrom};
31use mz_ore::fmt::FormatBuffer;
32use mz_ore::lex::LexBuf;
33use mz_ore::option::OptionExt;
34use mz_pgrepr::Type;
35use mz_pgtz::timezone::{Timezone, TimezoneSpec};
36use mz_repr::adt::array::{Array, ArrayDimension};
37use mz_repr::adt::date::Date;
38use mz_repr::adt::interval::{Interval, RoundBehavior};
39use mz_repr::adt::jsonb::JsonbRef;
40use mz_repr::adt::mz_acl_item::{AclMode, MzAclItem};
41use mz_repr::adt::numeric::{self, Numeric};
42use mz_repr::adt::range::Range;
43use mz_repr::adt::regex::Regex;
44use mz_repr::adt::timestamp::{CheckedTimestamp, TimestampLike};
45use mz_repr::{
46    ArrayRustType, Datum, DatumList, DatumMap, ExcludeNull, FromDatum, InputDatumType, Row,
47    RowArena, SqlScalarType, strconv,
48};
49use mz_sql_parser::ast::display::{AstDisplay, FormatMode};
50use mz_sql_pretty::{PrettyConfig, pretty_str};
51use num::traits::CheckedNeg;
52
53use crate::scalar::func::format::DateTimeFormat;
54use crate::{EvalError, like_pattern};
55
56#[macro_use]
57mod macros;
58mod binary;
59mod encoding;
60pub(crate) mod format;
61pub(crate) mod impls;
62#[cfg(feature = "func-registry")]
63pub mod registry;
64mod unary;
65mod unmaterializable;
66pub mod variadic;
67
68pub use binary::BinaryFunc;
69pub use impls::*;
70pub use unary::{EagerUnaryFunc, LazyUnaryFunc, UnaryFunc};
71pub use unmaterializable::UnmaterializableFunc;
72pub use variadic::VariadicFunc;
73
74/// The canonical name of a scalar function, and its source when it is
75/// generated by `#[sqlfunc]`.
76///
77/// For functions generated by `#[sqlfunc]` the name is that of the underlying
78/// Rust function. Hand-written functions declare it explicitly. Test tooling
79/// refers to exact function variants by these names, via
80/// [`UnaryFunc::from_variant_name`] and its binary and variadic counterparts.
81pub trait FuncName {
82    const NAME: &'static str;
83
84    /// The `#[sqlfunc]` declaration this function was generated from, or
85    /// `None` for hand-written functions. Consumed by [`registry`].
86    #[cfg(feature = "func-registry")]
87    const SQLFUNC: Option<SqlFuncSource> = None;
88
89    /// The column types a `#[sqlfunc]` function naturally consumes, one per
90    /// argument, when every parameter type maps to one. `None` for
91    /// hand-written functions and for parameters typed as a bare `Datum`.
92    /// Consumed by [`registry`] to probe the output type.
93    #[cfg(feature = "func-registry")]
94    fn sqlfunc_input_types() -> Option<Vec<mz_repr::SqlColumnType>> {
95        None
96    }
97}
98
99/// The source of a `#[sqlfunc]`-generated function, as captured by the macro.
100///
101/// Comments and formatting are invisible to both fields, because they are
102/// derived from token streams. Helpers the body calls are not covered, so an
103/// unchanged fingerprint does not prove unchanged behavior.
104#[cfg(feature = "func-registry")]
105#[derive(Clone, Copy, Debug, PartialEq, Eq)]
106pub struct SqlFuncSource {
107    /// The attribute arguments and the function signature, rendered as
108    /// compact source text: `#[sqlfunc(<args>)] fn <name>(<params>) -> <ret>`.
109    pub decl: &'static str,
110    /// The parameter and return types alone, as written: `fn(<types>) -> <ret>`.
111    pub signature: &'static str,
112    /// FNV-1a fingerprint of the function body.
113    pub body_fingerprint: u64,
114}
115
116/// Declares the canonical [`FuncName`] of hand-written function structs.
117/// Functions generated by `#[sqlfunc]` get theirs from that macro instead.
118/// Structs generic over an expression type declare a leading `<E>`.
119macro_rules! func_name {
120    () => {};
121    (<$param:ident> $ty:ty => $name:literal, $($rest:tt)*) => {
122        impl<$param> FuncName for $ty {
123            const NAME: &'static str = $name;
124        }
125        func_name!($($rest)*);
126    };
127    ($ty:ty => $name:literal, $($rest:tt)*) => {
128        impl FuncName for $ty {
129            const NAME: &'static str = $name;
130        }
131        func_name!($($rest)*);
132    };
133}
134
135func_name! {
136    AdjustNumericScale => "adjust_numeric_scale",
137    AdjustTimestampPrecision => "adjust_timestamp_precision",
138    AdjustTimestampTzPrecision => "adjust_timestamp_tz_precision",
139    CaseLiteral => "case_literal",
140    <E> CastArrayToArray<E> => "cast_array_to_array",
141    <E> CastArrayToJsonb<E> => "cast_array_to_jsonb",
142    CastArrayToString => "cast_array_to_string",
143    CastDateToTimestamp => "cast_date_to_timestamp",
144    CastDateToTimestampTz => "cast_date_to_timestamp_tz",
145    CastFloat32ToNumeric => "cast_float32_to_numeric",
146    CastFloat64ToNumeric => "cast_float64_to_numeric",
147    CastInt16ToNumeric => "cast_int16_to_numeric",
148    CastInt32ToNumeric => "cast_int32_to_numeric",
149    CastInt64ToNumeric => "cast_int64_to_numeric",
150    CastJsonbToNumeric => "cast_jsonb_to_numeric",
151    <E> CastList1ToList2<E> => "cast_list1_to_list2",
152    <E> CastListToJsonb<E> => "cast_list_to_jsonb",
153    CastListToString => "cast_list_to_string",
154    CastMapToString => "cast_map_to_string",
155    CastRangeToString => "cast_range_to_string",
156    <E> CastRecord1ToRecord2<E> => "cast_record1_to_record2",
157    CastRecordToString => "cast_record_to_string",
158    <E> CastStringToArray<E> => "cast_string_to_array",
159    CastStringToChar => "cast_string_to_char",
160    CastStringToInt2Vector => "cast_string_to_int2_vector",
161    <E> CastStringToList<E> => "cast_string_to_list",
162    <E> CastStringToMap<E> => "cast_string_to_map",
163    CastStringToNumeric => "cast_string_to_numeric",
164    <E> CastStringToRange<E> => "cast_string_to_range",
165    CastStringToTimestamp => "cast_string_to_timestamp",
166    CastStringToTimestampTz => "cast_string_to_timestamp_tz",
167    CastStringToVarChar => "cast_string_to_var_char",
168    CastTimestampToTimestampTz => "cast_timestamp_to_timestamp_tz",
169    CastTimestampTzToTimestamp => "cast_timestamp_tz_to_timestamp",
170    CastUint16ToNumeric => "cast_uint16_to_numeric",
171    CastUint32ToNumeric => "cast_uint32_to_numeric",
172    CastUint64ToNumeric => "cast_uint64_to_numeric",
173    DatePartInterval => "date_part_interval",
174    DatePartTime => "date_part_time",
175    DatePartTimestamp => "date_part_timestamp",
176    DatePartTimestampTz => "date_part_timestamp_tz",
177    DateTruncTimestamp => "date_trunc_timestamp",
178    DateTruncTimestampTz => "date_trunc_timestamp_tz",
179    variadic::ErrorIfNull => "error_if_null",
180    ExtractDate => "extract_date",
181    ExtractInterval => "extract_interval",
182    ExtractTime => "extract_time",
183    ExtractTimestamp => "extract_timestamp",
184    ExtractTimestampTz => "extract_timestamp_tz",
185    IsLikeMatch => "is_like_match",
186    IsRegexpMatch => "is_regexp_match",
187    ListLengthMax => "list_length_max",
188    MapBuildFromRecordList => "map_build_from_record_list",
189    PadChar => "pad_char",
190    variadic::RangeCreate => "range_create",
191    RecordGet => "record_get",
192    RegexpMatch => "regexp_match",
193    RegexpReplace => "regexp_replace",
194    RegexpSplitToArray => "regexp_split_to_array",
195    TimezoneTime => "timezone_time",
196    TimezoneTimestamp => "timezone_timestamp",
197    TimezoneTimestampTz => "timezone_timestamp_tz",
198    ToCharTimestamp => "to_char_timestamp",
199    ToCharTimestampTz => "to_char_timestamp_tz",
200    variadic::And => "and",
201    variadic::Coalesce => "coalesce",
202    variadic::Greatest => "greatest",
203    variadic::Least => "least",
204    variadic::Or => "or",
205}
206
207/// The maximum size of the result strings of certain string functions, such as `repeat` and `lpad`.
208/// Chosen to be the smallest number to keep our tests passing without changing. 100MiB is probably
209/// higher than what we want, but it's better than no limit.
210///
211/// Note: This number appears in our user-facing documentation in the function reference for every
212/// function where it applies.
213pub const MAX_STRING_FUNC_RESULT_BYTES: usize = 1024 * 1024 * 100;
214
215/// The largest result a string function may build into `temp_storage`.
216///
217/// [`MAX_STRING_FUNC_RESULT_BYTES`] unless the arena carries a tighter budget, which is how an
218/// expression evaluated in `environmentd` on behalf of a request (a webhook `CHECK`) is held to a
219/// size proportionate to that request rather than to the constant, which is sized for a cluster.
220///
221/// A function that can predict its result size must consult this *before* building the result: the
222/// arena's own budget is only observable after the bytes exist, which for an amplifying function is
223/// exactly too late.
224pub fn max_string_func_result_bytes(temp_storage: &RowArena) -> usize {
225    std::cmp::min(
226        MAX_STRING_FUNC_RESULT_BYTES,
227        temp_storage.budget_remaining(),
228    )
229}
230
231/// Refuses a collection `temp_storage`'s budget cannot afford, before it is packed.
232///
233/// The collection builders (`ARRAY[..]`, `LIST[..]`, `ROW(..)`, `MAP[..]`, `jsonb_build_*`) pack the
234/// datums they are handed straight into `temp_storage`, so the result is as large as those datums
235/// times however many times the expression names each one. Unlike the string functions bounded by
236/// [`max_string_func_result_bytes`], `ARRAY[body, body, ..]` has no ceiling of its own.
237///
238/// Like that ceiling, this must be consulted *before* the result is built, since the arena's budget
239/// is only observable once the bytes exist. [`mz_repr::datum_size`] is the size a datum occupies
240/// once packed, so summing it bounds the result without allocating anything. Without a budget
241/// `budget_remaining` is `usize::MAX` and nothing is refused.
242pub fn check_datums_fit_budget<'a>(
243    datums: impl IntoIterator<Item = Datum<'a>>,
244    temp_storage: &RowArena,
245) -> Result<(), EvalError> {
246    let need: usize = datums
247        .into_iter()
248        .map(|d| mz_repr::datum_size(&d))
249        .fold(0, usize::saturating_add);
250    if need > temp_storage.budget_remaining() {
251        return Err(EvalError::TempStorageBudgetExceeded);
252    }
253    Ok(())
254}
255
256/// Refuses an input-scaled transient the arena's budget cannot afford, before it is built.
257///
258/// A function that gathers `n_elems` of `elem_size` bytes into its own `Vec` before packing them
259/// holds that allocation on the stack, where the arena never sees it. The evaluator's post-call
260/// check counts arena bytes only, so a transient wider than the result it becomes can slip past a
261/// budget the packed result fits under. A `Vec<&str>` of split chunks is one: a 16-byte fat pointer
262/// per chunk against the ~2 bytes an empty chunk packs to. Pair this with
263/// [`check_datums_fit_budget`], which bounds the packed result.
264///
265/// `n_elems` is a closure so an unbudgeted arena, which is every arena in a dataflow, never pays for
266/// a count that can cost a pass over the input.
267///
268/// NOTE: the reformatters (`jsonb_pretty`, `pretty_sql`, `redact_sql`) are the known exception.
269/// Sizing their output needs the same walk that produces it, so only the post-call check bounds
270/// them.
271pub fn check_build_fits_budget(
272    n_elems: impl FnOnce() -> usize,
273    elem_size: usize,
274    temp_storage: &RowArena,
275) -> Result<(), EvalError> {
276    let budget_remaining = temp_storage.budget_remaining();
277    // `usize::MAX` is the unbudgeted sentinel. Return before the count so an unbudgeted arena never
278    // pays for it.
279    if budget_remaining == usize::MAX {
280        return Ok(());
281    }
282    if n_elems().saturating_mul(elem_size) > budget_remaining {
283        return Err(EvalError::TempStorageBudgetExceeded);
284    }
285    Ok(())
286}
287
288pub fn jsonb_stringify<'a>(a: Datum<'a>, temp_storage: &'a RowArena) -> Option<&'a str> {
289    match a {
290        Datum::JsonNull => None,
291        Datum::String(s) => Some(s),
292        _ => {
293            let s = cast_jsonb_to_string(JsonbRef::from_datum(a));
294            Some(temp_storage.push_string(s))
295        }
296    }
297}
298
299#[sqlfunc(
300    is_monotone = "(true, true)",
301    is_infix_op = true,
302    sqlname = "+",
303    propagates_nulls = true
304)]
305fn add_int16(a: i16, b: i16) -> Result<i16, EvalError> {
306    a.checked_add(b).ok_or(EvalError::NumericFieldOverflow)
307}
308
309#[sqlfunc(
310    is_monotone = "(true, true)",
311    is_infix_op = true,
312    sqlname = "+",
313    propagates_nulls = true
314)]
315fn add_int32(a: i32, b: i32) -> Result<i32, EvalError> {
316    a.checked_add(b).ok_or(EvalError::NumericFieldOverflow)
317}
318
319#[sqlfunc(
320    is_monotone = "(true, true)",
321    is_infix_op = true,
322    sqlname = "+",
323    propagates_nulls = true
324)]
325fn add_int64(a: i64, b: i64) -> Result<i64, EvalError> {
326    a.checked_add(b).ok_or(EvalError::NumericFieldOverflow)
327}
328
329#[sqlfunc(
330    is_monotone = "(true, true)",
331    is_infix_op = true,
332    sqlname = "+",
333    propagates_nulls = true
334)]
335fn add_uint16(a: u16, b: u16) -> Result<u16, EvalError> {
336    a.checked_add(b)
337        .ok_or_else(|| EvalError::UInt16OutOfRange(format!("{a} + {b}").into()))
338}
339
340#[sqlfunc(
341    is_monotone = "(true, true)",
342    is_infix_op = true,
343    sqlname = "+",
344    propagates_nulls = true
345)]
346fn add_uint32(a: u32, b: u32) -> Result<u32, EvalError> {
347    a.checked_add(b)
348        .ok_or_else(|| EvalError::UInt32OutOfRange(format!("{a} + {b}").into()))
349}
350
351#[sqlfunc(
352    is_monotone = "(true, true)",
353    is_infix_op = true,
354    sqlname = "+",
355    propagates_nulls = true
356)]
357fn add_uint64(a: u64, b: u64) -> Result<u64, EvalError> {
358    a.checked_add(b)
359        .ok_or_else(|| EvalError::UInt64OutOfRange(format!("{a} + {b}").into()))
360}
361
362#[sqlfunc(
363    is_monotone = "(true, true)",
364    is_infix_op = true,
365    sqlname = "+",
366    propagates_nulls = true
367)]
368fn add_float32(a: f32, b: f32) -> Result<f32, EvalError> {
369    let sum = a + b;
370    if sum.is_infinite() && !a.is_infinite() && !b.is_infinite() {
371        Err(EvalError::FloatOverflow)
372    } else {
373        Ok(sum)
374    }
375}
376
377#[sqlfunc(
378    is_monotone = "(true, true)",
379    is_infix_op = true,
380    sqlname = "+",
381    propagates_nulls = true
382)]
383fn add_float64(a: f64, b: f64) -> Result<f64, EvalError> {
384    let sum = a + b;
385    if sum.is_infinite() && !a.is_infinite() && !b.is_infinite() {
386        Err(EvalError::FloatOverflow)
387    } else {
388        Ok(sum)
389    }
390}
391
392// `Interval` is lex-ordered (months, days, micros), but adding an interval to a
393// timestamp adds *calendar* months (with day-clamping) which does not respect
394// that ordering: e.g. `i1 = {0 months, 31 days}` is lex-less than
395// `i2 = {1 month, 0 days}`, but `2024-01-31 + i1 = 2024-03-02` is greater than
396// `2024-01-31 + i2 = 2024-02-29`. Day-clamping plus preserved sub-day time also
397// breaks monotonicity in the first argument near month boundaries.
398#[sqlfunc(is_monotone = "(false, false)", is_infix_op = true, sqlname = "+")]
399fn add_timestamp_interval(
400    a: CheckedTimestamp<NaiveDateTime>,
401    b: Interval,
402) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
403    add_timestamplike_interval(a, b)
404}
405
406#[sqlfunc(is_monotone = "(false, false)", is_infix_op = true, sqlname = "+")]
407fn add_timestamp_tz_interval(
408    a: CheckedTimestamp<DateTime<Utc>>,
409    b: Interval,
410) -> Result<CheckedTimestamp<DateTime<Utc>>, EvalError> {
411    add_timestamplike_interval(a, b)
412}
413
414fn add_timestamplike_interval<T>(
415    a: CheckedTimestamp<T>,
416    b: Interval,
417) -> Result<CheckedTimestamp<T>, EvalError>
418where
419    T: TimestampLike,
420{
421    let dt = a.date_time();
422    let dt = add_timestamp_months(&dt, b.months)?;
423    let dt = dt
424        .checked_add_signed(b.duration_as_chrono())
425        .ok_or(EvalError::TimestampOutOfRange)?;
426    Ok(CheckedTimestamp::from_timestamplike(T::from_date_time(dt))?)
427}
428
429// See `add_timestamp_interval` for why this is not monotone.
430#[sqlfunc(is_monotone = "(false, false)", is_infix_op = true, sqlname = "-")]
431fn sub_timestamp_interval(
432    a: CheckedTimestamp<NaiveDateTime>,
433    b: Interval,
434) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
435    sub_timestamplike_interval(a, b)
436}
437
438#[sqlfunc(is_monotone = "(false, false)", is_infix_op = true, sqlname = "-")]
439fn sub_timestamp_tz_interval(
440    a: CheckedTimestamp<DateTime<Utc>>,
441    b: Interval,
442) -> Result<CheckedTimestamp<DateTime<Utc>>, EvalError> {
443    sub_timestamplike_interval(a, b)
444}
445
446fn sub_timestamplike_interval<T>(
447    a: CheckedTimestamp<T>,
448    b: Interval,
449) -> Result<CheckedTimestamp<T>, EvalError>
450where
451    T: TimestampLike,
452{
453    neg_interval_inner(b).and_then(|i| add_timestamplike_interval(a, i))
454}
455
456#[sqlfunc(is_monotone = "(true, true)", is_infix_op = true, sqlname = "+")]
457fn add_date_time(
458    date: Date,
459    time: chrono::NaiveTime,
460) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
461    // A leap-second TIME (nanos >= 1e9) rolls over into the next minute,
462    // matching what parsing the equivalent timestamp literal produces. The
463    // leap representation must not enter a timestamp: it sorts before the
464    // next second while epoch-style conversions count it at or past it,
465    // breaking the monotonicity contracts filter pushdown relies on.
466    let (extra_sec, nanos) = match time.nanosecond().checked_sub(1_000_000_000) {
467        Some(nanos) => (1, nanos),
468        None => (0, time.nanosecond()),
469    };
470    let dt = NaiveDate::from(date)
471        .and_hms_nano_opt(time.hour(), time.minute(), time.second(), nanos)
472        .unwrap()
473        .checked_add_signed(chrono::Duration::try_seconds(extra_sec).unwrap())
474        .ok_or(EvalError::TimestampOutOfRange)?;
475    Ok(CheckedTimestamp::from_timestamplike(dt)?)
476}
477
478// Monotone in `date` (dates have no sub-day component, so day-clamping at month
479// boundaries only causes results to collapse, never to reverse), but not in
480// `interval`: e.g. `{0 months, 31 days}` is lex-less than `{1 month, 0 days}`,
481// but adding the former to `2024-01-31` gives `2024-03-02` while the latter
482// gives `2024-02-29`.
483#[sqlfunc(is_monotone = "(true, false)", is_infix_op = true, sqlname = "+")]
484fn add_date_interval(
485    date: Date,
486    interval: Interval,
487) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
488    let dt = NaiveDate::from(date).and_hms_opt(0, 0, 0).unwrap();
489    let dt = add_timestamp_months(&dt, interval.months)?;
490    let dt = dt
491        .checked_add_signed(interval.duration_as_chrono())
492        .ok_or(EvalError::TimestampOutOfRange)?;
493    Ok(CheckedTimestamp::from_timestamplike(dt)?)
494}
495
496#[sqlfunc(
497    // <time> + <interval> wraps!
498    is_monotone = "(false, false)",
499    is_infix_op = true,
500    sqlname = "+",
501    propagates_nulls = true
502)]
503fn add_time_interval(time: chrono::NaiveTime, interval: Interval) -> chrono::NaiveTime {
504    let (t, _) = time.overflowing_add_signed(interval.duration_as_chrono());
505    t
506}
507
508#[sqlfunc(
509    is_monotone = "(true, false)",
510    output_type = "Numeric",
511    sqlname = "round",
512    propagates_nulls = true
513)]
514fn round_numeric_binary(a: OrderedDecimal<Numeric>, mut b: i32) -> Result<Numeric, EvalError> {
515    let mut a = a.0;
516    let mut cx = numeric::cx_datum();
517    let a_scale = numeric::get_scale(&a);
518    if a.is_finite() && i64::from(b) > i64::from(a_scale) {
519        // Rounding at or past `a`'s scale cannot change the value: right-pad
520        // with zeroes via rescale. The rounding path below shifts left by `b`
521        // first and overflows for large `b`.
522        //
523        // NOTE: Equal values can reach here with different scales, since
524        // `Row` encoding folds trailing zeroes into the exponent. The result,
525        // value or error, must not depend on which representation arrives: the
526        // abstract interpreter reads its datums back out of a `Row`, and if it
527        // calls infallible what the evaluator fails on, persist filter
528        // pushdown discards parts it has to keep.
529        //
530        // `Infinity` and `NaN` report a scale of zero, but `rescale` on an
531        // infinity yields `NaN` via invalid_operation, not the overflow
532        // checked below, so the specials take the rounding path, which
533        // propagates them unchanged, as PostgreSQL does.
534
535        // Ensure rescale doesn't exceed max precision by putting a ceiling on
536        // b equal to the maximum remaining scale the value can support.
537        let max_remaining_scale = u32::from(numeric::NUMERIC_DATUM_MAX_PRECISION)
538            - (numeric::get_precision(&a) - a_scale);
539        b = match i32::try_from(max_remaining_scale) {
540            Ok(max_remaining_scale) => std::cmp::min(b, max_remaining_scale),
541            Err(_) => b,
542        };
543        cx.rescale(&mut a, &numeric::Numeric::from(-b));
544    } else {
545        // To avoid invalid operations, clamp b to be within 1 more than the
546        // precision limit.
547        const MAX_P_LIMIT: i32 = 1 + cast::u8_to_i32(numeric::NUMERIC_DATUM_MAX_PRECISION);
548        b = std::cmp::min(MAX_P_LIMIT, b);
549        b = std::cmp::max(-MAX_P_LIMIT, b);
550        let mut b = numeric::Numeric::from(b);
551        // Shift by 10^b; this put digit to round to in the one's place.
552        cx.scaleb(&mut a, &b);
553        cx.round(&mut a);
554        // Negate exponent for shift back
555        cx.neg(&mut b);
556        cx.scaleb(&mut a, &b);
557    }
558
559    if cx.status().overflow() {
560        Err(EvalError::FloatOverflow)
561    } else if a.is_zero() {
562        // simpler than handling cases where exponent has gotten set to some
563        // value greater than the max precision, but all significant digits
564        // were rounded away.
565        Ok(numeric::Numeric::zero())
566    } else {
567        numeric::munge_numeric(&mut a).unwrap();
568        Ok(a)
569    }
570}
571
572#[sqlfunc(sqlname = "convert_from", propagates_nulls = true)]
573fn convert_from<'a>(a: &'a [u8], b: &str) -> Result<&'a str, EvalError> {
574    // Convert PostgreSQL-style encoding names[1] to WHATWG-style encoding names[2],
575    // which the encoding library uses[3].
576    // [1]: https://www.postgresql.org/docs/9.5/multibyte.html
577    // [2]: https://encoding.spec.whatwg.org/
578    // [3]: https://github.com/lifthrasiir/rust-encoding/blob/4e79c35ab6a351881a86dbff565c4db0085cc113/src/label.rs
579    let encoding_name = b.to_lowercase().replace('_', "-").into_boxed_str();
580
581    // Supporting other encodings is tracked by database-issues#797.
582    if encoding_from_whatwg_label(&encoding_name).map(|e| e.name()) != Some("utf-8") {
583        return Err(EvalError::InvalidEncodingName(encoding_name));
584    }
585
586    match str::from_utf8(a) {
587        // Match PostgreSQL, which rejects NUL bytes because text values must
588        // never contain them.
589        Ok(from) if from.contains('\0') => Err(EvalError::InvalidByteSequence {
590            byte_sequence: "0x00".into(),
591            encoding_name,
592        }),
593        Ok(from) => Ok(from),
594        Err(e) => Err(EvalError::InvalidByteSequence {
595            byte_sequence: e.to_string().into(),
596            encoding_name,
597        }),
598    }
599}
600
601#[sqlfunc]
602fn encode(bytes: &[u8], format: &str) -> Result<String, EvalError> {
603    let format = encoding::lookup_format(format)?;
604    Ok(format.encode(bytes))
605}
606
607#[sqlfunc]
608fn decode(string: &str, format: &str, temp_storage: &RowArena) -> Result<Vec<u8>, EvalError> {
609    let format = encoding::lookup_format(format)?;
610    let out = format.decode(string)?;
611    if out.len() > max_string_func_result_bytes(temp_storage) {
612        Err(EvalError::LengthTooLarge)
613    } else {
614        Ok(out)
615    }
616}
617
618#[sqlfunc(sqlname = "length", propagates_nulls = true)]
619fn encoded_bytes_char_length(a: &[u8], b: &str) -> Result<i32, EvalError> {
620    // Convert PostgreSQL-style encoding names[1] to WHATWG-style encoding names[2],
621    // which the encoding library uses[3].
622    // [1]: https://www.postgresql.org/docs/9.5/multibyte.html
623    // [2]: https://encoding.spec.whatwg.org/
624    // [3]: https://github.com/lifthrasiir/rust-encoding/blob/4e79c35ab6a351881a86dbff565c4db0085cc113/src/label.rs
625    let encoding_name = b.to_lowercase().replace('_', "-").into_boxed_str();
626
627    let enc = match encoding_from_whatwg_label(&encoding_name) {
628        Some(enc) => enc,
629        None => return Err(EvalError::InvalidEncodingName(encoding_name)),
630    };
631
632    let decoded_string = match enc.decode(a, DecoderTrap::Strict) {
633        Ok(s) => s,
634        Err(e) => {
635            return Err(EvalError::InvalidByteSequence {
636                byte_sequence: e.into(),
637                encoding_name,
638            });
639        }
640    };
641
642    let count = decoded_string.chars().count();
643    i32::try_from(count).map_err(|_| EvalError::Int32OutOfRange(count.to_string().into()))
644}
645
646// TODO(benesch): remove potentially dangerous usage of `as`.
647#[allow(clippy::as_conversions)]
648pub fn add_timestamp_months<T: TimestampLike>(
649    dt: &T,
650    mut months: i32,
651) -> Result<CheckedTimestamp<T>, EvalError> {
652    if months == 0 {
653        return Ok(CheckedTimestamp::from_timestamplike(dt.clone())?);
654    }
655
656    let (mut year, mut month, mut day) = (dt.year(), dt.month0() as i32, dt.day());
657    let years = months / 12;
658    year = year
659        .checked_add(years)
660        .ok_or(EvalError::TimestampOutOfRange)?;
661
662    months %= 12;
663    // positive modulus is easier to reason about
664    if months < 0 {
665        year -= 1;
666        months += 12;
667    }
668    year += (month + months) / 12;
669    month = (month + months) % 12;
670    // account for dt.month0
671    month += 1;
672
673    // handle going from January 31st to February by saturation
674    let mut new_d = chrono::NaiveDate::from_ymd_opt(year, month as u32, day);
675    while new_d.is_none() {
676        // If we have decremented day past 28 and are still receiving `None`,
677        // then we have generally overflowed `NaiveDate`.
678        if day < 28 {
679            return Err(EvalError::TimestampOutOfRange);
680        }
681        day -= 1;
682        new_d = chrono::NaiveDate::from_ymd_opt(year, month as u32, day);
683    }
684    let new_d = new_d.unwrap();
685
686    // Neither postgres nor mysql support leap seconds, so this should be safe.
687    //
688    // Both my testing and https://dba.stackexchange.com/a/105829 support the
689    // idea that we should ignore leap seconds
690    let new_dt = new_d
691        .and_hms_nano_opt(dt.hour(), dt.minute(), dt.second(), dt.nanosecond())
692        .unwrap();
693    let new_dt = T::from_date_time(new_dt);
694    Ok(CheckedTimestamp::from_timestamplike(new_dt)?)
695}
696
697#[sqlfunc(
698    is_monotone = "(true, true)",
699    is_infix_op = true,
700    sqlname = "+",
701    propagates_nulls = true
702)]
703fn add_numeric(
704    a: OrderedDecimal<Numeric>,
705    b: OrderedDecimal<Numeric>,
706) -> Result<Numeric, EvalError> {
707    let mut cx = numeric::cx_datum();
708    let mut a = a.0;
709    cx.add(&mut a, &b.0);
710    if cx.status().overflow() {
711        Err(EvalError::FloatOverflow)
712    } else {
713        Ok(a)
714    }
715}
716
717#[sqlfunc(
718    is_monotone = "(true, true)",
719    is_infix_op = true,
720    sqlname = "+",
721    propagates_nulls = true
722)]
723fn add_interval(a: Interval, b: Interval) -> Result<Interval, EvalError> {
724    a.checked_add(&b)
725        .ok_or_else(|| EvalError::IntervalOutOfRange(format!("{a} + {b}").into()))
726}
727
728#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
729fn bit_and_int16(a: i16, b: i16) -> i16 {
730    a & b
731}
732
733#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
734fn bit_and_int32(a: i32, b: i32) -> i32 {
735    a & b
736}
737
738#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
739fn bit_and_int64(a: i64, b: i64) -> i64 {
740    a & b
741}
742
743#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
744fn bit_and_uint16(a: u16, b: u16) -> u16 {
745    a & b
746}
747
748#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
749fn bit_and_uint32(a: u32, b: u32) -> u32 {
750    a & b
751}
752
753#[sqlfunc(is_infix_op = true, sqlname = "&", propagates_nulls = true)]
754fn bit_and_uint64(a: u64, b: u64) -> u64 {
755    a & b
756}
757
758#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
759fn bit_or_int16(a: i16, b: i16) -> i16 {
760    a | b
761}
762
763#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
764fn bit_or_int32(a: i32, b: i32) -> i32 {
765    a | b
766}
767
768#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
769fn bit_or_int64(a: i64, b: i64) -> i64 {
770    a | b
771}
772
773#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
774fn bit_or_uint16(a: u16, b: u16) -> u16 {
775    a | b
776}
777
778#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
779fn bit_or_uint32(a: u32, b: u32) -> u32 {
780    a | b
781}
782
783#[sqlfunc(is_infix_op = true, sqlname = "|", propagates_nulls = true)]
784fn bit_or_uint64(a: u64, b: u64) -> u64 {
785    a | b
786}
787
788#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
789fn bit_xor_int16(a: i16, b: i16) -> i16 {
790    a ^ b
791}
792
793#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
794fn bit_xor_int32(a: i32, b: i32) -> i32 {
795    a ^ b
796}
797
798#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
799fn bit_xor_int64(a: i64, b: i64) -> i64 {
800    a ^ b
801}
802
803#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
804fn bit_xor_uint16(a: u16, b: u16) -> u16 {
805    a ^ b
806}
807
808#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
809fn bit_xor_uint32(a: u32, b: u32) -> u32 {
810    a ^ b
811}
812
813#[sqlfunc(is_infix_op = true, sqlname = "#", propagates_nulls = true)]
814fn bit_xor_uint64(a: u64, b: u64) -> u64 {
815    a ^ b
816}
817
818#[sqlfunc(is_infix_op = true, sqlname = "<<", propagates_nulls = true)]
819// TODO(benesch): remove potentially dangerous usage of `as`.
820#[allow(clippy::as_conversions)]
821fn bit_shift_left_int16(a: i16, b: i32) -> i16 {
822    // widen to i32 and then cast back to i16 in order emulate the C promotion rules used in by Postgres
823    // when the rhs in the 16-31 range, e.g. (1 << 17 should evaluate to 0)
824    // see https://github.com/postgres/postgres/blob/REL_14_STABLE/src/backend/utils/adt/int.c#L1460-L1476
825    let lhs: i32 = a as i32;
826    let rhs: u32 = b as u32;
827    lhs.wrapping_shl(rhs) as i16
828}
829
830#[sqlfunc(is_infix_op = true, sqlname = "<<", propagates_nulls = true)]
831// TODO(benesch): remove potentially dangerous usage of `as`.
832#[allow(clippy::as_conversions)]
833fn bit_shift_left_int32(lhs: i32, rhs: i32) -> i32 {
834    let rhs = rhs as u32;
835    lhs.wrapping_shl(rhs)
836}
837
838#[sqlfunc(is_infix_op = true, sqlname = "<<", propagates_nulls = true)]
839// TODO(benesch): remove potentially dangerous usage of `as`.
840#[allow(clippy::as_conversions)]
841fn bit_shift_left_int64(lhs: i64, rhs: i32) -> i64 {
842    let rhs = rhs as u32;
843    lhs.wrapping_shl(rhs)
844}
845
846#[sqlfunc(is_infix_op = true, sqlname = "<<", propagates_nulls = true)]
847// TODO(benesch): remove potentially dangerous usage of `as`.
848#[allow(clippy::as_conversions)]
849fn bit_shift_left_uint16(a: u16, b: u32) -> u16 {
850    // widen to u32 and then cast back to u16 in order emulate the C promotion rules used in by Postgres
851    // when the rhs in the 16-31 range, e.g. (1 << 17 should evaluate to 0)
852    // see https://github.com/postgres/postgres/blob/REL_14_STABLE/src/backend/utils/adt/int.c#L1460-L1476
853    let lhs: u32 = a as u32;
854    let rhs: u32 = b;
855    lhs.wrapping_shl(rhs) as u16
856}
857
858#[sqlfunc(is_infix_op = true, sqlname = "<<", propagates_nulls = true)]
859fn bit_shift_left_uint32(a: u32, b: u32) -> u32 {
860    let lhs = a;
861    let rhs = b;
862    lhs.wrapping_shl(rhs)
863}
864
865#[sqlfunc(
866    output_type = "u64",
867    is_infix_op = true,
868    sqlname = "<<",
869    propagates_nulls = true
870)]
871fn bit_shift_left_uint64(lhs: u64, rhs: u32) -> u64 {
872    lhs.wrapping_shl(rhs)
873}
874
875#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
876// TODO(benesch): remove potentially dangerous usage of `as`.
877#[allow(clippy::as_conversions)]
878fn bit_shift_right_int16(lhs: i16, rhs: i32) -> i16 {
879    // widen to i32 and then cast back to i16 in order emulate the C promotion rules used in by Postgres
880    // when the rhs in the 16-31 range, e.g. (-32767 >> 17 should evaluate to -1)
881    // see https://github.com/postgres/postgres/blob/REL_14_STABLE/src/backend/utils/adt/int.c#L1460-L1476
882    let lhs = lhs as i32;
883    let rhs = rhs as u32;
884    lhs.wrapping_shr(rhs) as i16
885}
886
887#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
888// TODO(benesch): remove potentially dangerous usage of `as`.
889#[allow(clippy::as_conversions)]
890fn bit_shift_right_int32(lhs: i32, rhs: i32) -> i32 {
891    lhs.wrapping_shr(rhs as u32)
892}
893
894#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
895// TODO(benesch): remove potentially dangerous usage of `as`.
896#[allow(clippy::as_conversions)]
897fn bit_shift_right_int64(lhs: i64, rhs: i32) -> i64 {
898    lhs.wrapping_shr(rhs as u32)
899}
900
901#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
902// TODO(benesch): remove potentially dangerous usage of `as`.
903#[allow(clippy::as_conversions)]
904fn bit_shift_right_uint16(lhs: u16, rhs: u32) -> u16 {
905    // widen to u32 and then cast back to u16 in order emulate the C promotion rules used in by Postgres
906    // when the rhs in the 16-31 range, e.g. (-32767 >> 17 should evaluate to -1)
907    // see https://github.com/postgres/postgres/blob/REL_14_STABLE/src/backend/utils/adt/int.c#L1460-L1476
908    let lhs = lhs as u32;
909    lhs.wrapping_shr(rhs) as u16
910}
911
912#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
913fn bit_shift_right_uint32(lhs: u32, rhs: u32) -> u32 {
914    lhs.wrapping_shr(rhs)
915}
916
917#[sqlfunc(is_infix_op = true, sqlname = ">>", propagates_nulls = true)]
918fn bit_shift_right_uint64(lhs: u64, rhs: u32) -> u64 {
919    lhs.wrapping_shr(rhs)
920}
921
922#[sqlfunc(
923    is_monotone = "(true, true)",
924    is_infix_op = true,
925    sqlname = "-",
926    propagates_nulls = true
927)]
928fn sub_int16(a: i16, b: i16) -> Result<i16, EvalError> {
929    a.checked_sub(b).ok_or(EvalError::NumericFieldOverflow)
930}
931
932#[sqlfunc(
933    is_monotone = "(true, true)",
934    is_infix_op = true,
935    sqlname = "-",
936    propagates_nulls = true
937)]
938fn sub_int32(a: i32, b: i32) -> Result<i32, EvalError> {
939    a.checked_sub(b).ok_or(EvalError::NumericFieldOverflow)
940}
941
942#[sqlfunc(
943    is_monotone = "(true, true)",
944    is_infix_op = true,
945    sqlname = "-",
946    propagates_nulls = true
947)]
948fn sub_int64(a: i64, b: i64) -> Result<i64, EvalError> {
949    a.checked_sub(b).ok_or(EvalError::NumericFieldOverflow)
950}
951
952#[sqlfunc(
953    is_monotone = "(true, true)",
954    is_infix_op = true,
955    sqlname = "-",
956    propagates_nulls = true
957)]
958fn sub_uint16(a: u16, b: u16) -> Result<u16, EvalError> {
959    a.checked_sub(b)
960        .ok_or_else(|| EvalError::UInt16OutOfRange(format!("{a} - {b}").into()))
961}
962
963#[sqlfunc(
964    is_monotone = "(true, true)",
965    is_infix_op = true,
966    sqlname = "-",
967    propagates_nulls = true
968)]
969fn sub_uint32(a: u32, b: u32) -> Result<u32, EvalError> {
970    a.checked_sub(b)
971        .ok_or_else(|| EvalError::UInt32OutOfRange(format!("{a} - {b}").into()))
972}
973
974#[sqlfunc(
975    is_monotone = "(true, true)",
976    is_infix_op = true,
977    sqlname = "-",
978    propagates_nulls = true
979)]
980fn sub_uint64(a: u64, b: u64) -> Result<u64, EvalError> {
981    a.checked_sub(b)
982        .ok_or_else(|| EvalError::UInt64OutOfRange(format!("{a} - {b}").into()))
983}
984
985#[sqlfunc(
986    is_monotone = "(true, true)",
987    is_infix_op = true,
988    sqlname = "-",
989    propagates_nulls = true
990)]
991fn sub_float32(a: f32, b: f32) -> Result<f32, EvalError> {
992    let difference = a - b;
993    if difference.is_infinite() && !a.is_infinite() && !b.is_infinite() {
994        Err(EvalError::FloatOverflow)
995    } else {
996        Ok(difference)
997    }
998}
999
1000#[sqlfunc(
1001    is_monotone = "(true, true)",
1002    is_infix_op = true,
1003    sqlname = "-",
1004    propagates_nulls = true
1005)]
1006fn sub_float64(a: f64, b: f64) -> Result<f64, EvalError> {
1007    let difference = a - b;
1008    if difference.is_infinite() && !a.is_infinite() && !b.is_infinite() {
1009        Err(EvalError::FloatOverflow)
1010    } else {
1011        Ok(difference)
1012    }
1013}
1014
1015#[sqlfunc(
1016    is_monotone = "(true, true)",
1017    is_infix_op = true,
1018    sqlname = "-",
1019    propagates_nulls = true
1020)]
1021fn sub_numeric(
1022    a: OrderedDecimal<Numeric>,
1023    b: OrderedDecimal<Numeric>,
1024) -> Result<Numeric, EvalError> {
1025    let mut cx = numeric::cx_datum();
1026    let mut a = a.0;
1027    cx.sub(&mut a, &b.0);
1028    if cx.status().overflow() {
1029        Err(EvalError::FloatOverflow)
1030    } else {
1031        Ok(a)
1032    }
1033}
1034
1035// `age(a, b)` is non-monotone in *both* arguments:
1036//
1037// * Lex order on `Interval` is `(months, days, micros)`, but the Postgres
1038//   `age` algorithm independently subtracts year/month/day/... fields and
1039//   then *borrows* across boundaries when a lower field goes negative. With
1040//   `b = 2024-02-15` fixed:
1041//     a = 2024-03-31  →  age = {1 month, 16 days}
1042//     a = 2024-04-01  →  age = {1 month, 15 days}
1043//     a = 2024-05-01  →  age = {2 months, 15 days}
1044//   As `a` increases past a month boundary, `months` jumps by 1 and `days`
1045//   drops, producing a lex-smaller interval than the previous step.
1046//
1047// * Holding `a` fixed and varying `b`, the result has a V-shape at `a == b`
1048//   (sign is flipped when `a < b`):
1049//     a = 2024-02-15, b = 2024-02-14  →  age = {0 months, 1 day}
1050//     a = 2024-02-15, b = 2024-02-15  →  age = {0 months, 0 days}
1051//     a = 2024-02-15, b = 2024-02-16  →  age = {0 months, 1 day}
1052#[sqlfunc(sqlname = "age")]
1053fn age_timestamp(
1054    a: CheckedTimestamp<chrono::NaiveDateTime>,
1055    b: CheckedTimestamp<chrono::NaiveDateTime>,
1056) -> Result<Interval, EvalError> {
1057    Ok(a.age(&b)?)
1058}
1059
1060// See `age_timestamp` for why this is not monotone in either argument.
1061#[sqlfunc(sqlname = "age")]
1062fn age_timestamp_tz(
1063    a: CheckedTimestamp<chrono::DateTime<Utc>>,
1064    b: CheckedTimestamp<chrono::DateTime<Utc>>,
1065) -> Result<Interval, EvalError> {
1066    Ok(a.age(&b)?)
1067}
1068
1069#[sqlfunc(is_monotone = "(true, true)", is_infix_op = true, sqlname = "-")]
1070fn sub_timestamp(
1071    a: CheckedTimestamp<NaiveDateTime>,
1072    b: CheckedTimestamp<NaiveDateTime>,
1073) -> Result<Interval, EvalError> {
1074    Interval::from_chrono_duration(a - b)
1075        .map_err(|e| EvalError::IntervalOutOfRange(e.to_string().into()))
1076}
1077
1078#[sqlfunc(is_monotone = "(true, true)", is_infix_op = true, sqlname = "-")]
1079fn sub_timestamp_tz(
1080    a: CheckedTimestamp<chrono::DateTime<Utc>>,
1081    b: CheckedTimestamp<chrono::DateTime<Utc>>,
1082) -> Result<Interval, EvalError> {
1083    Interval::from_chrono_duration(a - b)
1084        .map_err(|e| EvalError::IntervalOutOfRange(e.to_string().into()))
1085}
1086
1087#[sqlfunc(
1088    is_monotone = "(true, true)",
1089    is_infix_op = true,
1090    sqlname = "-",
1091    propagates_nulls = true
1092)]
1093fn sub_date(a: Date, b: Date) -> i32 {
1094    a - b
1095}
1096
1097#[sqlfunc(is_monotone = "(true, true)", is_infix_op = true, sqlname = "-")]
1098fn sub_time(a: chrono::NaiveTime, b: chrono::NaiveTime) -> Result<Interval, EvalError> {
1099    Interval::from_chrono_duration(a - b)
1100        .map_err(|e| EvalError::IntervalOutOfRange(e.to_string().into()))
1101}
1102
1103#[sqlfunc(
1104    is_monotone = "(true, true)",
1105    output_type = "Interval",
1106    is_infix_op = true,
1107    sqlname = "-",
1108    propagates_nulls = true
1109)]
1110fn sub_interval(a: Interval, b: Interval) -> Result<Interval, EvalError> {
1111    b.checked_neg()
1112        .and_then(|b| b.checked_add(&a))
1113        .ok_or_else(|| EvalError::IntervalOutOfRange(format!("{a} - {b}").into()))
1114}
1115
1116// See `add_date_interval` for why this is not monotone in `interval`.
1117#[sqlfunc(
1118    is_monotone = "(true, false)",
1119    is_infix_op = true,
1120    sqlname = "-",
1121    propagates_nulls = true
1122)]
1123fn sub_date_interval(
1124    date: Date,
1125    interval: Interval,
1126) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
1127    let dt = NaiveDate::from(date).and_hms_opt(0, 0, 0).unwrap();
1128    let dt = interval
1129        .months
1130        .checked_neg()
1131        .ok_or_else(|| EvalError::IntervalOutOfRange(interval.months.to_string().into()))
1132        .and_then(|months| add_timestamp_months(&dt, months))?;
1133    let dt = dt
1134        .checked_sub_signed(interval.duration_as_chrono())
1135        .ok_or(EvalError::TimestampOutOfRange)?;
1136    Ok(dt.try_into()?)
1137}
1138
1139#[sqlfunc(
1140    is_monotone = "(false, false)",
1141    is_infix_op = true,
1142    sqlname = "-",
1143    propagates_nulls = true
1144)]
1145fn sub_time_interval(time: chrono::NaiveTime, interval: Interval) -> chrono::NaiveTime {
1146    let (t, _) = time.overflowing_sub_signed(interval.duration_as_chrono());
1147    t
1148}
1149
1150#[sqlfunc(
1151    is_monotone = "(true, true)",
1152    is_infix_op = true,
1153    sqlname = "*",
1154    propagates_nulls = true
1155)]
1156fn mul_int16(a: i16, b: i16) -> Result<i16, EvalError> {
1157    a.checked_mul(b).ok_or(EvalError::NumericFieldOverflow)
1158}
1159
1160#[sqlfunc(
1161    is_monotone = "(true, true)",
1162    is_infix_op = true,
1163    sqlname = "*",
1164    propagates_nulls = true
1165)]
1166fn mul_int32(a: i32, b: i32) -> Result<i32, EvalError> {
1167    a.checked_mul(b).ok_or(EvalError::NumericFieldOverflow)
1168}
1169
1170#[sqlfunc(
1171    is_monotone = "(true, true)",
1172    is_infix_op = true,
1173    sqlname = "*",
1174    propagates_nulls = true
1175)]
1176fn mul_int64(a: i64, b: i64) -> Result<i64, EvalError> {
1177    a.checked_mul(b).ok_or(EvalError::NumericFieldOverflow)
1178}
1179
1180#[sqlfunc(
1181    is_monotone = "(true, true)",
1182    is_infix_op = true,
1183    sqlname = "*",
1184    propagates_nulls = true
1185)]
1186fn mul_uint16(a: u16, b: u16) -> Result<u16, EvalError> {
1187    a.checked_mul(b)
1188        .ok_or_else(|| EvalError::UInt16OutOfRange(format!("{a} * {b}").into()))
1189}
1190
1191#[sqlfunc(
1192    is_monotone = "(true, true)",
1193    is_infix_op = true,
1194    sqlname = "*",
1195    propagates_nulls = true
1196)]
1197fn mul_uint32(a: u32, b: u32) -> Result<u32, EvalError> {
1198    a.checked_mul(b)
1199        .ok_or_else(|| EvalError::UInt32OutOfRange(format!("{a} * {b}").into()))
1200}
1201
1202#[sqlfunc(
1203    is_monotone = "(true, true)",
1204    is_infix_op = true,
1205    sqlname = "*",
1206    propagates_nulls = true
1207)]
1208fn mul_uint64(a: u64, b: u64) -> Result<u64, EvalError> {
1209    a.checked_mul(b)
1210        .ok_or_else(|| EvalError::UInt64OutOfRange(format!("{a} * {b}").into()))
1211}
1212
1213#[sqlfunc(
1214    is_monotone = (true, true),
1215    is_infinity_monotone = false,
1216    is_infix_op = true,
1217    sqlname = "*",
1218    propagates_nulls = true
1219)]
1220fn mul_float32(a: f32, b: f32) -> Result<f32, EvalError> {
1221    let product = a * b;
1222    if product.is_infinite() && !a.is_infinite() && !b.is_infinite() {
1223        Err(EvalError::FloatOverflow)
1224    } else if product == 0.0f32 && a != 0.0f32 && b != 0.0f32 {
1225        Err(EvalError::FloatUnderflow)
1226    } else {
1227        Ok(product)
1228    }
1229}
1230
1231#[sqlfunc(
1232    is_monotone = "(true, true)",
1233    is_infinity_monotone = false,
1234    is_infix_op = true,
1235    sqlname = "*",
1236    propagates_nulls = true
1237)]
1238fn mul_float64(a: f64, b: f64) -> Result<f64, EvalError> {
1239    let product = a * b;
1240    if product.is_infinite() && !a.is_infinite() && !b.is_infinite() {
1241        Err(EvalError::FloatOverflow)
1242    } else if product == 0.0f64 && a != 0.0f64 && b != 0.0f64 {
1243        Err(EvalError::FloatUnderflow)
1244    } else {
1245        Ok(product)
1246    }
1247}
1248
1249#[sqlfunc(
1250    is_monotone = "(true, true)",
1251    is_infinity_monotone = false,
1252    is_infix_op = true,
1253    sqlname = "*",
1254    propagates_nulls = true
1255)]
1256fn mul_numeric(mut a: Numeric, b: Numeric) -> Result<Numeric, EvalError> {
1257    let mut cx = numeric::cx_datum();
1258    cx.mul(&mut a, &b);
1259    let cx_status = cx.status();
1260    if cx_status.overflow() {
1261        Err(EvalError::FloatOverflow)
1262    } else if cx_status.subnormal() {
1263        Err(EvalError::FloatUnderflow)
1264    } else {
1265        numeric::munge_numeric(&mut a).unwrap();
1266        Ok(a)
1267    }
1268}
1269
1270#[sqlfunc(
1271    is_monotone = "(false, false)",
1272    is_infix_op = true,
1273    sqlname = "*",
1274    propagates_nulls = true
1275)]
1276fn mul_interval(a: Interval, b: f64) -> Result<Interval, EvalError> {
1277    a.checked_mul(b)
1278        .ok_or_else(|| EvalError::IntervalOutOfRange(format!("{a} * {b}").into()))
1279}
1280
1281#[sqlfunc(
1282    is_monotone = "(true, false)",
1283    is_infix_op = true,
1284    sqlname = "/",
1285    propagates_nulls = true
1286)]
1287fn div_int16(a: i16, b: i16) -> Result<i16, EvalError> {
1288    if b == 0 {
1289        Err(EvalError::DivisionByZero)
1290    } else {
1291        a.checked_div(b)
1292            .ok_or_else(|| EvalError::Int16OutOfRange(format!("{a} / {b}").into()))
1293    }
1294}
1295
1296#[sqlfunc(
1297    is_monotone = "(true, false)",
1298    is_infix_op = true,
1299    sqlname = "/",
1300    propagates_nulls = true
1301)]
1302fn div_int32(a: i32, b: i32) -> Result<i32, EvalError> {
1303    if b == 0 {
1304        Err(EvalError::DivisionByZero)
1305    } else {
1306        a.checked_div(b)
1307            .ok_or_else(|| EvalError::Int32OutOfRange(format!("{a} / {b}").into()))
1308    }
1309}
1310
1311#[sqlfunc(
1312    is_monotone = "(true, false)",
1313    is_infix_op = true,
1314    sqlname = "/",
1315    propagates_nulls = true
1316)]
1317fn div_int64(a: i64, b: i64) -> Result<i64, EvalError> {
1318    if b == 0 {
1319        Err(EvalError::DivisionByZero)
1320    } else {
1321        a.checked_div(b)
1322            .ok_or_else(|| EvalError::Int64OutOfRange(format!("{a} / {b}").into()))
1323    }
1324}
1325
1326#[sqlfunc(
1327    is_monotone = "(true, false)",
1328    is_infix_op = true,
1329    sqlname = "/",
1330    propagates_nulls = true
1331)]
1332fn div_uint16(a: u16, b: u16) -> Result<u16, EvalError> {
1333    if b == 0 {
1334        Err(EvalError::DivisionByZero)
1335    } else {
1336        Ok(a / b)
1337    }
1338}
1339
1340#[sqlfunc(
1341    is_monotone = "(true, false)",
1342    is_infix_op = true,
1343    sqlname = "/",
1344    propagates_nulls = true
1345)]
1346fn div_uint32(a: u32, b: u32) -> Result<u32, EvalError> {
1347    if b == 0 {
1348        Err(EvalError::DivisionByZero)
1349    } else {
1350        Ok(a / b)
1351    }
1352}
1353
1354#[sqlfunc(
1355    is_monotone = "(true, false)",
1356    is_infix_op = true,
1357    sqlname = "/",
1358    propagates_nulls = true
1359)]
1360fn div_uint64(a: u64, b: u64) -> Result<u64, EvalError> {
1361    if b == 0 {
1362        Err(EvalError::DivisionByZero)
1363    } else {
1364        Ok(a / b)
1365    }
1366}
1367
1368#[sqlfunc(
1369    is_monotone = "(true, false)",
1370    is_infinity_monotone = false,
1371    is_infix_op = true,
1372    sqlname = "/",
1373    propagates_nulls = true
1374)]
1375fn div_float32(a: f32, b: f32) -> Result<f32, EvalError> {
1376    if b == 0.0f32 && !a.is_nan() {
1377        Err(EvalError::DivisionByZero)
1378    } else {
1379        let quotient = a / b;
1380        if quotient.is_infinite() && !a.is_infinite() {
1381            Err(EvalError::FloatOverflow)
1382        } else if quotient == 0.0f32 && a != 0.0f32 && !b.is_infinite() {
1383            Err(EvalError::FloatUnderflow)
1384        } else {
1385            Ok(quotient)
1386        }
1387    }
1388}
1389
1390#[sqlfunc(
1391    is_monotone = "(true, false)",
1392    is_infinity_monotone = false,
1393    is_infix_op = true,
1394    sqlname = "/",
1395    propagates_nulls = true
1396)]
1397fn div_float64(a: f64, b: f64) -> Result<f64, EvalError> {
1398    if b == 0.0f64 && !a.is_nan() {
1399        Err(EvalError::DivisionByZero)
1400    } else {
1401        let quotient = a / b;
1402        if quotient.is_infinite() && !a.is_infinite() {
1403            Err(EvalError::FloatOverflow)
1404        } else if quotient == 0.0f64 && a != 0.0f64 && !b.is_infinite() {
1405            Err(EvalError::FloatUnderflow)
1406        } else {
1407            Ok(quotient)
1408        }
1409    }
1410}
1411
1412#[sqlfunc(
1413    is_monotone = "(true, false)",
1414    is_infinity_monotone = false,
1415    is_infix_op = true,
1416    sqlname = "/",
1417    propagates_nulls = true
1418)]
1419fn div_numeric(mut a: Numeric, b: Numeric) -> Result<Numeric, EvalError> {
1420    let mut cx = numeric::cx_datum();
1421
1422    cx.div(&mut a, &b);
1423    let cx_status = cx.status();
1424
1425    // checking the status for division by zero errors is insufficient because
1426    // the underlying library treats 0/0 as undefined and not division by zero.
1427    if b.is_zero() {
1428        Err(EvalError::DivisionByZero)
1429    } else if cx_status.overflow() {
1430        Err(EvalError::FloatOverflow)
1431    } else if cx_status.subnormal() {
1432        Err(EvalError::FloatUnderflow)
1433    } else {
1434        numeric::munge_numeric(&mut a).unwrap();
1435        Ok(a)
1436    }
1437}
1438
1439#[sqlfunc(
1440    is_monotone = "(false, false)",
1441    is_infix_op = true,
1442    sqlname = "/",
1443    propagates_nulls = true
1444)]
1445fn div_interval(a: Interval, b: f64) -> Result<Interval, EvalError> {
1446    if b == 0.0 {
1447        Err(EvalError::DivisionByZero)
1448    } else {
1449        a.checked_div(b)
1450            .ok_or_else(|| EvalError::IntervalOutOfRange(format!("{a} / {b}").into()))
1451    }
1452}
1453
1454#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1455fn mod_int16(a: i16, b: i16) -> Result<i16, EvalError> {
1456    if b == 0 {
1457        Err(EvalError::DivisionByZero)
1458    } else {
1459        Ok(a.checked_rem(b).unwrap_or(0))
1460    }
1461}
1462
1463#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1464fn mod_int32(a: i32, b: i32) -> Result<i32, EvalError> {
1465    if b == 0 {
1466        Err(EvalError::DivisionByZero)
1467    } else {
1468        Ok(a.checked_rem(b).unwrap_or(0))
1469    }
1470}
1471
1472#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1473fn mod_int64(a: i64, b: i64) -> Result<i64, EvalError> {
1474    if b == 0 {
1475        Err(EvalError::DivisionByZero)
1476    } else {
1477        Ok(a.checked_rem(b).unwrap_or(0))
1478    }
1479}
1480
1481#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1482fn mod_uint16(a: u16, b: u16) -> Result<u16, EvalError> {
1483    if b == 0 {
1484        Err(EvalError::DivisionByZero)
1485    } else {
1486        Ok(a % b)
1487    }
1488}
1489
1490#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1491fn mod_uint32(a: u32, b: u32) -> Result<u32, EvalError> {
1492    if b == 0 {
1493        Err(EvalError::DivisionByZero)
1494    } else {
1495        Ok(a % b)
1496    }
1497}
1498
1499#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1500fn mod_uint64(a: u64, b: u64) -> Result<u64, EvalError> {
1501    if b == 0 {
1502        Err(EvalError::DivisionByZero)
1503    } else {
1504        Ok(a % b)
1505    }
1506}
1507
1508#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1509fn mod_float32(a: f32, b: f32) -> Result<f32, EvalError> {
1510    if b == 0.0 {
1511        Err(EvalError::DivisionByZero)
1512    } else {
1513        Ok(a % b)
1514    }
1515}
1516
1517#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1518fn mod_float64(a: f64, b: f64) -> Result<f64, EvalError> {
1519    if b == 0.0 {
1520        Err(EvalError::DivisionByZero)
1521    } else {
1522        Ok(a % b)
1523    }
1524}
1525
1526#[sqlfunc(is_infix_op = true, sqlname = "%", propagates_nulls = true)]
1527fn mod_numeric(mut a: Numeric, b: Numeric) -> Result<Numeric, EvalError> {
1528    if b.is_zero() {
1529        return Err(EvalError::DivisionByZero);
1530    }
1531    let mut cx = numeric::cx_datum();
1532    // Postgres does _not_ use IEEE 754-style remainder
1533    cx.rem(&mut a, &b);
1534    numeric::munge_numeric(&mut a).unwrap();
1535    Ok(a)
1536}
1537
1538fn neg_interval_inner(a: Interval) -> Result<Interval, EvalError> {
1539    a.checked_neg()
1540        .ok_or_else(|| EvalError::IntervalOutOfRange(a.to_string().into()))
1541}
1542
1543fn log_guard_numeric(val: &Numeric, function_name: &str) -> Result<(), EvalError> {
1544    if val.is_negative() {
1545        return Err(EvalError::NegativeOutOfDomain(function_name.into()));
1546    }
1547    if val.is_zero() {
1548        return Err(EvalError::ZeroOutOfDomain(function_name.into()));
1549    }
1550    Ok(())
1551}
1552
1553#[sqlfunc(sqlname = "log", propagates_nulls = true)]
1554fn log_base_numeric(mut a: Numeric, mut b: Numeric) -> Result<Numeric, EvalError> {
1555    log_guard_numeric(&a, "log")?;
1556    log_guard_numeric(&b, "log")?;
1557    let mut cx = numeric::cx_datum();
1558    cx.ln(&mut a);
1559    cx.ln(&mut b);
1560    cx.div(&mut b, &a);
1561    if a.is_zero() {
1562        Err(EvalError::DivisionByZero)
1563    } else {
1564        // This division can result in slightly wrong answers due to the
1565        // limitation of dividing irrational numbers. To correct that, see if
1566        // rounding off the value from its `numeric::NUMERIC_DATUM_MAX_PRECISION
1567        // - 1`th position results in an integral value.
1568        cx.set_precision(usize::from(numeric::NUMERIC_DATUM_MAX_PRECISION - 1))
1569            .expect("reducing precision below max always succeeds");
1570        let mut integral_check = b.clone();
1571
1572        // `reduce` rounds to the context's final digit when the number of
1573        // digits in its argument exceeds its precision. We've contrived that to
1574        // happen by shrinking the context's precision by 1.
1575        cx.reduce(&mut integral_check);
1576
1577        // Reduced integral values always have a non-negative exponent.
1578        let mut b = if integral_check.exponent() >= 0 {
1579            // We believe our result should have been an integral
1580            integral_check
1581        } else {
1582            b
1583        };
1584
1585        numeric::munge_numeric(&mut b).unwrap();
1586        Ok(b)
1587    }
1588}
1589
1590#[sqlfunc(propagates_nulls = true)]
1591fn power(a: f64, b: f64) -> Result<f64, EvalError> {
1592    if a == 0.0 && b.is_sign_negative() {
1593        return Err(EvalError::Undefined(
1594            "zero raised to a negative power".into(),
1595        ));
1596    }
1597    if a.is_sign_negative() && b.fract() != 0.0 {
1598        // Equivalent to PG error:
1599        // > a negative number raised to a non-integer power yields a complex result
1600        return Err(EvalError::ComplexOutOfRange("pow".into()));
1601    }
1602    let res = a.powf(b);
1603    if res.is_infinite() {
1604        return Err(EvalError::FloatOverflow);
1605    }
1606    if res == 0.0 && a != 0.0 {
1607        return Err(EvalError::FloatUnderflow);
1608    }
1609    Ok(res)
1610}
1611
1612#[sqlfunc(propagates_nulls = true)]
1613fn uuid_generate_v5(a: uuid::Uuid, b: &str) -> uuid::Uuid {
1614    uuid::Uuid::new_v5(&a, b.as_bytes())
1615}
1616
1617#[sqlfunc(output_type = "Numeric", propagates_nulls = true)]
1618fn power_numeric(mut a: Numeric, b: Numeric) -> Result<Numeric, EvalError> {
1619    if a.is_zero() {
1620        if b.is_zero() {
1621            return Ok(Numeric::from(1));
1622        }
1623        if b.is_negative() {
1624            return Err(EvalError::Undefined(
1625                "zero raised to a negative power".into(),
1626            ));
1627        }
1628    }
1629    if a.is_negative() && b.exponent() < 0 {
1630        // Equivalent to PG error:
1631        // > a negative number raised to a non-integer power yields a complex result
1632        return Err(EvalError::ComplexOutOfRange("pow".into()));
1633    }
1634    let mut cx = numeric::cx_datum();
1635    cx.pow(&mut a, &b);
1636    let cx_status = cx.status();
1637    if cx_status.overflow() || (cx_status.invalid_operation() && !b.is_negative()) {
1638        Err(EvalError::FloatOverflow)
1639    } else if cx_status.subnormal() || cx_status.invalid_operation() {
1640        Err(EvalError::FloatUnderflow)
1641    } else {
1642        numeric::munge_numeric(&mut a).unwrap();
1643        Ok(a)
1644    }
1645}
1646
1647#[sqlfunc(propagates_nulls = true)]
1648fn get_bit(bytes: &[u8], index: i32) -> Result<i32, EvalError> {
1649    let err = EvalError::IndexOutOfRange {
1650        provided: index,
1651        valid_end: i32::try_from(bytes.len().saturating_mul(8)).unwrap_or(i32::MAX) - 1,
1652    };
1653
1654    let index = usize::try_from(index).map_err(|_| err.clone())?;
1655
1656    let byte_index = index / 8;
1657    let bit_index = index % 8;
1658
1659    let i = bytes
1660        .get(byte_index)
1661        .map(|b| (*b >> bit_index) & 1)
1662        .ok_or(err)?;
1663    assert!(i == 0 || i == 1);
1664    Ok(i32::from(i))
1665}
1666
1667#[sqlfunc(propagates_nulls = true)]
1668fn get_byte(bytes: &[u8], index: i32) -> Result<i32, EvalError> {
1669    let err = EvalError::IndexOutOfRange {
1670        provided: index,
1671        valid_end: i32::try_from(bytes.len()).unwrap_or(i32::MAX) - 1,
1672    };
1673    let i: &u8 = bytes
1674        .get(usize::try_from(index).map_err(|_| err.clone())?)
1675        .ok_or(err)?;
1676    Ok(i32::from(*i))
1677}
1678
1679#[sqlfunc(sqlname = "constant_time_compare_bytes", propagates_nulls = true)]
1680pub fn constant_time_eq_bytes(a: &[u8], b: &[u8]) -> bool {
1681    verify_slices_are_equal(a, b).is_ok()
1682}
1683
1684#[sqlfunc(sqlname = "constant_time_compare_strings", propagates_nulls = true)]
1685pub fn constant_time_eq_string(a: &str, b: &str) -> bool {
1686    verify_slices_are_equal(a.as_bytes(), b.as_bytes()).is_ok()
1687}
1688
1689#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1690fn range_contains_i32<'a>(a: Range<Datum<'a>>, b: i32) -> bool {
1691    a.contains_elem(&b)
1692}
1693
1694#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1695fn range_contains_i64<'a>(a: Range<Datum<'a>>, elem: i64) -> bool {
1696    a.contains_elem(&elem)
1697}
1698
1699#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1700fn range_contains_date<'a>(a: Range<Datum<'a>>, elem: Date) -> bool {
1701    a.contains_elem(&elem)
1702}
1703
1704#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1705fn range_contains_numeric<'a>(a: Range<Datum<'a>>, elem: OrderedDecimal<Numeric>) -> bool {
1706    a.contains_elem(&elem)
1707}
1708
1709#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1710fn range_contains_timestamp<'a>(
1711    a: Range<Datum<'a>>,
1712    elem: CheckedTimestamp<NaiveDateTime>,
1713) -> bool {
1714    a.contains_elem(&elem)
1715}
1716
1717#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
1718fn range_contains_timestamp_tz<'a>(
1719    a: Range<Datum<'a>>,
1720    elem: CheckedTimestamp<DateTime<Utc>>,
1721) -> bool {
1722    a.contains_elem(&elem)
1723}
1724
1725#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1726fn range_contains_i32_rev<'a>(a: Range<Datum<'a>>, b: i32) -> bool {
1727    a.contains_elem(&b)
1728}
1729
1730#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1731fn range_contains_i64_rev<'a>(a: Range<Datum<'a>>, elem: i64) -> bool {
1732    a.contains_elem(&elem)
1733}
1734
1735#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1736fn range_contains_date_rev<'a>(a: Range<Datum<'a>>, elem: Date) -> bool {
1737    a.contains_elem(&elem)
1738}
1739
1740#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1741fn range_contains_numeric_rev<'a>(a: Range<Datum<'a>>, elem: OrderedDecimal<Numeric>) -> bool {
1742    a.contains_elem(&elem)
1743}
1744
1745#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1746fn range_contains_timestamp_rev<'a>(
1747    a: Range<Datum<'a>>,
1748    elem: CheckedTimestamp<NaiveDateTime>,
1749) -> bool {
1750    a.contains_elem(&elem)
1751}
1752
1753#[sqlfunc(is_infix_op = true, sqlname = "<@", propagates_nulls = true)]
1754fn range_contains_timestamp_tz_rev<'a>(
1755    a: Range<Datum<'a>>,
1756    elem: CheckedTimestamp<DateTime<Utc>>,
1757) -> bool {
1758    a.contains_elem(&elem)
1759}
1760
1761/// Macro to define binary function for various range operations.
1762/// Parameters:
1763/// 1. Unique binary function symbol.
1764/// 2. Range function symbol.
1765/// 3. SQL name for the function.
1766macro_rules! range_fn {
1767    ($fn:expr, $range_fn:expr, $sqlname:expr) => {
1768        paste::paste! {
1769
1770            #[sqlfunc(
1771                output_type = "bool",
1772                is_infix_op = true,
1773                sqlname = $sqlname,
1774                propagates_nulls = true
1775            )]
1776            fn [< range_ $fn >]<'a>(a: Datum<'a>, b: Datum<'a>) -> Datum<'a>
1777            {
1778                if a.is_null() || b.is_null() { return Datum::Null }
1779                let l = a.unwrap_range();
1780                let r = b.unwrap_range();
1781                Datum::from(Range::<Datum<'a>>::$range_fn(&l, &r))
1782            }
1783        }
1784    };
1785}
1786
1787// RangeContainsRange is either @> or <@ depending on the order of the arguments.
1788// It doesn't influence the result, but it does influence the display string.
1789range_fn!(contains_range, contains_range, "@>");
1790range_fn!(contains_range_rev, contains_range, "<@");
1791range_fn!(overlaps, overlaps, "&&");
1792range_fn!(after, after, ">>");
1793range_fn!(before, before, "<<");
1794range_fn!(overleft, overleft, "&<");
1795range_fn!(overright, overright, "&>");
1796range_fn!(adjacent, adjacent, "-|-");
1797
1798#[sqlfunc(is_infix_op = true, sqlname = "+")]
1799fn range_union<T: Copy + Ord>(l: Range<T>, r: Range<T>) -> Result<Range<T>, EvalError> {
1800    Ok(l.union(&r)?)
1801}
1802
1803#[sqlfunc(is_infix_op = true, sqlname = "*")]
1804fn range_intersection<T: Copy + Ord>(l: Range<T>, r: Range<T>) -> Range<T> {
1805    l.intersection(&r)
1806}
1807
1808#[sqlfunc(
1809    output_type_expr = "input_types[0].scalar_type.without_modifiers().nullable(true)",
1810    is_infix_op = true,
1811    sqlname = "-",
1812    propagates_nulls = true,
1813    introduces_nulls = false
1814)]
1815fn range_difference<'a>(
1816    l: Range<Datum<'a>>,
1817    r: Range<Datum<'a>>,
1818) -> Result<Range<Datum<'a>>, EvalError> {
1819    Ok(l.difference(&r)?)
1820}
1821
1822#[sqlfunc(is_infix_op = true, sqlname = "=", negate = "Some(NotEq.into())")]
1823fn eq<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1824    // SQL equality demands that if either input is null, then the result should be null. However,
1825    // we don't need to handle this case here; it is handled when `BinaryFunc::eval` checks
1826    // `propagates_nulls`.
1827    a == b
1828}
1829
1830#[sqlfunc(is_infix_op = true, sqlname = "!=", negate = "Some(Eq.into())")]
1831fn not_eq<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1832    a != b
1833}
1834
1835#[sqlfunc(
1836    is_monotone = "(true, true)",
1837    is_infix_op = true,
1838    sqlname = "<",
1839    negate = "Some(Gte.into())"
1840)]
1841fn lt<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1842    a < b
1843}
1844
1845#[sqlfunc(
1846    is_monotone = "(true, true)",
1847    is_infix_op = true,
1848    sqlname = "<=",
1849    negate = "Some(Gt.into())"
1850)]
1851fn lte<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1852    a <= b
1853}
1854
1855#[sqlfunc(
1856    is_monotone = "(true, true)",
1857    is_infix_op = true,
1858    sqlname = ">",
1859    negate = "Some(Lte.into())"
1860)]
1861fn gt<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1862    a > b
1863}
1864
1865#[sqlfunc(
1866    is_monotone = "(true, true)",
1867    is_infix_op = true,
1868    sqlname = ">=",
1869    negate = "Some(Lt.into())"
1870)]
1871fn gte<'a>(a: ExcludeNull<Datum<'a>>, b: ExcludeNull<Datum<'a>>) -> bool {
1872    a >= b
1873}
1874
1875#[sqlfunc(sqlname = "tocharts", propagates_nulls = true)]
1876fn to_char_timestamp_format(ts: CheckedTimestamp<chrono::NaiveDateTime>, format: &str) -> String {
1877    let fmt = DateTimeFormat::compile(format);
1878    fmt.render(&*ts)
1879}
1880
1881#[sqlfunc(sqlname = "tochartstz", propagates_nulls = true)]
1882fn to_char_timestamp_tz_format(
1883    ts: CheckedTimestamp<chrono::DateTime<Utc>>,
1884    format: &str,
1885) -> String {
1886    let fmt = DateTimeFormat::compile(format);
1887    fmt.render(&*ts)
1888}
1889
1890#[sqlfunc(sqlname = "->", is_infix_op = true)]
1891fn jsonb_get_int64<'a>(a: JsonbRef<'a>, i: i64) -> Option<JsonbRef<'a>> {
1892    match a.into_datum() {
1893        Datum::List(list) => {
1894            let i = if i >= 0 {
1895                usize::cast_from(i.unsigned_abs())
1896            } else {
1897                // index backwards from the end
1898                let i = usize::cast_from(i.unsigned_abs());
1899                (list.iter().count()).wrapping_sub(i)
1900            };
1901            let v = list.iter().nth(i)?;
1902            // `v` should be valid jsonb because it came from a jsonb list, but we don't
1903            // panic on mismatch to avoid bringing down the whole system on corrupt data.
1904            // Instead, we'll return None.
1905            JsonbRef::try_from_result(Ok::<_, ()>(v)).ok()
1906        }
1907        Datum::Map(_) => None,
1908        _ => {
1909            // I have no idea why postgres does this, but we're stuck with it
1910            (i == 0 || i == -1).then_some(a)
1911        }
1912    }
1913}
1914
1915#[sqlfunc(sqlname = "->>", is_infix_op = true)]
1916fn jsonb_get_int64_stringify<'a>(
1917    a: JsonbRef<'a>,
1918    i: i64,
1919    temp_storage: &'a RowArena,
1920) -> Option<&'a str> {
1921    let json = jsonb_get_int64(a, i)?;
1922    jsonb_stringify(json.into_datum(), temp_storage)
1923}
1924
1925#[sqlfunc(sqlname = "->", is_infix_op = true)]
1926fn jsonb_get_string<'a>(a: JsonbRef<'a>, k: &str) -> Option<JsonbRef<'a>> {
1927    let dict = DatumMap::try_from_result(Ok::<_, ()>(a.into_datum())).ok()?;
1928    let v = dict.iter().find(|(k2, _v)| k == *k2).map(|(_k, v)| v)?;
1929    JsonbRef::try_from_result(Ok::<_, ()>(v)).ok()
1930}
1931
1932#[sqlfunc(sqlname = "->>", is_infix_op = true)]
1933fn jsonb_get_string_stringify<'a>(
1934    a: JsonbRef<'a>,
1935    k: &str,
1936    temp_storage: &'a RowArena,
1937) -> Option<&'a str> {
1938    let v = jsonb_get_string(a, k)?;
1939    jsonb_stringify(v.into_datum(), temp_storage)
1940}
1941
1942#[sqlfunc(sqlname = "#>", is_infix_op = true)]
1943fn jsonb_get_path<'a>(mut json: JsonbRef<'a>, b: Array<'a>) -> Option<JsonbRef<'a>> {
1944    let path = b.elements();
1945    for key in path.iter() {
1946        let key = match key {
1947            Datum::String(s) => s,
1948            Datum::Null => return None,
1949            _ => unreachable!("keys in jsonb_get_path known to be strings"),
1950        };
1951        let v = match json.into_datum() {
1952            Datum::Map(map) => map.iter().find(|(k, _)| key == *k).map(|(_k, v)| v),
1953            Datum::List(list) => {
1954                let i = strconv::parse_int64(key).ok()?;
1955                let i = if i >= 0 {
1956                    usize::cast_from(i.unsigned_abs())
1957                } else {
1958                    // index backwards from the end
1959                    let i = usize::cast_from(i.unsigned_abs());
1960                    (list.iter().count()).wrapping_sub(i)
1961                };
1962                list.iter().nth(i)
1963            }
1964            _ => return None,
1965        }?;
1966        json = JsonbRef::try_from_result(Ok::<_, ()>(v)).ok()?;
1967    }
1968    Some(json)
1969}
1970
1971#[sqlfunc(sqlname = "#>>", is_infix_op = true)]
1972fn jsonb_get_path_stringify<'a>(
1973    a: JsonbRef<'a>,
1974    b: Array<'a>,
1975    temp_storage: &'a RowArena,
1976) -> Option<&'a str> {
1977    let json = jsonb_get_path(a, b)?;
1978    jsonb_stringify(json.into_datum(), temp_storage)
1979}
1980
1981#[sqlfunc(is_infix_op = true, sqlname = "?")]
1982fn jsonb_contains_string<'a>(a: JsonbRef<'a>, k: &str) -> bool {
1983    // https://www.postgresql.org/docs/current/datatype-json.html#JSON-CONTAINMENT
1984    // When the left operand is SQL NULL (NULL::jsonb), JsonbRef::try_from_result rejects it,
1985    // so the binary evaluator never calls this function and returns NULL (see binary.rs).
1986    // So, this function only runs for non-null jsonb; a.into_datum() never sees Datum::Null.
1987    match a.into_datum() {
1988        Datum::List(list) => list.iter().any(|k2| Datum::from(k) == k2),
1989        Datum::Map(dict) => dict.iter().any(|(k2, _v)| k == k2),
1990        Datum::String(string) => string == k,
1991        _ => false,
1992    }
1993}
1994
1995#[sqlfunc(is_infix_op = true, sqlname = "?", propagates_nulls = true)]
1996// Map keys are always text.
1997fn map_contains_key<'a>(map: DatumMap<'a>, k: &str) -> bool {
1998    map.iter().any(|(k2, _v)| k == k2)
1999}
2000
2001#[sqlfunc(is_infix_op = true, sqlname = "?&")]
2002fn map_contains_all_keys<'a>(map: DatumMap<'a>, keys: Array<'a>) -> bool {
2003    keys.elements()
2004        .iter()
2005        .all(|key| !key.is_null() && map.iter().any(|(k, _v)| k == key.unwrap_str()))
2006}
2007
2008#[sqlfunc(is_infix_op = true, sqlname = "?|", propagates_nulls = true)]
2009fn map_contains_any_keys<'a>(map: DatumMap<'a>, keys: Array<'a>) -> bool {
2010    keys.elements()
2011        .iter()
2012        .any(|key| !key.is_null() && map.iter().any(|(k, _v)| k == key.unwrap_str()))
2013}
2014
2015#[sqlfunc(is_infix_op = true, sqlname = "@>", propagates_nulls = true)]
2016fn map_contains_map<'a>(map_a: DatumMap<'a>, b: DatumMap<'a>) -> bool {
2017    b.iter().all(|(b_key, b_val)| {
2018        map_a
2019            .iter()
2020            .any(|(a_key, a_val)| (a_key == b_key) && (a_val == b_val))
2021    })
2022}
2023
2024#[sqlfunc(is_infix_op = true, sqlname = "->", propagates_nulls = true)]
2025fn map_get_value<'a, T: FromDatum<'a>>(a: DatumMap<'a, T>, target_key: &str) -> Option<T> {
2026    a.typed_iter()
2027        .find(|(key, _v)| target_key == *key)
2028        .map(|(_k, v)| v)
2029}
2030
2031#[sqlfunc(is_infix_op = true, sqlname = "@>")]
2032fn list_contains_list<'a>(a: ExcludeNull<DatumList<'a>>, b: ExcludeNull<DatumList<'a>>) -> bool {
2033    // NULL is never equal to NULL. If NULL is an element of b, b cannot be contained in a, even if a contains NULL.
2034    if b.iter().contains(&Datum::Null) {
2035        false
2036    } else {
2037        b.iter()
2038            .all(|item_b| a.iter().any(|item_a| item_a == item_b))
2039    }
2040}
2041
2042#[sqlfunc(is_infix_op = true, sqlname = "<@")]
2043fn list_contains_list_rev<'a>(
2044    a: ExcludeNull<DatumList<'a>>,
2045    b: ExcludeNull<DatumList<'a>>,
2046) -> bool {
2047    list_contains_list(b, a)
2048}
2049
2050// TODO(jamii) nested loops are possibly not the fastest way to do this
2051#[sqlfunc(is_infix_op = true, sqlname = "@>")]
2052fn jsonb_contains_jsonb<'a>(a: JsonbRef<'a>, b: JsonbRef<'a>) -> bool {
2053    // https://www.postgresql.org/docs/current/datatype-json.html#JSON-CONTAINMENT
2054    fn contains(a: Datum, b: Datum, at_top_level: bool) -> bool {
2055        match (a, b) {
2056            (Datum::JsonNull, Datum::JsonNull) => true,
2057            (Datum::False, Datum::False) => true,
2058            (Datum::True, Datum::True) => true,
2059            (Datum::Numeric(a), Datum::Numeric(b)) => a == b,
2060            (Datum::String(a), Datum::String(b)) => a == b,
2061            (Datum::List(a), Datum::List(b)) => b
2062                .iter()
2063                .all(|b_elem| a.iter().any(|a_elem| contains(a_elem, b_elem, false))),
2064            (Datum::Map(a), Datum::Map(b)) => b.iter().all(|(b_key, b_val)| {
2065                a.iter()
2066                    .any(|(a_key, a_val)| (a_key == b_key) && contains(a_val, b_val, false))
2067            }),
2068
2069            // fun special case
2070            (Datum::List(a), b) => {
2071                at_top_level && a.iter().any(|a_elem| contains(a_elem, b, false))
2072            }
2073
2074            _ => false,
2075        }
2076    }
2077    contains(a.into_datum(), b.into_datum(), true)
2078}
2079
2080#[sqlfunc(is_infix_op = true, sqlname = "||")]
2081fn jsonb_concat<'a>(
2082    a: JsonbRef<'a>,
2083    b: JsonbRef<'a>,
2084    temp_storage: &'a RowArena,
2085) -> Option<JsonbRef<'a>> {
2086    let res = match (a.into_datum(), b.into_datum()) {
2087        (Datum::Map(dict_a), Datum::Map(dict_b)) => {
2088            let mut pairs = dict_b.iter().chain(dict_a.iter()).collect::<Vec<_>>();
2089            // stable sort, so if keys collide dedup prefers dict_b
2090            pairs.sort_by(|(k1, _v1), (k2, _v2)| k1.cmp(k2));
2091            pairs.dedup_by(|(k1, _v1), (k2, _v2)| k1 == k2);
2092            temp_storage.make_datum(|packer| packer.push_dict(pairs))
2093        }
2094        (Datum::List(list_a), Datum::List(list_b)) => {
2095            let elems = list_a.iter().chain(list_b.iter());
2096            temp_storage.make_datum(|packer| packer.push_list(elems))
2097        }
2098        (Datum::List(list_a), b) => {
2099            let elems = list_a.iter().chain(Some(b));
2100            temp_storage.make_datum(|packer| packer.push_list(elems))
2101        }
2102        (a, Datum::List(list_b)) => {
2103            let elems = Some(a).into_iter().chain(list_b.iter());
2104            temp_storage.make_datum(|packer| packer.push_list(elems))
2105        }
2106        _ => return None,
2107    };
2108    Some(JsonbRef::from_datum(res))
2109}
2110
2111#[sqlfunc(
2112    output_type_expr = "SqlScalarType::Jsonb.nullable(true)",
2113    is_infix_op = true,
2114    sqlname = "-",
2115    propagates_nulls = true,
2116    introduces_nulls = true
2117)]
2118fn jsonb_delete_int64<'a>(a: Datum<'a>, i: i64, temp_storage: &'a RowArena) -> Datum<'a> {
2119    match a {
2120        Datum::List(list) => {
2121            let i = if i >= 0 {
2122                usize::cast_from(i.unsigned_abs())
2123            } else {
2124                // index backwards from the end
2125                let i = usize::cast_from(i.unsigned_abs());
2126                (list.iter().count()).wrapping_sub(i)
2127            };
2128            let elems = list
2129                .iter()
2130                .enumerate()
2131                .filter(|(i2, _e)| i != *i2)
2132                .map(|(_, e)| e);
2133            temp_storage.make_datum(|packer| packer.push_list(elems))
2134        }
2135        _ => Datum::Null,
2136    }
2137}
2138
2139#[sqlfunc(
2140    output_type_expr = "SqlScalarType::Jsonb.nullable(true)",
2141    is_infix_op = true,
2142    sqlname = "-",
2143    propagates_nulls = true,
2144    introduces_nulls = true
2145)]
2146fn jsonb_delete_string<'a>(a: Datum<'a>, k: &str, temp_storage: &'a RowArena) -> Datum<'a> {
2147    match a {
2148        Datum::List(list) => {
2149            let elems = list.iter().filter(|e| Datum::from(k) != *e);
2150            temp_storage.make_datum(|packer| packer.push_list(elems))
2151        }
2152        Datum::Map(dict) => {
2153            let pairs = dict.iter().filter(|(k2, _v)| k != *k2);
2154            temp_storage.make_datum(|packer| packer.push_dict(pairs))
2155        }
2156        _ => Datum::Null,
2157    }
2158}
2159
2160#[sqlfunc(
2161    sqlname = "extractiv",
2162    propagates_nulls = true,
2163    introduces_nulls = false
2164)]
2165fn date_part_interval_numeric(units: &str, b: Interval) -> Result<Numeric, EvalError> {
2166    match units.parse() {
2167        Ok(units) => Ok(date_part_interval_inner::<Numeric>(units, b)?),
2168        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2169    }
2170}
2171
2172#[sqlfunc(
2173    sqlname = "date_partiv",
2174    propagates_nulls = true,
2175    introduces_nulls = false
2176)]
2177fn date_part_interval_f64(units: &str, b: Interval) -> Result<f64, EvalError> {
2178    match units.parse() {
2179        Ok(units) => Ok(date_part_interval_inner::<f64>(units, b)?),
2180        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2181    }
2182}
2183
2184#[sqlfunc(
2185    sqlname = "extractt",
2186    propagates_nulls = true,
2187    introduces_nulls = false
2188)]
2189fn date_part_time_numeric(units: &str, b: chrono::NaiveTime) -> Result<Numeric, EvalError> {
2190    match units.parse() {
2191        Ok(units) => Ok(date_part_time_inner::<Numeric>(units, b)?),
2192        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2193    }
2194}
2195
2196#[sqlfunc(
2197    sqlname = "date_partt",
2198    propagates_nulls = true,
2199    introduces_nulls = false
2200)]
2201fn date_part_time_f64(units: &str, b: chrono::NaiveTime) -> Result<f64, EvalError> {
2202    match units.parse() {
2203        Ok(units) => Ok(date_part_time_inner::<f64>(units, b)?),
2204        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2205    }
2206}
2207
2208#[sqlfunc(sqlname = "extractts", propagates_nulls = true)]
2209fn date_part_timestamp_timestamp_numeric(
2210    units: &str,
2211    ts: CheckedTimestamp<NaiveDateTime>,
2212) -> Result<Numeric, EvalError> {
2213    match units.parse() {
2214        Ok(units) => Ok(date_part_timestamp_inner::<_, Numeric>(units, &*ts)?),
2215        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2216    }
2217}
2218
2219#[sqlfunc(sqlname = "extracttstz", propagates_nulls = true)]
2220fn date_part_timestamp_timestamp_tz_numeric(
2221    units: &str,
2222    ts: CheckedTimestamp<DateTime<Utc>>,
2223) -> Result<Numeric, EvalError> {
2224    match units.parse() {
2225        Ok(units) => Ok(date_part_timestamp_inner::<_, Numeric>(units, &*ts)?),
2226        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2227    }
2228}
2229
2230#[sqlfunc(sqlname = "date_partts", propagates_nulls = true)]
2231fn date_part_timestamp_timestamp_f64(
2232    units: &str,
2233    ts: CheckedTimestamp<NaiveDateTime>,
2234) -> Result<f64, EvalError> {
2235    match units.parse() {
2236        Ok(units) => date_part_timestamp_inner(units, &*ts),
2237        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2238    }
2239}
2240
2241#[sqlfunc(sqlname = "date_parttstz", propagates_nulls = true)]
2242fn date_part_timestamp_timestamp_tz_f64(
2243    units: &str,
2244    ts: CheckedTimestamp<DateTime<Utc>>,
2245) -> Result<f64, EvalError> {
2246    match units.parse() {
2247        Ok(units) => date_part_timestamp_inner(units, &*ts),
2248        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2249    }
2250}
2251
2252#[sqlfunc(sqlname = "extractd", propagates_nulls = true)]
2253fn extract_date_units(units: &str, b: Date) -> Result<Numeric, EvalError> {
2254    match units.parse() {
2255        Ok(units) => Ok(extract_date_inner(units, b.into())?),
2256        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2257    }
2258}
2259
2260pub fn date_bin<T>(
2261    stride: Interval,
2262    source: CheckedTimestamp<T>,
2263    origin: CheckedTimestamp<T>,
2264) -> Result<CheckedTimestamp<T>, EvalError>
2265where
2266    T: TimestampLike,
2267{
2268    if stride.months != 0 {
2269        return Err(EvalError::DateBinOutOfRange(
2270            "timestamps cannot be binned into intervals containing months or years".into(),
2271        ));
2272    }
2273
2274    let stride_ns = match stride.duration_as_chrono().num_nanoseconds() {
2275        Some(ns) if ns <= 0 => Err(EvalError::DateBinOutOfRange(
2276            "stride must be greater than zero".into(),
2277        )),
2278        Some(ns) => Ok(ns),
2279        None => Err(EvalError::DateBinOutOfRange(
2280            format!("stride cannot exceed {}/{} nanoseconds", i64::MAX, i64::MIN,).into(),
2281        )),
2282    }?;
2283
2284    // Make sure the returned timestamp is at the start of the bin, even if the
2285    // origin is in the future. We do this here because `T` is not `Copy` and
2286    // gets moved by its subtraction operation.
2287    let sub_stride = origin > source;
2288
2289    let tm_diff = (source - origin.clone()).num_nanoseconds().ok_or_else(|| {
2290        EvalError::DateBinOutOfRange(
2291            "source and origin must not differ more than 2^63 nanoseconds".into(),
2292        )
2293    })?;
2294
2295    let remainder = tm_diff % stride_ns;
2296    let mut tm_delta = tm_diff - remainder;
2297
2298    if sub_stride && remainder != 0 {
2299        tm_delta = tm_delta.checked_sub(stride_ns).ok_or_else(|| {
2300            EvalError::DateBinOutOfRange(
2301                "source and origin must not differ more than 2^63 nanoseconds".into(),
2302            )
2303        })?;
2304    }
2305
2306    let res = origin
2307        .checked_add_signed(Duration::nanoseconds(tm_delta))
2308        .ok_or(EvalError::TimestampOutOfRange)?;
2309    Ok(CheckedTimestamp::from_timestamplike(res)?)
2310}
2311
2312// Non-monotone in `stride`: the result is `origin + floor((source - origin) /
2313// stride) * stride`. For a fixed source like `2024-01-01 12:00:00`, a 1-day
2314// stride bins to `2024-01-01 00:00:00`, but a 2-day stride bins to
2315// `2023-12-31 00:00:00` — i.e. the lex-larger interval produces an earlier
2316// timestamp. Monotone in `source`.
2317#[sqlfunc(is_monotone = "(false, true)", sqlname = "bin_unix_epoch_timestamp")]
2318fn date_bin_timestamp(
2319    stride: Interval,
2320    source: CheckedTimestamp<NaiveDateTime>,
2321) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
2322    let origin =
2323        CheckedTimestamp::from_timestamplike(DateTime::from_timestamp(0, 0).unwrap().naive_utc())
2324            .expect("must fit");
2325    date_bin(stride, source, origin)
2326}
2327
2328// See `date_bin_timestamp` for why this is not monotone in `stride`.
2329#[sqlfunc(is_monotone = "(false, true)", sqlname = "bin_unix_epoch_timestamptz")]
2330fn date_bin_timestamp_tz(
2331    stride: Interval,
2332    source: CheckedTimestamp<DateTime<Utc>>,
2333) -> Result<CheckedTimestamp<DateTime<Utc>>, EvalError> {
2334    let origin = CheckedTimestamp::from_timestamplike(DateTime::from_timestamp(0, 0).unwrap())
2335        .expect("must fit");
2336    date_bin(stride, source, origin)
2337}
2338
2339#[sqlfunc(sqlname = "date_truncts", propagates_nulls = true)]
2340fn date_trunc_units_timestamp(
2341    units: &str,
2342    ts: CheckedTimestamp<NaiveDateTime>,
2343) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
2344    match units.parse() {
2345        Ok(units) => Ok(date_trunc_inner(units, &*ts)?.try_into()?),
2346        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2347    }
2348}
2349
2350#[sqlfunc(sqlname = "date_trunctstz", propagates_nulls = true)]
2351fn date_trunc_units_timestamp_tz(
2352    units: &str,
2353    ts: CheckedTimestamp<DateTime<Utc>>,
2354) -> Result<CheckedTimestamp<DateTime<Utc>>, EvalError> {
2355    match units.parse() {
2356        Ok(units) => Ok(date_trunc_inner(units, &*ts)?.try_into()?),
2357        Err(_) => Err(EvalError::UnknownUnits(units.into())),
2358    }
2359}
2360
2361#[sqlfunc(sqlname = "date_trunciv", propagates_nulls = true)]
2362fn date_trunc_interval(units: &str, mut interval: Interval) -> Result<Interval, EvalError> {
2363    let dtf = units
2364        .parse()
2365        .map_err(|_| EvalError::UnknownUnits(units.into()))?;
2366
2367    interval
2368        .truncate_low_fields(dtf, Some(0), RoundBehavior::Truncate)
2369        .expect(
2370            "truncate_low_fields should not fail with max_precision 0 and RoundBehavior::Truncate",
2371        );
2372    Ok(interval)
2373}
2374
2375/// Parses a named timezone like `EST` or `America/New_York`, or a fixed-offset timezone like `-05:00`.
2376///
2377/// The interpretation of fixed offsets depend on whether the POSIX or ISO 8601 standard is being
2378/// used.
2379pub(crate) fn parse_timezone(tz: &str, spec: TimezoneSpec) -> Result<Timezone, EvalError> {
2380    Timezone::parse(tz, spec).map_err(|_| EvalError::InvalidTimezone(tz.into()))
2381}
2382
2383/// Converts the time datum `b`, which is assumed to be in UTC, to the timezone that the interval datum `a` is assumed
2384/// to represent. The interval is not allowed to hold months, but there are no limits on the amount of seconds.
2385/// The interval acts like a `chrono::FixedOffset`, without the `-86,400 < x < 86,400` limitation.
2386#[sqlfunc(sqlname = "timezoneit")]
2387fn timezone_interval_time_binary(
2388    interval: Interval,
2389    time: chrono::NaiveTime,
2390) -> Result<chrono::NaiveTime, EvalError> {
2391    if interval.months != 0 {
2392        Err(EvalError::InvalidTimezoneInterval)
2393    } else {
2394        Ok(time.overflowing_add_signed(interval.duration_as_chrono()).0)
2395    }
2396}
2397
2398/// Converts the timestamp datum `b`, which is assumed to be in the time of the timezone datum `a` to a timestamptz
2399/// in UTC. The interval is not allowed to hold months, but there are no limits on the amount of seconds.
2400/// The interval acts like a `chrono::FixedOffset`, without the `-86,400 < x < 86,400` limitation.
2401#[sqlfunc(sqlname = "timezoneits")]
2402fn timezone_interval_timestamp_binary(
2403    interval: Interval,
2404    ts: CheckedTimestamp<NaiveDateTime>,
2405) -> Result<CheckedTimestamp<DateTime<Utc>>, EvalError> {
2406    if interval.months != 0 {
2407        Err(EvalError::InvalidTimezoneInterval)
2408    } else {
2409        match ts.checked_sub_signed(interval.duration_as_chrono()) {
2410            Some(sub) => Ok(DateTime::from_naive_utc_and_offset(sub, Utc).try_into()?),
2411            None => Err(EvalError::TimestampOutOfRange),
2412        }
2413    }
2414}
2415
2416/// Converts the UTC timestamptz datum `b`, to the local timestamp of the timezone datum `a`.
2417/// The interval is not allowed to hold months, but there are no limits on the amount of seconds.
2418/// The interval acts like a `chrono::FixedOffset`, without the `-86,400 < x < 86,400` limitation.
2419#[sqlfunc(sqlname = "timezoneitstz")]
2420fn timezone_interval_timestamp_tz_binary(
2421    interval: Interval,
2422    tstz: CheckedTimestamp<DateTime<Utc>>,
2423) -> Result<CheckedTimestamp<NaiveDateTime>, EvalError> {
2424    if interval.months != 0 {
2425        return Err(EvalError::InvalidTimezoneInterval);
2426    }
2427    match tstz
2428        .naive_utc()
2429        .checked_add_signed(interval.duration_as_chrono())
2430    {
2431        Some(dt) => Ok(dt.try_into()?),
2432        None => Err(EvalError::TimestampOutOfRange),
2433    }
2434}
2435
2436#[sqlfunc(
2437    output_type_expr = r#"SqlScalarType::Record {
2438                fields: [
2439                    ("abbrev".into(), SqlScalarType::String.nullable(false)),
2440                    ("base_utc_offset".into(), SqlScalarType::Interval.nullable(false)),
2441                    ("dst_offset".into(), SqlScalarType::Interval.nullable(false)),
2442                ].into(),
2443                custom_id: None,
2444            }.nullable(true)"#,
2445    propagates_nulls = true,
2446    introduces_nulls = false
2447)]
2448fn timezone_offset<'a>(
2449    tz_str: &str,
2450    b: CheckedTimestamp<chrono::DateTime<Utc>>,
2451    temp_storage: &'a RowArena,
2452) -> Result<Datum<'a>, EvalError> {
2453    let tz = match Tz::from_str_insensitive(tz_str) {
2454        Ok(tz) => tz,
2455        Err(_) => return Err(EvalError::InvalidIanaTimezoneId(tz_str.into())),
2456    };
2457    let offset = tz.offset_from_utc_datetime(&b.naive_utc());
2458    // Zones without an alphabetic abbreviation get the numeric form rendered
2459    // from the offset, e.g. "+05". This matches PostgreSQL, whose tzdata files
2460    // have the same rendering applied by zic's %z expansion.
2461    let abbrev = match offset.abbreviation() {
2462        Some(abbrev) => abbrev.to_string(),
2463        None => {
2464            const SECONDS_PER_MINUTE: i64 = 60;
2465            const MINUTES_PER_HOUR: i64 = 60;
2466            let secs = (offset.base_utc_offset() + offset.dst_offset()).num_seconds();
2467            let sign = if secs < 0 { '-' } else { '+' };
2468            let (mins, s) = (
2469                secs.abs() / SECONDS_PER_MINUTE,
2470                secs.abs() % SECONDS_PER_MINUTE,
2471            );
2472            let (h, m) = (mins / MINUTES_PER_HOUR, mins % MINUTES_PER_HOUR);
2473            if s != 0 {
2474                // Unreachable for current tzdata: sub-minute offsets exist
2475                // only for pre-standardization history (e.g. Africa/Monrovia
2476                // until 1972), and tzdata gives all of them alphabetic names.
2477                // We render rather than trust that invariant forever.
2478                // chrono-tz's own Display impl instead asserts, and a scalar
2479                // function must not panic.
2480                format!("{sign}{h:02}{m:02}{s:02}")
2481            } else if m != 0 {
2482                // Fractional-hour zones, e.g. Asia/Kathmandu renders "+0545".
2483                format!("{sign}{h:02}{m:02}")
2484            } else {
2485                // Whole-hour zones, e.g. Asia/Almaty renders "+05".
2486                format!("{sign}{h:02}")
2487            }
2488        }
2489    };
2490    Ok(temp_storage.make_datum(|packer| {
2491        packer.push_list_with(|packer| {
2492            packer.push(Datum::from(abbrev.as_str()));
2493            packer.push(Datum::from(offset.base_utc_offset()));
2494            packer.push(Datum::from(offset.dst_offset()));
2495        });
2496    }))
2497}
2498
2499/// Determines if an mz_aclitem contains one of the specified privileges. This will return true if
2500/// any of the listed privileges are contained in the mz_aclitem.
2501#[sqlfunc(
2502    sqlname = "mz_aclitem_contains_privilege",
2503    output_type = "bool",
2504    propagates_nulls = true
2505)]
2506fn mz_acl_item_contains_privilege(
2507    mz_acl_item: MzAclItem,
2508    privileges: &str,
2509) -> Result<bool, EvalError> {
2510    let acl_mode = AclMode::parse_multiple_privileges(privileges)
2511        .map_err(|e: anyhow::Error| EvalError::InvalidPrivileges(e.to_string().into()))?;
2512    let contains = !mz_acl_item.acl_mode.intersection(acl_mode).is_empty();
2513    Ok(contains)
2514}
2515
2516#[sqlfunc]
2517// transliterated from postgres/src/backend/utils/adt/misc.c
2518fn parse_ident<'a>(ident: &'a str, strict: bool) -> Result<ArrayRustType<Cow<'a, str>>, EvalError> {
2519    fn is_ident_start(c: char) -> bool {
2520        matches!(c, 'A'..='Z' | 'a'..='z' | '_' | '\u{80}'..=char::MAX)
2521    }
2522
2523    fn is_ident_cont(c: char) -> bool {
2524        matches!(c, '0'..='9' | '$') || is_ident_start(c)
2525    }
2526
2527    let mut elems = vec![];
2528    let buf = &mut LexBuf::new(ident);
2529
2530    let mut after_dot = false;
2531
2532    buf.take_while(|ch| ch.is_ascii_whitespace());
2533
2534    loop {
2535        let mut missing_ident = true;
2536
2537        let c = buf.next();
2538
2539        if c == Some('"') {
2540            let s = buf.take_while(|ch| !matches!(ch, '"'));
2541
2542            if buf.next() != Some('"') {
2543                return Err(EvalError::InvalidIdentifier {
2544                    ident: ident.into(),
2545                    detail: Some("String has unclosed double quotes.".into()),
2546                });
2547            }
2548            elems.push(Cow::Borrowed(s));
2549            missing_ident = false;
2550        } else if c.map(is_ident_start).unwrap_or(false) {
2551            buf.prev();
2552            let s = buf.take_while(is_ident_cont);
2553            elems.push(Cow::Owned(s.to_ascii_lowercase()));
2554            missing_ident = false;
2555        }
2556
2557        if missing_ident {
2558            if c == Some('.') {
2559                return Err(EvalError::InvalidIdentifier {
2560                    ident: ident.into(),
2561                    detail: Some("No valid identifier before \".\".".into()),
2562                });
2563            } else if after_dot {
2564                return Err(EvalError::InvalidIdentifier {
2565                    ident: ident.into(),
2566                    detail: Some("No valid identifier after \".\".".into()),
2567                });
2568            } else {
2569                return Err(EvalError::InvalidIdentifier {
2570                    ident: ident.into(),
2571                    detail: None,
2572                });
2573            }
2574        }
2575
2576        buf.take_while(|ch| ch.is_ascii_whitespace());
2577
2578        match buf.next() {
2579            Some('.') => {
2580                after_dot = true;
2581
2582                buf.take_while(|ch| ch.is_ascii_whitespace());
2583            }
2584            Some(_) if strict => {
2585                return Err(EvalError::InvalidIdentifier {
2586                    ident: ident.into(),
2587                    detail: None,
2588                });
2589            }
2590            _ => break,
2591        }
2592    }
2593
2594    Ok(elems.into())
2595}
2596
2597fn regexp_split_to_array_re<'a>(
2598    text: &str,
2599    regexp: &Regex,
2600    temp_storage: &'a RowArena,
2601) -> Result<Datum<'a>, EvalError> {
2602    // Bound the transient `Vec<&str>` before the split builds it. The count follows the split's own
2603    // zero-length-match rule, so it refuses exactly the calls the split would build.
2604    check_build_fits_budget(
2605        || mz_regexp::regexp_split_to_array_count(text, regexp),
2606        std::mem::size_of::<&str>(),
2607        temp_storage,
2608    )?;
2609    let found = mz_regexp::regexp_split_to_array(text, regexp);
2610    // Splitting amplifies: each chunk is packed with its own tag and length, so a pattern that
2611    // splits per character costs several times the input.
2612    check_datums_fit_budget(found.iter().copied().map(Datum::String), temp_storage)?;
2613    let mut row = Row::default();
2614    let mut packer = row.packer();
2615    packer.try_push_array(
2616        &[ArrayDimension {
2617            lower_bound: 1,
2618            length: found.len(),
2619        }],
2620        found.into_iter().map(Datum::String),
2621    )?;
2622    Ok(temp_storage.push_unary_row(row))
2623}
2624
2625// NOTE: no budget pre-check, see the exception on `check_build_fits_budget`.
2626#[sqlfunc(propagates_nulls = true)]
2627fn pretty_sql<'a>(sql: &str, width: i32, temp_storage: &'a RowArena) -> Result<&'a str, EvalError> {
2628    let width =
2629        usize::try_from(width).map_err(|_| EvalError::PrettyError("invalid width".into()))?;
2630    let pretty = pretty_str(
2631        sql,
2632        PrettyConfig {
2633            width,
2634            format_mode: FormatMode::Simple,
2635        },
2636    )
2637    .map_err(|e| EvalError::PrettyError(e.to_string().into()))?;
2638    let pretty = temp_storage.push_string(pretty);
2639    Ok(pretty)
2640}
2641
2642// NOTE: no budget pre-check, see the exception on `check_build_fits_budget`.
2643#[sqlfunc]
2644fn redact_sql(sql: &str) -> Result<String, EvalError> {
2645    let stmts = mz_sql_parser::parser::parse_statements(sql)
2646        .map_err(|e| EvalError::RedactError(e.to_string().into()))?;
2647    match stmts.len() {
2648        1 => Ok(stmts[0].ast.to_ast_string_redacted()),
2649        n => Err(EvalError::RedactError(
2650            format!("expected a single statement, found {n}").into(),
2651        )),
2652    }
2653}
2654
2655#[sqlfunc(propagates_nulls = true)]
2656fn starts_with(a: &str, b: &str) -> bool {
2657    a.starts_with(b)
2658}
2659
2660#[sqlfunc(
2661    sqlname = "||",
2662    is_infix_op = true,
2663    propagates_nulls = true,
2664    // Text concatenation is monotonic in its second argument, because if I change the
2665    // second argument but don't change the first argument, then we won't find a difference
2666    // in that part of the concatenation result that came from the first argument, so we'll
2667    // find the difference that comes from changing the second argument.
2668    // (It's not monotonic in its first argument, because e.g.,
2669    // 'A' < 'AA' but 'AZ' > 'AAZ'.)
2670    is_monotone = (false, true),
2671)]
2672fn text_concat_binary(a: &str, b: &str, temp_storage: &RowArena) -> Result<String, EvalError> {
2673    if a.len() + b.len() > max_string_func_result_bytes(temp_storage) {
2674        return Err(EvalError::LengthTooLarge);
2675    }
2676    let mut buf = String::with_capacity(a.len() + b.len());
2677    buf.push_str(a);
2678    buf.push_str(b);
2679    Ok(buf)
2680}
2681
2682#[sqlfunc(propagates_nulls = true, introduces_nulls = false)]
2683fn like_escape<'a>(
2684    pattern: &str,
2685    b: &str,
2686    temp_storage: &'a RowArena,
2687) -> Result<&'a str, EvalError> {
2688    let escape = like_pattern::EscapeBehavior::from_str(b)?;
2689    let normalized = like_pattern::normalize_pattern(pattern, escape)?;
2690    Ok(temp_storage.push_string(normalized))
2691}
2692
2693#[sqlfunc(is_infix_op = true, sqlname = "like")]
2694fn is_like_match_case_sensitive(haystack: &str, pattern: &str) -> Result<bool, EvalError> {
2695    like_pattern::compile(pattern, false).map(|needle| needle.is_match(haystack))
2696}
2697
2698#[sqlfunc(is_infix_op = true, sqlname = "ilike")]
2699fn is_like_match_case_insensitive(haystack: &str, pattern: &str) -> Result<bool, EvalError> {
2700    like_pattern::compile(pattern, true).map(|needle| needle.is_match(haystack))
2701}
2702
2703#[sqlfunc(is_infix_op = true, sqlname = "~")]
2704fn is_regexp_match_case_sensitive(haystack: &str, needle: &str) -> Result<bool, EvalError> {
2705    let regex = build_regex(needle, "")?;
2706    Ok(regex.is_match(haystack))
2707}
2708
2709#[sqlfunc(is_infix_op = true, sqlname = "~*")]
2710fn is_regexp_match_case_insensitive(haystack: &str, needle: &str) -> Result<bool, EvalError> {
2711    let regex = build_regex(needle, "i")?;
2712    Ok(regex.is_match(haystack))
2713}
2714
2715fn regexp_match_static<'a>(
2716    haystack: Datum<'a>,
2717    temp_storage: &'a RowArena,
2718    needle: &regex::Regex,
2719) -> Result<Datum<'a>, EvalError> {
2720    let mut row = Row::default();
2721    let mut packer = row.packer();
2722    if needle.captures_len() > 1 {
2723        // The regex contains capture groups, so return an array containing the
2724        // matched text in each capture group, unless the entire match fails.
2725        // Individual capture groups may also be null if that group did not
2726        // participate in the match.
2727        match needle.captures(haystack.unwrap_str()) {
2728            None => packer.push(Datum::Null),
2729            Some(captures) => packer.try_push_array(
2730                &[ArrayDimension {
2731                    lower_bound: 1,
2732                    length: captures.len() - 1,
2733                }],
2734                // Skip the 0th capture group, which is the whole match.
2735                captures.iter().skip(1).map(|mtch| match mtch {
2736                    None => Datum::Null,
2737                    Some(mtch) => Datum::String(mtch.as_str()),
2738                }),
2739            )?,
2740        }
2741    } else {
2742        // The regex contains no capture groups, so return a one-element array
2743        // containing the match, or null if there is no match.
2744        match needle.find(haystack.unwrap_str()) {
2745            None => packer.push(Datum::Null),
2746            Some(mtch) => packer.try_push_array(
2747                &[ArrayDimension {
2748                    lower_bound: 1,
2749                    length: 1,
2750                }],
2751                iter::once(Datum::String(mtch.as_str())),
2752            )?,
2753        };
2754    };
2755    Ok(temp_storage.push_unary_row(row))
2756}
2757
2758/// Sets `limit` based on the presence of 'g' in `flags` for use in `Regex::replacen`,
2759/// and removes 'g' from `flags` if present.
2760pub(crate) fn regexp_replace_parse_flags(flags: &str) -> (usize, Cow<'_, str>) {
2761    // 'g' means to replace all instead of the first. Use a Cow to avoid allocating in the fast
2762    // path. We could switch build_regex to take an iter which would also achieve that.
2763    let (limit, flags) = if flags.contains('g') {
2764        let flags = flags.replace('g', "");
2765        (0, Cow::Owned(flags))
2766    } else {
2767        (1, Cow::Borrowed(flags))
2768    };
2769    (limit, flags)
2770}
2771
2772pub fn build_regex(needle: &str, flags: &str) -> Result<Regex, EvalError> {
2773    let mut case_insensitive = false;
2774    // Note: Postgres accepts it when both flags are present, taking the last one. We do the same.
2775    for f in flags.chars() {
2776        match f {
2777            'i' => {
2778                case_insensitive = true;
2779            }
2780            'c' => {
2781                case_insensitive = false;
2782            }
2783            _ => return Err(EvalError::InvalidRegexFlag(f)),
2784        }
2785    }
2786    Ok(Regex::new(needle, case_insensitive)?)
2787}
2788
2789#[sqlfunc(sqlname = "repeat")]
2790fn repeat_string(string: &str, count: i32, temp_storage: &RowArena) -> Result<String, EvalError> {
2791    let len = usize::try_from(count).unwrap_or(0);
2792    if len.saturating_mul(string.len()) > max_string_func_result_bytes(temp_storage) {
2793        return Err(EvalError::LengthTooLarge);
2794    }
2795    Ok(string.repeat(len))
2796}
2797
2798/// Constructs a new zero or one dimensional array out of an arbitrary number of
2799/// scalars.
2800///
2801/// If `datums` is empty, constructs a zero-dimensional array. Otherwise,
2802/// constructs a one dimensional array whose lower bound is one and whose length
2803/// is equal to `datums.len()`.
2804fn array_create_scalar<'a>(
2805    datums: &[Datum<'a>],
2806    temp_storage: &'a RowArena,
2807) -> Result<Datum<'a>, EvalError> {
2808    let mut dims = &[ArrayDimension {
2809        lower_bound: 1,
2810        length: datums.len(),
2811    }][..];
2812    if datums.is_empty() {
2813        // Per PostgreSQL, empty arrays are represented with zero dimensions,
2814        // not one dimension of zero length. We write this condition a little
2815        // strangely to satisfy the borrow checker while avoiding an allocation.
2816        dims = &[];
2817    }
2818    check_datums_fit_budget(datums.iter().copied(), temp_storage)?;
2819    let datum = temp_storage.try_make_datum(|packer| packer.try_push_array(dims, datums))?;
2820    Ok(datum)
2821}
2822
2823fn stringify_datum<'a, B>(
2824    buf: &mut B,
2825    d: Datum<'a>,
2826    ty: &SqlScalarType,
2827) -> Result<strconv::Nestable, EvalError>
2828where
2829    B: FormatBuffer,
2830{
2831    use SqlScalarType::*;
2832    match &ty {
2833        AclItem => Ok(strconv::format_acl_item(buf, d.unwrap_acl_item())),
2834        Bool => Ok(strconv::format_bool(buf, d.unwrap_bool())),
2835        Int16 => Ok(strconv::format_int16(buf, d.unwrap_int16())),
2836        Int32 => Ok(strconv::format_int32(buf, d.unwrap_int32())),
2837        Int64 => Ok(strconv::format_int64(buf, d.unwrap_int64())),
2838        UInt16 => Ok(strconv::format_uint16(buf, d.unwrap_uint16())),
2839        UInt32 | Oid | RegClass | RegProc | RegType => {
2840            Ok(strconv::format_uint32(buf, d.unwrap_uint32()))
2841        }
2842        UInt64 => Ok(strconv::format_uint64(buf, d.unwrap_uint64())),
2843        Float32 => Ok(strconv::format_float32(buf, d.unwrap_float32())),
2844        Float64 => Ok(strconv::format_float64(buf, d.unwrap_float64())),
2845        Numeric { .. } => Ok(strconv::format_numeric(buf, &d.unwrap_numeric())),
2846        Date => Ok(strconv::format_date(buf, d.unwrap_date())),
2847        Time => Ok(strconv::format_time(buf, d.unwrap_time())),
2848        Timestamp { .. } => Ok(strconv::format_timestamp(buf, &d.unwrap_timestamp())),
2849        TimestampTz { .. } => Ok(strconv::format_timestamptz(buf, &d.unwrap_timestamptz())),
2850        Interval => Ok(strconv::format_interval(buf, d.unwrap_interval())),
2851        Bytes => Ok(strconv::format_bytes(buf, d.unwrap_bytes())),
2852        String | VarChar { .. } | PgLegacyName => Ok(strconv::format_string(buf, d.unwrap_str())),
2853        Char { length } => Ok(strconv::format_string(
2854            buf,
2855            &mz_repr::adt::char::format_str_pad(d.unwrap_str(), *length),
2856        )),
2857        PgLegacyChar => {
2858            format_pg_legacy_char(buf, d.unwrap_uint8())?;
2859            Ok(strconv::Nestable::MayNeedEscaping)
2860        }
2861        Jsonb => Ok(strconv::format_jsonb(buf, JsonbRef::from_datum(d))),
2862        Uuid => Ok(strconv::format_uuid(buf, d.unwrap_uuid())),
2863        Record { fields, .. } => {
2864            let mut fields = fields.iter();
2865            strconv::format_record(buf, d.unwrap_list(), |buf, d| {
2866                let (_name, ty) = fields.next().unwrap();
2867                if d.is_null() {
2868                    Ok(buf.write_null())
2869                } else {
2870                    stringify_datum(buf.nonnull_buffer(), d, &ty.scalar_type)
2871                }
2872            })
2873        }
2874        Array(elem_type) => strconv::format_array(
2875            buf,
2876            &d.unwrap_array().dims().into_iter().collect::<Vec<_>>(),
2877            d.unwrap_array().elements(),
2878            |buf, d| {
2879                if d.is_null() {
2880                    Ok(buf.write_null())
2881                } else {
2882                    stringify_datum(buf.nonnull_buffer(), d, elem_type)
2883                }
2884            },
2885        ),
2886        List { element_type, .. } => strconv::format_list(buf, d.unwrap_list(), |buf, d| {
2887            if d.is_null() {
2888                Ok(buf.write_null())
2889            } else {
2890                stringify_datum(buf.nonnull_buffer(), d, element_type)
2891            }
2892        }),
2893        Map { value_type, .. } => strconv::format_map(buf, &d.unwrap_map(), |buf, d| {
2894            if d.is_null() {
2895                Ok(buf.write_null())
2896            } else {
2897                stringify_datum(buf.nonnull_buffer(), d, value_type)
2898            }
2899        }),
2900        Int2Vector => strconv::format_legacy_vector(buf, d.unwrap_array().elements(), |buf, d| {
2901            stringify_datum(buf.nonnull_buffer(), d, &SqlScalarType::Int16)
2902        }),
2903        MzTimestamp { .. } => Ok(strconv::format_mz_timestamp(buf, d.unwrap_mz_timestamp())),
2904        Range { element_type } => strconv::format_range(buf, &d.unwrap_range(), |buf, d| match d {
2905            Some(d) => stringify_datum(buf.nonnull_buffer(), *d, element_type),
2906            None => Ok::<_, EvalError>(buf.write_null()),
2907        }),
2908        MzAclItem => Ok(strconv::format_mz_acl_item(buf, d.unwrap_mz_acl_item())),
2909    }
2910}
2911
2912#[sqlfunc]
2913fn position(substring: &str, string: &str) -> Result<i32, EvalError> {
2914    let char_index = string.find(substring);
2915
2916    if let Some(char_index) = char_index {
2917        // find the index in char space
2918        let string_prefix = &string[0..char_index];
2919
2920        let num_prefix_chars = string_prefix.chars().count();
2921        let num_prefix_chars = i32::try_from(num_prefix_chars)
2922            .map_err(|_| EvalError::Int32OutOfRange(num_prefix_chars.to_string().into()))?;
2923
2924        Ok(num_prefix_chars + 1)
2925    } else {
2926        Ok(0)
2927    }
2928}
2929
2930#[sqlfunc]
2931fn strpos(string: &str, substring: &str) -> Result<i32, EvalError> {
2932    position(substring, string)
2933}
2934
2935#[sqlfunc(
2936    propagates_nulls = true,
2937    // `left` is unfortunately not monotonic (at least for negative second arguments),
2938    // because 'aa' < 'z', but `left(_, -1)` makes 'a' > ''.
2939    is_monotone = (false, false)
2940)]
2941fn left<'a>(string: &'a str, b: i32) -> Result<&'a str, EvalError> {
2942    let n = i64::from(b);
2943
2944    let mut byte_indices = string.char_indices().map(|(i, _)| i);
2945
2946    let end_in_bytes = match n.cmp(&0) {
2947        Ordering::Equal => 0,
2948        Ordering::Greater => {
2949            let n = usize::try_from(n).map_err(|_| {
2950                EvalError::InvalidParameterValue(format!("invalid parameter n: {:?}", n).into())
2951            })?;
2952            // nth from the back
2953            byte_indices.nth(n).unwrap_or(string.len())
2954        }
2955        Ordering::Less => {
2956            let n = usize::try_from(n.abs() - 1).map_err(|_| {
2957                EvalError::InvalidParameterValue(format!("invalid parameter n: {:?}", n).into())
2958            })?;
2959            byte_indices.rev().nth(n).unwrap_or(0)
2960        }
2961    };
2962
2963    Ok(&string[..end_in_bytes])
2964}
2965
2966#[sqlfunc(propagates_nulls = true)]
2967fn right<'a>(string: &'a str, n: i32) -> Result<&'a str, EvalError> {
2968    let mut byte_indices = string.char_indices().map(|(i, _)| i);
2969
2970    let start_in_bytes = if n == 0 {
2971        string.len()
2972    } else if n > 0 {
2973        let n = usize::try_from(n - 1).map_err(|_| {
2974            EvalError::InvalidParameterValue(format!("invalid parameter n: {:?}", n).into())
2975        })?;
2976        // nth from the back
2977        byte_indices.rev().nth(n).unwrap_or(0)
2978    } else if n == i32::MIN {
2979        // this seems strange but Postgres behaves like this
2980        0
2981    } else {
2982        let n = n.abs();
2983        let n = usize::try_from(n).map_err(|_| {
2984            EvalError::InvalidParameterValue(format!("invalid parameter n: {:?}", n).into())
2985        })?;
2986        byte_indices.nth(n).unwrap_or(string.len())
2987    };
2988
2989    Ok(&string[start_in_bytes..])
2990}
2991
2992#[sqlfunc(sqlname = "btrim", propagates_nulls = true)]
2993fn trim<'a>(a: &'a str, trim_chars: &str) -> &'a str {
2994    a.trim_matches(|c| trim_chars.contains(c))
2995}
2996
2997#[sqlfunc(sqlname = "ltrim", propagates_nulls = true)]
2998fn trim_leading<'a>(a: &'a str, trim_chars: &str) -> &'a str {
2999    a.trim_start_matches(|c| trim_chars.contains(c))
3000}
3001
3002#[sqlfunc(sqlname = "rtrim", propagates_nulls = true)]
3003fn trim_trailing<'a>(a: &'a str, trim_chars: &str) -> &'a str {
3004    a.trim_end_matches(|c| trim_chars.contains(c))
3005}
3006
3007#[sqlfunc(
3008    sqlname = "array_length",
3009    propagates_nulls = true,
3010    introduces_nulls = true
3011)]
3012fn array_length<'a>(a: Array<'a>, b: i64) -> Result<Option<i32>, EvalError> {
3013    let i = match usize::try_from(b) {
3014        Ok(0) | Err(_) => return Ok(None),
3015        Ok(n) => n - 1,
3016    };
3017    Ok(match a.dims().into_iter().nth(i) {
3018        None => None,
3019        Some(dim) => Some(
3020            dim.length
3021                .try_into()
3022                .map_err(|_| EvalError::Int32OutOfRange(dim.length.to_string().into()))?,
3023        ),
3024    })
3025}
3026
3027#[sqlfunc(is_infix_op = true)]
3028// TODO(benesch): remove potentially dangerous usage of `as`.
3029#[allow(clippy::as_conversions)]
3030fn array_lower<'a>(a: Array<'a>, i: i64) -> Result<Option<i32>, EvalError> {
3031    if i < 1 {
3032        return Ok(None);
3033    }
3034    a.dims()
3035        .into_iter()
3036        .nth(i as usize - 1)
3037        .map(|dim| {
3038            let (lower, _upper) = dim.dimension_bounds();
3039            lower
3040                .try_into()
3041                .map_err(|_| EvalError::Int32OutOfRange(lower.to_string().into()))
3042        })
3043        .transpose()
3044}
3045
3046#[sqlfunc(
3047    output_type_expr = "input_types[0].scalar_type.without_modifiers().nullable(true)",
3048    sqlname = "array_remove",
3049    propagates_nulls = false,
3050    introduces_nulls = false
3051)]
3052fn array_remove<'a>(
3053    arr: Array<'a>,
3054    b: Datum<'a>,
3055    temp_storage: &'a RowArena,
3056) -> Result<Datum<'a>, EvalError> {
3057    // Zero-dimensional arrays are empty by definition
3058    if arr.dims().len() == 0 {
3059        return Ok(Datum::Array(arr));
3060    }
3061
3062    // array_remove only supports one-dimensional arrays
3063    if arr.dims().len() > 1 {
3064        return Err(EvalError::MultidimensionalArrayRemovalNotSupported);
3065    }
3066
3067    let mut dims = arr.dims().into_iter().collect::<Vec<_>>();
3068    // Removal can't grow the result, but the transient `Vec<Datum>` it filters into is a fresh
3069    // input-scaled allocation. One-dimensional by the check above, so dim 0's length is the count.
3070    check_build_fits_budget(
3071        || dims[0].length,
3072        std::mem::size_of::<Datum<'a>>(),
3073        temp_storage,
3074    )?;
3075
3076    let elems: Vec<_> = arr.elements().iter().filter(|v| v != &b).collect();
3077    // This access is safe because `dims` is guaranteed to be non-empty
3078    dims[0] = ArrayDimension {
3079        lower_bound: 1,
3080        length: elems.len(),
3081    };
3082
3083    Ok(temp_storage.try_make_datum(|packer| packer.try_push_array(&dims, elems))?)
3084}
3085
3086#[sqlfunc(is_infix_op = true)]
3087// TODO(benesch): remove potentially dangerous usage of `as`.
3088#[allow(clippy::as_conversions)]
3089fn array_upper<'a>(a: Array<'a>, i: i64) -> Result<Option<i32>, EvalError> {
3090    if i < 1 {
3091        return Ok(None);
3092    }
3093    a.dims()
3094        .into_iter()
3095        .nth(i as usize - 1)
3096        .map(|dim| {
3097            let (_lower, upper) = dim.dimension_bounds();
3098            upper
3099                .try_into()
3100                .map_err(|_| EvalError::Int32OutOfRange(upper.to_string().into()))
3101        })
3102        .transpose()
3103}
3104
3105#[sqlfunc(
3106    is_infix_op = true,
3107    sqlname = "array_contains",
3108    propagates_nulls = true,
3109    introduces_nulls = false
3110)]
3111fn array_contains<'a>(a: Datum<'a>, array: Array<'a>) -> bool {
3112    array.elements().iter().any(|e| e == a)
3113}
3114
3115#[sqlfunc(is_infix_op = true, sqlname = "@>")]
3116fn array_contains_array<'a>(a: Array<'a>, b: Array<'a>) -> bool {
3117    let a = a.elements();
3118    let b = b.elements();
3119
3120    // NULL is never equal to NULL. If NULL is an element of b, b cannot be contained in a, even if a contains NULL.
3121    if b.iter().contains(&Datum::Null) {
3122        false
3123    } else {
3124        b.iter()
3125            .all(|item_b| a.iter().any(|item_a| item_a == item_b))
3126    }
3127}
3128
3129#[sqlfunc(is_infix_op = true, sqlname = "<@")]
3130fn array_contains_array_rev<'a>(a: Array<'a>, b: Array<'a>) -> bool {
3131    array_contains_array(b, a)
3132}
3133
3134#[sqlfunc(
3135    output_type_expr = "input_types[0].scalar_type.without_modifiers().nullable(true)",
3136    is_infix_op = true,
3137    sqlname = "||",
3138    propagates_nulls = false,
3139    introduces_nulls = false
3140)]
3141fn array_array_concat<'a>(
3142    a: Option<Array<'a>>,
3143    b: Option<Array<'a>>,
3144    temp_storage: &'a RowArena,
3145) -> Result<Option<Array<'a>>, EvalError> {
3146    let Some(a_array) = a else {
3147        return Ok(b);
3148    };
3149    let Some(b_array) = b else {
3150        return Ok(a);
3151    };
3152
3153    let a_dims: Vec<ArrayDimension> = a_array.dims().into_iter().collect();
3154    let b_dims: Vec<ArrayDimension> = b_array.dims().into_iter().collect();
3155
3156    let a_ndims = a_dims.len();
3157    let b_ndims = b_dims.len();
3158
3159    // Per PostgreSQL, if either of the input arrays is zero dimensional,
3160    // the output is the other array, no matter their dimensions.
3161    if a_ndims == 0 {
3162        return Ok(b);
3163    } else if b_ndims == 0 {
3164        return Ok(a);
3165    }
3166
3167    // Postgres supports concatenating arrays of different dimensions,
3168    // as long as one of the arrays has the same type as an element of
3169    // the other array, i.e. `int[2][4] || int[4]` (or `int[4] || int[2][4]`)
3170    // works, because each element of `int[2][4]` is an `int[4]`.
3171    // This check is separate from the one below because Postgres gives a
3172    // specific error message if the number of dimensions differs by more
3173    // than one.
3174    // This cast is safe since MAX_ARRAY_DIMENSIONS is 6
3175    // Can be replaced by .abs_diff once it is stabilized
3176    // TODO(benesch): remove potentially dangerous usage of `as`.
3177    #[allow(clippy::as_conversions)]
3178    if (a_ndims as isize - b_ndims as isize).abs() > 1 {
3179        return Err(EvalError::IncompatibleArrayDimensions {
3180            dims: Some((a_ndims, b_ndims)),
3181        });
3182    }
3183
3184    let mut dims;
3185
3186    // After the checks above, we are certain that:
3187    // - neither array is zero dimensional nor empty
3188    // - both arrays have the same number of dimensions, or differ
3189    //   at most by one.
3190    match a_ndims.cmp(&b_ndims) {
3191        // If both arrays have the same number of dimensions, validate
3192        // that their inner dimensions are the same and concatenate the
3193        // arrays.
3194        Ordering::Equal => {
3195            if &a_dims[1..] != &b_dims[1..] {
3196                return Err(EvalError::IncompatibleArrayDimensions { dims: None });
3197            }
3198            dims = vec![ArrayDimension {
3199                lower_bound: a_dims[0].lower_bound,
3200                length: a_dims[0].length + b_dims[0].length,
3201            }];
3202            dims.extend(&a_dims[1..]);
3203        }
3204        // If `a` has less dimensions than `b`, this is an element-array
3205        // concatenation, which requires that `a` has the same dimensions
3206        // as an element of `b`.
3207        Ordering::Less => {
3208            if &a_dims[..] != &b_dims[1..] {
3209                return Err(EvalError::IncompatibleArrayDimensions { dims: None });
3210            }
3211            dims = vec![ArrayDimension {
3212                lower_bound: b_dims[0].lower_bound,
3213                // Since `a` is treated as an element of `b`, the length of
3214                // the first dimension of `b` is incremented by one, as `a` is
3215                // non-empty.
3216                length: b_dims[0].length + 1,
3217            }];
3218            dims.extend(a_dims);
3219        }
3220        // If `a` has more dimensions than `b`, this is an array-element
3221        // concatenation, which requires that `b` has the same dimensions
3222        // as an element of `a`.
3223        Ordering::Greater => {
3224            if &a_dims[1..] != &b_dims[..] {
3225                return Err(EvalError::IncompatibleArrayDimensions { dims: None });
3226            }
3227            dims = vec![ArrayDimension {
3228                lower_bound: a_dims[0].lower_bound,
3229                // Since `b` is treated as an element of `a`, the length of
3230                // the first dimension of `a` is incremented by one, as `b`
3231                // is non-empty.
3232                length: a_dims[0].length + 1,
3233            }];
3234            dims.extend(b_dims);
3235        }
3236    }
3237
3238    let elems = a_array.elements().iter().chain(b_array.elements().iter());
3239
3240    let datum = temp_storage.try_make_datum(|packer| packer.try_push_array(&dims, elems))?;
3241    Ok(Some(datum.unwrap_array()))
3242}
3243
3244#[sqlfunc(
3245    is_infix_op = true,
3246    sqlname = "||",
3247    propagates_nulls = false,
3248    introduces_nulls = false
3249)]
3250fn list_list_concat<'a, T: FromDatum<'a>>(
3251    a: Option<DatumList<'a, T>>,
3252    b: Option<DatumList<'a, T>>,
3253    temp_storage: &'a RowArena,
3254) -> Option<DatumList<'a, T>> {
3255    let Some(a) = a else {
3256        return b;
3257    };
3258    let Some(b) = b else {
3259        return Some(a);
3260    };
3261
3262    Some(temp_storage.make_datum_list(a.typed_iter().chain(b.typed_iter())))
3263}
3264
3265#[sqlfunc(is_infix_op = true, sqlname = "||", propagates_nulls = false)]
3266fn list_element_concat<'a, T: FromDatum<'a>>(
3267    a: Option<DatumList<'a, T>>,
3268    b: T,
3269    temp_storage: &'a RowArena,
3270) -> DatumList<'a, T> {
3271    let a_elems = a.into_iter().flat_map(|a| a.typed_iter());
3272    temp_storage.make_datum_list(a_elems.chain(std::iter::once(b)))
3273}
3274
3275// Note that the output type corresponds to the _second_ parameter's input type.
3276#[sqlfunc(is_infix_op = true, sqlname = "||", propagates_nulls = false)]
3277fn element_list_concat<'a, T: FromDatum<'a>>(
3278    a: T,
3279    b: Option<DatumList<'a, T>>,
3280    temp_storage: &'a RowArena,
3281) -> DatumList<'a, T> {
3282    let b_elems = b.into_iter().flat_map(|b| b.typed_iter());
3283    temp_storage.make_datum_list(std::iter::once(a).chain(b_elems))
3284}
3285
3286#[sqlfunc(sqlname = "list_remove")]
3287fn list_remove<'a, T: FromDatum<'a>>(
3288    a: DatumList<'a, T>,
3289    b: T,
3290    temp_storage: &'a RowArena,
3291) -> DatumList<'a, T> {
3292    temp_storage.make_datum_list(a.typed_iter().filter(|elem| *elem != b))
3293}
3294
3295#[sqlfunc(sqlname = "digest")]
3296fn digest_string(to_digest: &str, digest_fn: &str) -> Result<Vec<u8>, EvalError> {
3297    digest_inner(to_digest.as_bytes(), digest_fn)
3298}
3299
3300#[sqlfunc(sqlname = "digest")]
3301fn digest_bytes(to_digest: &[u8], digest_fn: &str) -> Result<Vec<u8>, EvalError> {
3302    digest_inner(to_digest, digest_fn)
3303}
3304
3305fn digest_inner(bytes: &[u8], digest_fn: &str) -> Result<Vec<u8>, EvalError> {
3306    match digest_fn {
3307        "md5" => Ok(Md5::digest(bytes).to_vec()),
3308        "sha1" => Ok(digest::digest(&digest::SHA1_FOR_LEGACY_USE_ONLY, bytes)
3309            .as_ref()
3310            .to_vec()),
3311        "sha224" => Ok(digest::digest(&digest::SHA224, bytes).as_ref().to_vec()),
3312        "sha256" => Ok(digest::digest(&digest::SHA256, bytes).as_ref().to_vec()),
3313        "sha384" => Ok(digest::digest(&digest::SHA384, bytes).as_ref().to_vec()),
3314        "sha512" => Ok(digest::digest(&digest::SHA512, bytes).as_ref().to_vec()),
3315        other => Err(EvalError::InvalidHashAlgorithm(other.into())),
3316    }
3317}
3318
3319#[sqlfunc]
3320fn mz_render_typmod(oid: u32, typmod: i32) -> String {
3321    match Type::from_oid_and_typmod(oid, typmod) {
3322        Ok(typ) => typ.constraint().display_or("").to_string(),
3323        // Match dubious PostgreSQL behavior of outputting the unmodified
3324        // `typmod` when positive if the type OID/typmod is invalid.
3325        Err(_) if typmod >= 0 => format!("({typmod})"),
3326        Err(_) => "".into(),
3327    }
3328}
3329
3330#[cfg(test)]
3331mod test {
3332    use mz_repr::PropDatum;
3333    use proptest::prelude::*;
3334
3335    use super::*;
3336    use crate::{Eval, MirScalarExpr};
3337
3338    #[mz_ore::test]
3339    fn variant_names_unique() {
3340        // `from_variant_name` resolves the first variant with a matching
3341        // canonical name, so a duplicate name (from a typo in `func_name!` or
3342        // colliding function names) would silently shadow a variant.
3343        fn assert_unique(enum_name: &str, names: impl Iterator<Item = &'static str>) {
3344            let mut seen = std::collections::BTreeSet::new();
3345            for name in names {
3346                assert!(seen.insert(name), "duplicate {enum_name} name: {name}");
3347            }
3348        }
3349        assert_unique("UnaryFunc", UnaryFunc::variant_names());
3350        assert_unique("BinaryFunc", BinaryFunc::variant_names());
3351        assert_unique("VariadicFunc", VariadicFunc::variant_names());
3352    }
3353
3354    #[mz_ore::test]
3355    fn add_interval_months() {
3356        let dt = ym(2000, 1);
3357
3358        assert_eq!(add_timestamp_months(&*dt, 0).unwrap(), dt);
3359        assert_eq!(add_timestamp_months(&*dt, 1).unwrap(), ym(2000, 2));
3360        assert_eq!(add_timestamp_months(&*dt, 12).unwrap(), ym(2001, 1));
3361        assert_eq!(add_timestamp_months(&*dt, 13).unwrap(), ym(2001, 2));
3362        assert_eq!(add_timestamp_months(&*dt, 24).unwrap(), ym(2002, 1));
3363        assert_eq!(add_timestamp_months(&*dt, 30).unwrap(), ym(2002, 7));
3364
3365        // and negatives
3366        assert_eq!(add_timestamp_months(&*dt, -1).unwrap(), ym(1999, 12));
3367        assert_eq!(add_timestamp_months(&*dt, -12).unwrap(), ym(1999, 1));
3368        assert_eq!(add_timestamp_months(&*dt, -13).unwrap(), ym(1998, 12));
3369        assert_eq!(add_timestamp_months(&*dt, -24).unwrap(), ym(1998, 1));
3370        assert_eq!(add_timestamp_months(&*dt, -30).unwrap(), ym(1997, 7));
3371
3372        // and going over a year boundary by less than a year
3373        let dt = ym(1999, 12);
3374        assert_eq!(add_timestamp_months(&*dt, 1).unwrap(), ym(2000, 1));
3375        let end_of_month_dt = NaiveDate::from_ymd_opt(1999, 12, 31)
3376            .unwrap()
3377            .and_hms_opt(9, 9, 9)
3378            .unwrap();
3379        assert_eq!(
3380            // leap year
3381            add_timestamp_months(&end_of_month_dt, 2).unwrap(),
3382            NaiveDate::from_ymd_opt(2000, 2, 29)
3383                .unwrap()
3384                .and_hms_opt(9, 9, 9)
3385                .unwrap()
3386                .try_into()
3387                .unwrap(),
3388        );
3389        assert_eq!(
3390            // not leap year
3391            add_timestamp_months(&end_of_month_dt, 14).unwrap(),
3392            NaiveDate::from_ymd_opt(2001, 2, 28)
3393                .unwrap()
3394                .and_hms_opt(9, 9, 9)
3395                .unwrap()
3396                .try_into()
3397                .unwrap(),
3398        );
3399    }
3400
3401    fn ym(year: i32, month: u32) -> CheckedTimestamp<NaiveDateTime> {
3402        NaiveDate::from_ymd_opt(year, month, 1)
3403            .unwrap()
3404            .and_hms_opt(9, 9, 9)
3405            .unwrap()
3406            .try_into()
3407            .unwrap()
3408    }
3409
3410    #[mz_ore::test]
3411    fn array_lower_upper_respect_lower_bound() {
3412        use mz_repr::adt::array::ArrayDimension;
3413        use mz_repr::{Datum, RowArena};
3414
3415        let arena = RowArena::new();
3416
3417        // Builds a one-dimensional array with the given lower bound and length,
3418        // then returns (array_lower(_, 1), array_upper(_, 1)).
3419        let bounds = |lower_bound: isize, length: usize| {
3420            let dims = [ArrayDimension {
3421                lower_bound,
3422                length,
3423            }];
3424            let elems = vec![Datum::Int32(0); length];
3425            let datum = arena.make_datum(|packer| packer.try_push_array(&dims, elems).unwrap());
3426            let arr = match datum {
3427                Datum::Array(arr) => arr,
3428                other => panic!("expected array, got {other:?}"),
3429            };
3430            (array_lower(arr, 1).unwrap(), array_upper(arr, 1).unwrap())
3431        };
3432
3433        // Default lower bound of 1: array_fill(0, ARRAY[3]).
3434        assert_eq!(bounds(1, 3), (Some(1), Some(3)));
3435        // Lower bound of 5: array_fill(0, ARRAY[3], ARRAY[5]) => [5:7].
3436        assert_eq!(bounds(5, 3), (Some(5), Some(7)));
3437        // Negative lower bound: array_fill(0, ARRAY[3], ARRAY[-3]) => [-3:-1].
3438        assert_eq!(bounds(-3, 3), (Some(-3), Some(-1)));
3439
3440        // Out-of-range dimensions return None rather than the bound.
3441        let dims = [ArrayDimension {
3442            lower_bound: 5,
3443            length: 3,
3444        }];
3445        let elems = vec![Datum::Int32(0); 3];
3446        let datum = arena.make_datum(|packer| packer.try_push_array(&dims, elems).unwrap());
3447        let arr = match datum {
3448            Datum::Array(arr) => arr,
3449            other => panic!("expected array, got {other:?}"),
3450        };
3451        assert_eq!(array_lower(arr, 0).unwrap(), None);
3452        assert_eq!(array_upper(arr, 0).unwrap(), None);
3453        assert_eq!(array_lower(arr, 2).unwrap(), None);
3454        assert_eq!(array_upper(arr, 2).unwrap(), None);
3455    }
3456
3457    #[mz_ore::test]
3458    #[cfg_attr(miri, ignore)] // unsupported operation: can't call foreign function `decNumberFromInt32` on OS `linux`
3459    fn test_is_monotone() {
3460        use proptest::prelude::*;
3461
3462        /// Asserts that the function is either monotonically increasing or decreasing over
3463        /// the given sets of arguments.
3464        fn assert_monotone<'a, const N: usize>(
3465            expr: &MirScalarExpr,
3466            arena: &'a RowArena,
3467            datums: &[[Datum<'a>; N]],
3468        ) {
3469            // TODO: assertions for nulls, errors
3470            let Ok(results) = datums
3471                .iter()
3472                .map(|args| expr.eval(args.as_slice(), arena))
3473                .collect::<Result<Vec<_>, _>>()
3474            else {
3475                return;
3476            };
3477
3478            let forward = results.iter().tuple_windows().all(|(a, b)| a <= b);
3479            let reverse = results.iter().tuple_windows().all(|(a, b)| a >= b);
3480            assert!(
3481                forward || reverse,
3482                "expected {expr} to be monotone, but passing {datums:?} returned {results:?}"
3483            );
3484        }
3485
3486        fn proptest_binary<'a>(
3487            func: BinaryFunc,
3488            arena: &'a RowArena,
3489            left: impl Strategy<Value = PropDatum>,
3490            right: impl Strategy<Value = PropDatum>,
3491        ) {
3492            let (left_monotone, right_monotone) = func.is_monotone();
3493            let expr = MirScalarExpr::CallBinary {
3494                func,
3495                expr1: Box::new(MirScalarExpr::column(0)),
3496                expr2: Box::new(MirScalarExpr::column(1)),
3497            };
3498            proptest!(|(
3499                mut left in proptest::array::uniform3(left),
3500                mut right in proptest::array::uniform3(right),
3501            )| {
3502                left.sort();
3503                right.sort();
3504                if left_monotone {
3505                    for r in &right {
3506                        let args: Vec<[_; 2]> = left
3507                            .iter()
3508                            .map(|l| [Datum::from(l), Datum::from(r)])
3509                            .collect();
3510                        assert_monotone(&expr, arena, &args);
3511                    }
3512                }
3513                if right_monotone {
3514                    for l in &left {
3515                        let args: Vec<[_; 2]> = right
3516                            .iter()
3517                            .map(|r| [Datum::from(l), Datum::from(r)])
3518                            .collect();
3519                        assert_monotone(&expr, arena, &args);
3520                    }
3521                }
3522            });
3523        }
3524
3525        let interesting_strs: Vec<_> = SqlScalarType::String.interesting_datums().collect();
3526        let str_datums = proptest::strategy::Union::new([
3527            proptest::string::string_regex("[A-Z]{0,10}")
3528                .expect("valid regex")
3529                .prop_map(|s| PropDatum::String(s.to_string()))
3530                .boxed(),
3531            (0..interesting_strs.len())
3532                .prop_map(move |i| {
3533                    let Datum::String(val) = interesting_strs[i] else {
3534                        unreachable!("interesting strings has non-strings")
3535                    };
3536                    PropDatum::String(val.to_string())
3537                })
3538                .boxed(),
3539        ]);
3540
3541        let interesting_i32s: Vec<Datum<'static>> =
3542            SqlScalarType::Int32.interesting_datums().collect();
3543        let i32_datums = proptest::strategy::Union::new([
3544            any::<i32>().prop_map(PropDatum::Int32).boxed(),
3545            (0..interesting_i32s.len())
3546                .prop_map(move |i| {
3547                    let Datum::Int32(val) = interesting_i32s[i] else {
3548                        unreachable!("interesting int32 has non-i32s")
3549                    };
3550                    PropDatum::Int32(val)
3551                })
3552                .boxed(),
3553            (-10i32..10).prop_map(PropDatum::Int32).boxed(),
3554        ]);
3555
3556        let arena = RowArena::new();
3557
3558        // It would be interesting to test all funcs here, but we currently need to hardcode
3559        // the generators for the argument types, which makes this tedious. Choose an interesting
3560        // subset for now.
3561        proptest_binary(
3562            BinaryFunc::AddInt32(AddInt32),
3563            &arena,
3564            &i32_datums,
3565            &i32_datums,
3566        );
3567        proptest_binary(SubInt32.into(), &arena, &i32_datums, &i32_datums);
3568        proptest_binary(MulInt32.into(), &arena, &i32_datums, &i32_datums);
3569        proptest_binary(DivInt32.into(), &arena, &i32_datums, &i32_datums);
3570        proptest_binary(TextConcatBinary.into(), &arena, &str_datums, &str_datums);
3571        proptest_binary(Left.into(), &arena, &str_datums, &i32_datums);
3572    }
3573}