Skip to main content

mz_expr_derive_impl/
sqlfunc.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
10use darling::FromMeta;
11use proc_macro2::{Delimiter, Ident, Spacing, TokenStream, TokenTree};
12use quote::{ToTokens, quote};
13use syn::spanned::Spanned;
14use syn::{Expr, Lifetime, Lit};
15
16/// Modifiers passed as key-value pairs to the `#[sqlfunc]` macro.
17#[derive(Debug, Default, darling::FromMeta)]
18pub(crate) struct Modifiers {
19    /// An optional expression that evaluates to a boolean indicating whether the function is
20    /// monotone with respect to its arguments. Defined for unary and binary functions.
21    is_monotone: Option<Expr>,
22    /// Optional expression evaluating to a boolean: whether `is_monotone`'s
23    /// endpoint-sampling guarantee still holds when an operand may be infinite.
24    /// Set `false` for multiplication and division. Applies to binary functions.
25    is_infinity_monotone: Option<Expr>,
26    /// The SQL name for the function. Applies to all functions.
27    sqlname: Option<SqlName>,
28    /// Whether the function preserves uniqueness. Applies to unary functions.
29    preserves_uniqueness: Option<Expr>,
30    /// The inverse of the function, if it exists. Applies to unary functions.
31    inverse: Option<Expr>,
32    /// The negated function, if it exists. Applies to binary functions.
33    negate: Option<Expr>,
34    /// Whether the function is an infix operator. Applies to binary functions, and needs to
35    /// be specified.
36    is_infix_op: Option<Expr>,
37    /// The output type of the function, if it cannot be inferred. Applies to all functions.
38    output_type: Option<syn::Path>,
39    /// The output type of the function as an expression. Applies to binary and variadic functions.
40    output_type_expr: Option<Expr>,
41    /// Optional expression evaluating to a boolean indicating whether the function could error.
42    /// Applies to all functions.
43    could_error: Option<Expr>,
44    /// Whether the function propagates nulls. Applies to binary and variadic functions.
45    propagates_nulls: Option<Expr>,
46    /// Whether the function introduces nulls. Applies to all functions.
47    introduces_nulls: Option<Expr>,
48    /// Whether the function is associative. Applies to variadic functions.
49    is_associative: Option<Expr>,
50    /// Whether the function is a noop cast. Applies to unary functions.
51    is_eliminable_cast: Option<Expr>,
52    /// Whether to generate a snapshot test for the function. Defaults to false.
53    test: Option<bool>,
54}
55
56/// A name for the SQL function. It can be either a literal or a macro, thus we
57/// can't use `String` or `syn::Expr` directly.
58#[derive(Debug)]
59enum SqlName {
60    /// A literal string.
61    Literal(syn::Lit),
62    /// A macro expression.
63    Macro(syn::ExprMacro),
64}
65
66impl quote::ToTokens for SqlName {
67    fn to_tokens(&self, tokens: &mut TokenStream) {
68        let name = match self {
69            SqlName::Literal(lit) => quote! { #lit },
70            SqlName::Macro(mac) => quote! { #mac },
71        };
72        tokens.extend(name);
73    }
74}
75
76impl darling::FromMeta for SqlName {
77    fn from_value(value: &Lit) -> darling::Result<Self> {
78        Ok(Self::Literal(value.clone()))
79    }
80    fn from_expr(expr: &Expr) -> darling::Result<Self> {
81        match expr {
82            Expr::Lit(lit) => Self::from_value(&lit.lit),
83            Expr::Macro(mac) => Ok(Self::Macro(mac.clone())),
84            // Syn sometimes inserts groups, see `FromMeta::from_expr` for
85            // details.
86            Expr::Group(mac) => Self::from_expr(&mac.expr),
87            _ => Err(darling::Error::unexpected_expr_type(expr)),
88        }
89    }
90}
91
92/// Implementation for the `#[sqlfunc]` macro. The first parameter is the attribute
93/// arguments, the second is the function body. The third parameter indicates
94/// whether to include the test function in the output.
95///
96/// The feature `test` must be enabled to include the test function.
97pub fn sqlfunc(
98    attr: TokenStream,
99    item: TokenStream,
100    include_test: bool,
101) -> darling::Result<TokenStream> {
102    let mut attr_args = darling::ast::NestedMeta::parse_meta_list(attr.clone())?;
103
104    // Check if the first attribute arg is a bare Path (struct name for variadic).
105    let struct_ty = match attr_args.first() {
106        Some(darling::ast::NestedMeta::Meta(syn::Meta::Path(_))) => {
107            let darling::ast::NestedMeta::Meta(syn::Meta::Path(path)) = attr_args.remove(0) else {
108                unreachable!()
109            };
110            Some(path)
111        }
112        _ => None,
113    };
114
115    let modifiers = Modifiers::from_list(&attr_args).unwrap();
116    let generate_tests = modifiers.test.unwrap_or(false);
117    let func = syn::parse2::<syn::ItemFn>(item.clone())?;
118    let tokens = match determine_arity(&func) {
119        Arity::Nullary => Err(darling::Error::custom("Nullary functions not supported")),
120        Arity::Unary { arena: false } => unary_func(&func, modifiers, &attr),
121        Arity::Unary { arena: true } => Err(darling::Error::custom(
122            "Unary functions do not yet support RowArena.",
123        )),
124        Arity::Binary { arena } => binary_func(&func, modifiers, arena, &attr),
125        Arity::Variadic { arena, has_self } => {
126            variadic_func(&func, modifiers, struct_ty, arena, has_self, &attr)
127        }
128    }?;
129
130    let test = (generate_tests && include_test).then(|| generate_test(attr, item, &func.sig.ident));
131
132    Ok(quote! {
133        #tokens
134        #test
135    })
136}
137
138/// Renders a token stream as compact, canonical source text.
139///
140/// The rendering walks the token trees rather than calling
141/// `TokenStream::to_string`, whose spacing differs between the compiler and
142/// the proc-macro2 fallback and has changed across rustc releases. Idents
143/// and literals are rendered verbatim, so whitespace inside a string literal
144/// is preserved and distinguishes two bodies. Spacing between tokens follows
145/// fixed rules that read like formatted code for signatures and attribute
146/// arguments, and is merely deterministic for bodies.
147fn render_tokens(tokens: &TokenStream) -> String {
148    let mut out = String::new();
149    render_into(tokens.clone(), &mut out);
150    out
151}
152
153fn render_into(tokens: TokenStream, out: &mut String) {
154    /// What the previous token was, as far as spacing cares.
155    enum Last {
156        Start,
157        Ident,
158        Punct { ch: char, joint: bool },
159        Other,
160    }
161    let mut last = Last::Start;
162    for tree in tokens {
163        let glue = match &last {
164            Last::Start => true,
165            Last::Punct { joint: true, .. } => true,
166            Last::Punct { ch, .. } if "&.!#<".contains(*ch) => true,
167            _ => out.ends_with("::"),
168        } || match &tree {
169            TokenTree::Punct(p) if ",;:.?>".contains(p.as_char()) => true,
170            // Generic parameter lists and macro invocations attach to the
171            // preceding ident.
172            TokenTree::Punct(p) if "<!".contains(p.as_char()) => matches!(last, Last::Ident),
173            // Call and index groups attach to the callee, including a closing
174            // generic angle bracket, but not to an arrow's `>`.
175            TokenTree::Group(g) => {
176                matches!(g.delimiter(), Delimiter::Parenthesis | Delimiter::Bracket)
177                    && (matches!(last, Last::Ident)
178                        || (matches!(last, Last::Punct { ch: '>', .. })
179                            && !out.ends_with("->")
180                            && !out.ends_with("=>")))
181            }
182            _ => false,
183        };
184        if !glue {
185            out.push(' ');
186        }
187        last = match tree {
188            TokenTree::Ident(ident) => {
189                out.push_str(&ident.to_string());
190                Last::Ident
191            }
192            TokenTree::Literal(lit) => {
193                out.push_str(&lit.to_string());
194                Last::Other
195            }
196            TokenTree::Punct(p) => {
197                out.push(p.as_char());
198                Last::Punct {
199                    ch: p.as_char(),
200                    joint: p.spacing() == Spacing::Joint,
201                }
202            }
203            TokenTree::Group(g) => {
204                let (open, close) = match g.delimiter() {
205                    Delimiter::Parenthesis => ("(", ")"),
206                    Delimiter::Bracket => ("[", "]"),
207                    Delimiter::Brace => ("{", "}"),
208                    Delimiter::None => ("", ""),
209                };
210                out.push_str(open);
211                render_into(g.stream(), out);
212                out.push_str(close);
213                Last::Other
214            }
215        };
216    }
217}
218
219/// FNV-1a, 64 bit. A fingerprint, not a secure hash: it only needs to be
220/// stable across builds and sensitive to any change in its input.
221fn fnv1a64(bytes: &[u8]) -> u64 {
222    bytes.iter().fold(0xcbf29ce484222325, |hash, byte| {
223        (hash ^ u64::from(*byte)).wrapping_mul(0x100000001b3)
224    })
225}
226
227/// Emits the source-derived members of the generated `FuncName` impl: the
228/// `SQLFUNC` const (declaration text, types-only signature, body
229/// fingerprint) and `sqlfunc_input_types`, which yields the column types the
230/// function naturally consumes when every parameter type has one. Both exist
231/// only under mz-expr's `func-registry` feature, like the trait members they
232/// implement.
233///
234/// Text comes from the token trees via [`render_tokens`], so formatting and
235/// comments do not affect it. `input_tys_raw` are the parameter types as
236/// written, for the signature. `input_tys` are the erased types the trait
237/// impl uses, for the probes.
238fn sqlfunc_source(
239    attr: &TokenStream,
240    func: &syn::ItemFn,
241    input_tys_raw: &[syn::Type],
242    output_ty_raw: &syn::Type,
243    input_tys: &[syn::Type],
244) -> TokenStream {
245    let attr = render_tokens(attr);
246    let attr = if attr.is_empty() {
247        String::new()
248    } else {
249        format!("({attr})")
250    };
251    let decl = format!(
252        "#[sqlfunc{attr}] {}",
253        render_tokens(&func.sig.to_token_stream())
254    );
255    let render_type = |ty: &syn::Type| render_tokens(&ty.to_token_stream());
256    let signature = format!(
257        "fn({}) -> {}",
258        input_tys_raw
259            .iter()
260            .map(render_type)
261            .collect::<Vec<_>>()
262            .join(", "),
263        render_type(output_ty_raw)
264    );
265    let body_fingerprint = fnv1a64(render_tokens(&func.block.to_token_stream()).as_bytes());
266    let probes: Vec<TokenStream> = input_tys.iter().flat_map(probe_column_types).collect();
267    quote! {
268        #[cfg(feature = "func-registry")]
269        const SQLFUNC: Option<crate::func::SqlFuncSource> = Some(crate::func::SqlFuncSource {
270            decl: #decl,
271            signature: #signature,
272            body_fingerprint: #body_fingerprint,
273        });
274
275        #[cfg(feature = "func-registry")]
276        fn sqlfunc_input_types() -> Option<Vec<mz_repr::SqlColumnType>> {
277            use crate::func::registry::{ProbeColumnType as _, ProbeColumnTypeFallback as _};
278            [#(#probes),*].into_iter().collect()
279        }
280    }
281}
282
283/// One probe expression per datum a parameter consumes. `Variadic<T>` stands
284/// for two `T` arguments and `OptionalArg<T>` for one present `T`.
285fn probe_column_types(ty: &syn::Type) -> Vec<TokenStream> {
286    if let Some((wrapper, inner)) = single_generic_arg(ty) {
287        match wrapper.as_str() {
288            "Variadic" => return vec![probe_column_type(inner), probe_column_type(inner)],
289            "OptionalArg" => return vec![probe_column_type(inner)],
290            _ => {}
291        }
292    }
293    vec![probe_column_type(ty)]
294}
295
296/// An expression of type `Option<SqlColumnType>`: the column type of `ty` if
297/// it implements `AsColumnType`, else `None`. Resolved by autoref
298/// specialization on `ColumnTypeProbe`, so it needs no trait bound the macro
299/// cannot check.
300fn probe_column_type(ty: &syn::Type) -> TokenStream {
301    let ty = staticize_lifetimes(ty);
302    quote! {
303        (&crate::func::registry::ColumnTypeProbe::<#ty>(::std::marker::PhantomData)).column_type()
304    }
305}
306
307/// The last path segment's name and its single type argument, for types
308/// shaped like `Wrapper<T>`.
309fn single_generic_arg(ty: &syn::Type) -> Option<(String, &syn::Type)> {
310    let syn::Type::Path(path) = ty else {
311        return None;
312    };
313    let segment = path.path.segments.last()?;
314    let syn::PathArguments::AngleBracketed(args) = &segment.arguments else {
315        return None;
316    };
317    let mut types = args.args.iter().filter_map(|arg| match arg {
318        syn::GenericArgument::Type(ty) => Some(ty),
319        _ => None,
320    });
321    let inner = types.next()?;
322    if types.next().is_some() {
323        return None;
324    }
325    Some((segment.ident.to_string(), inner))
326}
327
328/// Replaces every lifetime in `ty` with `'static`, so a type written against
329/// the trait impl's `'a` can be named in a function body.
330fn staticize_lifetimes(ty: &syn::Type) -> syn::Type {
331    let mut ty = ty.clone();
332    fn walk(ty: &mut syn::Type) {
333        match ty {
334            syn::Type::Reference(r) => {
335                r.lifetime = Some(Lifetime::new("'static", r.span()));
336                walk(&mut r.elem);
337            }
338            syn::Type::Tuple(t) => t.elems.iter_mut().for_each(walk),
339            syn::Type::Path(p) => {
340                for segment in &mut p.path.segments {
341                    if let syn::PathArguments::AngleBracketed(args) = &mut segment.arguments {
342                        for arg in &mut args.args {
343                            match arg {
344                                syn::GenericArgument::Lifetime(lt) => {
345                                    *lt = Lifetime::new("'static", lt.span());
346                                }
347                                syn::GenericArgument::Type(ty) => walk(ty),
348                                _ => {}
349                            }
350                        }
351                    }
352                }
353            }
354            _ => {}
355        }
356    }
357    walk(&mut ty);
358    ty
359}
360
361#[cfg(any(feature = "test", test))]
362fn generate_test(attr: TokenStream, item: TokenStream, name: &Ident) -> TokenStream {
363    let attr = attr.to_string();
364    let item = item.to_string();
365    let test_name = Ident::new(&format!("test_{}", name), name.span());
366    let fn_name = name.to_string();
367
368    quote! {
369        #[cfg(test)]
370        #[cfg_attr(miri, ignore)] // unsupported operation: extern static `pidfd_spawnp` is not supported by Miri
371        #[mz_ore::test]
372        fn #test_name() {
373            let (output, input) = mz_expr_derive_impl::test_sqlfunc_str(#attr, #item);
374            insta::assert_snapshot!(#fn_name, output, &input);
375        }
376    }
377}
378
379#[cfg(not(any(feature = "test", test)))]
380fn generate_test(_attr: TokenStream, _item: TokenStream, _name: &Ident) -> TokenStream {
381    quote! {}
382}
383
384/// Checks if the last parameter of the function is a `&RowArena`.
385fn last_is_arena(func: &syn::ItemFn) -> bool {
386    func.sig.inputs.last().map_or(false, |last| {
387        if let syn::FnArg::Typed(pat) = last {
388            if let syn::Type::Reference(reference) = &*pat.ty {
389                if let syn::Type::Path(path) = &*reference.elem {
390                    return path.path.is_ident("RowArena");
391                }
392            }
393        }
394        false
395    })
396}
397
398/// Arity classification for a function annotated with `#[sqlfunc]`.
399enum Arity {
400    Nullary,
401    Unary { arena: bool },
402    Binary { arena: bool },
403    Variadic { arena: bool, has_self: bool },
404}
405
406/// Checks whether a parameter's type is `Variadic<...>` or `OptionalArg<...>`,
407/// which indicates the function should be treated as variadic regardless of
408/// parameter count.
409fn is_variadic_arg(arg: &syn::FnArg) -> bool {
410    if let syn::FnArg::Typed(pat) = arg {
411        if let syn::Type::Path(path) = &*pat.ty {
412            if let Some(segment) = path.path.segments.last() {
413                let ident = segment.ident.to_string();
414                return ident == "Variadic" || ident == "OptionalArg";
415            }
416        }
417    }
418    false
419}
420
421/// Determines the arity of a function annotated with `#[sqlfunc]`.
422///
423/// Accounts for `&self` receivers, trailing `&RowArena` parameters, and
424/// parameter types like `Variadic<T>` or `OptionalArg<T>` that indicate
425/// variadic dispatch.
426fn determine_arity(func: &syn::ItemFn) -> Arity {
427    let arena = last_is_arena(func);
428    let has_self = matches!(func.sig.inputs.first(), Some(syn::FnArg::Receiver(_)));
429
430    let mut effective_count = func.sig.inputs.len();
431    if arena {
432        effective_count -= 1;
433    }
434    if has_self {
435        effective_count -= 1;
436    }
437
438    // Check if any effective parameter uses a variadic-typed wrapper.
439    let start = if has_self { 1 } else { 0 };
440    let end = if arena {
441        func.sig.inputs.len() - 1
442    } else {
443        func.sig.inputs.len()
444    };
445    let has_variadic_param = func
446        .sig
447        .inputs
448        .iter()
449        .skip(start)
450        .take(end - start)
451        .any(is_variadic_arg);
452
453    if has_variadic_param || effective_count >= 3 {
454        Arity::Variadic { arena, has_self }
455    } else {
456        match effective_count {
457            0 => Arity::Nullary,
458            1 => Arity::Unary { arena },
459            2 => Arity::Binary { arena },
460            _ => unreachable!(),
461        }
462    }
463}
464
465/// Convert an identifier to a camel-cased identifier.
466/// Checks if a parameter type accepts NULL.
467///
468/// `Option<T>` always accepts NULL. `OptionalArg<T>` delegates to `T`.
469/// `Datum` accepts NULL (it passes through raw values including null).
470/// Everything else (references, concrete types) rejects NULL.
471fn is_nullable_type(ty: &syn::Type) -> bool {
472    if let syn::Type::Path(type_path) = ty {
473        if let Some(last_segment) = type_path.path.segments.last() {
474            let ident = &last_segment.ident;
475            if ident == "Option" || ident == "Datum" {
476                return true;
477            }
478            if ident == "OptionalArg" {
479                // OptionalArg<T> delegates nullability to T.
480                if let syn::PathArguments::AngleBracketed(args) = &last_segment.arguments {
481                    if let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() {
482                        return is_nullable_type(inner_ty);
483                    }
484                }
485                return false;
486            }
487        }
488    }
489    false
490}
491
492/// Checks if a type is `Variadic<T>`.
493fn is_variadic_type(ty: &syn::Type) -> bool {
494    if let syn::Type::Path(type_path) = ty {
495        if let Some(last_segment) = type_path.path.segments.last() {
496            return last_segment.ident == "Variadic";
497        }
498    }
499    false
500}
501
502/// For a `Variadic<T>` type, checks if `T` accepts NULL.
503fn variadic_element_is_nullable(ty: &syn::Type) -> bool {
504    if let syn::Type::Path(type_path) = ty {
505        if let Some(last_segment) = type_path.path.segments.last() {
506            if let syn::PathArguments::AngleBracketed(args) = &last_segment.arguments {
507                if let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() {
508                    return is_nullable_type(inner_ty);
509                }
510            }
511        }
512    }
513    false
514}
515
516/// Generates per-position nullability checks for non-nullable parameters.
517///
518/// For each parameter that rejects NULL (not `Option`, not `OptionalArg<Option<..>>`),
519/// generates a check that the corresponding input position is nullable. For `Variadic<T>`
520/// with non-nullable `T`, generates a check over all remaining input positions.
521fn non_nullable_position_checks(param_types: &[syn::Type]) -> Vec<TokenStream> {
522    let mut checks = Vec::new();
523    for (i, ty) in param_types.iter().enumerate() {
524        if is_variadic_type(ty) {
525            if !variadic_element_is_nullable(ty) {
526                checks.push(quote! { || input_types.iter().skip(#i).any(|t| t.nullable) });
527            }
528        } else if !is_nullable_type(ty) {
529            checks.push(quote! { || input_types.get(#i).map_or(false, |t| t.nullable) });
530        }
531    }
532    checks
533}
534
535fn camel_case(ident: &Ident) -> Ident {
536    let mut result = String::new();
537    let mut capitalize_next = true;
538    for c in ident.to_string().chars() {
539        if c == '_' {
540            capitalize_next = true;
541        } else if capitalize_next {
542            result.push(c.to_ascii_uppercase());
543            capitalize_next = false;
544        } else {
545            result.push(c);
546        }
547    }
548    Ident::new(&result, ident.span())
549}
550
551/// Extracts generic type parameters from a function signature.
552/// Returns an empty Vec if there are no type parameters.
553fn find_generic_type_params(func: &syn::ItemFn) -> Vec<Ident> {
554    func.sig
555        .generics
556        .params
557        .iter()
558        .filter_map(|p| {
559            if let syn::GenericParam::Type(tp) = p {
560                Some(tp.ident.clone())
561            } else {
562                None
563            }
564        })
565        .collect()
566}
567
568/// How a generic type parameter `T` appears in a type.
569#[derive(Debug, Clone)]
570enum GenericUsage {
571    /// `T` does not appear in this type.
572    Absent,
573    /// `T` appears bare (possibly wrapped in `Option` or `Result`).
574    Bare,
575    /// `T` appears inside a container type (e.g. `DatumList<'a, T>`, `Array<'a, T>`).
576    /// The stored `syn::TypePath` is the container with `T` erased to `Datum<'a>`.
577    InContainer(syn::TypePath),
578}
579
580impl PartialEq for GenericUsage {
581    fn eq(&self, other: &Self) -> bool {
582        match (self, other) {
583            (GenericUsage::Absent, GenericUsage::Absent) => true,
584            (GenericUsage::Bare, GenericUsage::Bare) => true,
585            (GenericUsage::InContainer(a), GenericUsage::InContainer(b)) => {
586                container_idents_match(a, b)
587            }
588            _ => false,
589        }
590    }
591}
592
593impl Eq for GenericUsage {}
594
595/// Compare two container type paths by their ident segments (ignoring lifetimes
596/// and generic arguments). Two containers are "same" if their path idents match.
597///
598/// This is safe because after erasure all container types have the same generic
599/// arity (lifetimes + `Datum<'a>`), so ident equality implies structural equality.
600fn container_idents_match(a: &syn::TypePath, b: &syn::TypePath) -> bool {
601    let a_idents: Vec<_> = a.path.segments.iter().map(|s| &s.ident).collect();
602    let b_idents: Vec<_> = b.path.segments.iter().map(|s| &s.ident).collect();
603    a_idents == b_idents
604}
605
606/// Classifies how a generic type parameter appears in a type.
607///
608/// Strips `Option<...>`, `Result<..., E>`, and `ExcludeNull<...>` wrappers before
609/// inspecting the inner type. Any generic type wrapping `T` that isn't `Option`,
610/// `Result`, or `ExcludeNull` is treated as a container. If the container doesn't
611/// implement `SqlContainerType`, the generated code won't compile (a clear error).
612fn classify_generic_usage(ty: &syn::Type, generic_name: &Ident) -> GenericUsage {
613    match ty {
614        syn::Type::Path(type_path) => {
615            if type_path.path.is_ident(generic_name) {
616                return GenericUsage::Bare;
617            }
618            if let Some(last) = type_path.path.segments.last() {
619                let ident_str = last.ident.to_string();
620                // Unwrap Option, Result, ExcludeNull wrappers
621                if ident_str == "Option" || ident_str == "Result" || ident_str == "ExcludeNull" {
622                    if let syn::PathArguments::AngleBracketed(args) = &last.arguments {
623                        if let Some(syn::GenericArgument::Type(inner)) = args.args.first() {
624                            return classify_generic_usage(inner, generic_name);
625                        }
626                    }
627                }
628                // Check if any angle-bracketed arg contains the generic param.
629                // If so, treat this type as a container.
630                if let syn::PathArguments::AngleBracketed(args) = &last.arguments {
631                    let has_generic_arg = args.args.iter().any(|arg| {
632                        if let syn::GenericArgument::Type(inner) = arg {
633                            type_contains_ident(inner, generic_name)
634                        } else {
635                            false
636                        }
637                    });
638                    if has_generic_arg {
639                        // Build the erased container type path (T → Datum<'a>).
640                        let erased = erase_generic_param(ty, generic_name);
641                        if let syn::Type::Path(erased_path) = erased {
642                            return GenericUsage::InContainer(erased_path);
643                        }
644                    }
645                    // Recurse into args for nested containers
646                    // (e.g., Option<DatumList<'a, T>> was already handled by
647                    // the Option unwrapping above, but handle other nestings)
648                    for arg in &args.args {
649                        if let syn::GenericArgument::Type(inner) = arg {
650                            let inner_usage = classify_generic_usage(inner, generic_name);
651                            if inner_usage != GenericUsage::Absent {
652                                return inner_usage;
653                            }
654                        }
655                    }
656                }
657            }
658            GenericUsage::Absent
659        }
660        syn::Type::Reference(r) => classify_generic_usage(&r.elem, generic_name),
661        syn::Type::Tuple(t) => {
662            // Prefer container usages over bare. For example, `(T, DatumList<'_, T>)`
663            // should classify as `InDatumList`, not `Bare`.
664            let mut best = GenericUsage::Absent;
665            for elem in &t.elems {
666                let usage = classify_generic_usage(elem, generic_name);
667                match (&best, &usage) {
668                    (GenericUsage::Absent, _) => best = usage,
669                    (GenericUsage::Bare, u) if *u != GenericUsage::Absent => best = usage.clone(),
670                    _ => {
671                        if usage != GenericUsage::Absent && usage != best {
672                            // Conflicting container usages — cannot resolve.
673                            return GenericUsage::Bare;
674                        }
675                    }
676                }
677            }
678            best
679        }
680        _ => GenericUsage::Absent,
681    }
682}
683
684/// Returns whether a type syntactically contains an identifier.
685fn type_contains_ident(ty: &syn::Type, ident: &Ident) -> bool {
686    match ty {
687        syn::Type::Path(type_path) => {
688            if type_path.path.is_ident(ident) {
689                return true;
690            }
691            if let Some(last) = type_path.path.segments.last() {
692                if let syn::PathArguments::AngleBracketed(args) = &last.arguments {
693                    return args.args.iter().any(|arg| {
694                        if let syn::GenericArgument::Type(inner) = arg {
695                            type_contains_ident(inner, ident)
696                        } else {
697                            false
698                        }
699                    });
700                }
701            }
702            false
703        }
704        syn::Type::Reference(r) => type_contains_ident(&r.elem, ident),
705        syn::Type::Tuple(t) => t.elems.iter().any(|e| type_contains_ident(e, ident)),
706        _ => false,
707    }
708}
709
710/// Returns whether the outermost wrapper of a type is `Option`.
711fn is_option_wrapped(ty: &syn::Type) -> bool {
712    if let syn::Type::Path(type_path) = ty {
713        if let Some(last) = type_path.path.segments.last() {
714            return last.ident == "Option";
715        }
716    }
717    false
718}
719
720/// Derives an `output_type_expr` TokenStream from the structural relationship
721/// between input types and the output type, based on where generic parameters appear.
722///
723/// Finds the first generic parameter that appears in the output type, then looks for
724/// the first input parameter containing that generic in a container type to determine
725/// the unwrap strategy.
726///
727/// `is_unary` controls whether the generated expression uses `input_type`
728/// (singular, for unary functions) or `input_types[i]` (indexed, for binary/variadic).
729///
730/// Returns `None` if no generic parameter appears in the output type.
731fn derive_output_type_for_generics(
732    input_types: &[syn::Type],
733    output_ty: &syn::Type,
734    generic_names: &[Ident],
735    is_unary: bool,
736) -> darling::Result<Option<TokenStream>> {
737    // Find the first generic param that appears in the output.
738    let generic_name = match generic_names
739        .iter()
740        .find(|gn| classify_generic_usage(output_ty, gn) != GenericUsage::Absent)
741    {
742        Some(gn) => gn,
743        None => return Ok(None),
744    };
745    derive_output_type_for_generic(input_types, output_ty, generic_name, is_unary)
746}
747
748/// Derives an `output_type_expr` for a single generic parameter.
749///
750/// Uses `SqlContainerType` trait calls instead of matching on specific container
751/// type names. The generated code calls `<Container as SqlContainerType>::unwrap_element_type()`
752/// and `wrap_element_type()`, letting Rust's type system resolve the correct behavior.
753fn derive_output_type_for_generic(
754    input_types: &[syn::Type],
755    output_ty: &syn::Type,
756    generic_name: &Ident,
757    is_unary: bool,
758) -> darling::Result<Option<TokenStream>> {
759    let output_usage = classify_generic_usage(output_ty, generic_name);
760    if output_usage == GenericUsage::Absent {
761        return Ok(None);
762    }
763
764    let nullable = is_option_wrapped(output_ty);
765
766    // Find the first input parameter that has T in a container.
767    // Prefer container inputs over bare inputs.
768    let mut container_input: Option<(usize, GenericUsage)> = None;
769    for (i, ty) in input_types.iter().enumerate() {
770        let usage = classify_generic_usage(ty, generic_name);
771        match &usage {
772            GenericUsage::InContainer(_) => {
773                container_input = Some((i, usage));
774                break;
775            }
776            GenericUsage::Bare => {
777                // Bare T in input — not a container, keep looking for a container.
778                if container_input.is_none() {
779                    container_input = Some((i, usage));
780                }
781            }
782            GenericUsage::Absent => {}
783        }
784    }
785
786    let (pos, source_usage) = container_input.ok_or_else(|| {
787        darling::Error::custom(
788            "generic parameter T appears in the output type but not in any input type",
789        )
790    })?;
791
792    // Generate the base expression to access the input type.
793    let input_access = if is_unary {
794        quote! { input_type }
795    } else {
796        let pos_lit = syn::Index::from(pos);
797        quote! { input_types[#pos_lit] }
798    };
799
800    // For multi-input functions, generate soft assertions that all inputs
801    // carrying T agree on the SQL element type. This catches bugs in the
802    // planner's overload resolution or cast insertion.
803    let consistency_checks = if !is_unary {
804        let mut checks = Vec::new();
805        for (i, ty) in input_types.iter().enumerate() {
806            if i == pos {
807                continue;
808            }
809            let usage = classify_generic_usage(ty, generic_name);
810            if usage == GenericUsage::Absent {
811                continue;
812            }
813            let primary_elem = element_type_expr(&input_access, &source_usage);
814            let i_lit = syn::Index::from(i);
815            let other_access = quote! { input_types[#i_lit] };
816            let other_elem = element_type_expr(&other_access, &usage);
817            let generic_str = generic_name.to_string();
818            checks.push(quote! {
819                mz_ore::soft_assert_or_log!(
820                    #primary_elem.base_eq(#other_elem),
821                    "auto-derived sqlfunc output type inference found inconsistent \
822                     SQL types for generic {} across inputs: {:?} vs {:?}; \
823                     this indicates a bug in polymorphic coercion, builtin \
824                     declaration, or sqlfunc inference",
825                    #generic_str,
826                    #primary_elem,
827                    #other_elem,
828                );
829            });
830        }
831        quote! { #(#checks)* }
832    } else {
833        quote! {}
834    };
835
836    // Now generate the output_type_expr based on the combination of
837    // source container and output usage.
838    let expr = match (&output_usage, &source_usage) {
839        // Output is bare T, source is a container → unwrap element type via trait.
840        (GenericUsage::Bare, GenericUsage::InContainer(in_container)) => {
841            let in_c = elide_lifetimes(in_container);
842            quote! {
843                {
844                    #consistency_checks
845                    <#in_c as mz_repr::SqlContainerType>::unwrap_element_type(
846                        &#input_access.scalar_type
847                    ).clone().nullable(#nullable)
848                }
849            }
850        }
851        // Output is bare T, source is bare T → forward input type directly.
852        (GenericUsage::Bare, GenericUsage::Bare) => {
853            quote! {
854                {
855                    #consistency_checks
856                    #input_access.scalar_type.clone().nullable(#nullable)
857                }
858            }
859        }
860        // Output is a container, source is a container (same or different) →
861        // unwrap from input container, wrap into output container via traits.
862        (GenericUsage::InContainer(out_container), GenericUsage::InContainer(in_container)) => {
863            let out_c = elide_lifetimes(out_container);
864            let in_c = elide_lifetimes(in_container);
865            quote! {
866                {
867                    #consistency_checks
868                    <#out_c as mz_repr::SqlContainerType>::wrap_element_type(
869                        <#in_c as mz_repr::SqlContainerType>::unwrap_element_type(
870                            &#input_access.scalar_type
871                        ).clone()
872                    ).nullable(#nullable)
873                }
874            }
875        }
876        // Other cases — user must provide explicit output_type_expr.
877        _ => {
878            return Err(darling::Error::custom(format!(
879                "cannot auto-derive output_type_expr: output uses T as {:?} but \
880                 the first T-containing input uses T as {:?}",
881                output_usage, source_usage
882            )));
883        }
884    };
885
886    Ok(Some(expr))
887}
888
889/// Generates a token stream that extracts the T-level SQL type from an input
890/// access expression, based on how T is used in that input.
891fn element_type_expr(input_access: &TokenStream, usage: &GenericUsage) -> TokenStream {
892    match usage {
893        GenericUsage::Bare => {
894            quote! { &#input_access.scalar_type }
895        }
896        GenericUsage::InContainer(container) => {
897            let c = elide_lifetimes(container);
898            quote! {
899                <#c as mz_repr::SqlContainerType>::unwrap_element_type(
900                    &#input_access.scalar_type
901                )
902            }
903        }
904        GenericUsage::Absent => unreachable!("element_type_expr called with Absent usage"),
905    }
906}
907
908/// Replaces all lifetime parameters in a `syn::TypePath` with `'_`.
909///
910/// Used for container type paths in turbofish position (e.g.
911/// `<DatumList<'_, Datum<'_>> as SqlContainerType>::...`).
912/// The `output_sql_type` method's `&self` provides an implicit lifetime
913/// that the compiler can infer through `'_`.
914fn elide_lifetimes(tp: &syn::TypePath) -> syn::TypePath {
915    let mut tp = tp.clone();
916    for segment in &mut tp.path.segments {
917        if let syn::PathArguments::AngleBracketed(args) = &mut segment.arguments {
918            for arg in &mut args.args {
919                match arg {
920                    syn::GenericArgument::Lifetime(lt) => {
921                        *lt = Lifetime::new("'_", lt.span());
922                    }
923                    syn::GenericArgument::Type(ty) => {
924                        elide_lifetimes_in_type(ty);
925                    }
926                    _ => {}
927                }
928            }
929        }
930    }
931    tp
932}
933
934/// Recursively replaces all lifetime parameters in a `syn::Type` with `'_`.
935fn elide_lifetimes_in_type(ty: &mut syn::Type) {
936    match ty {
937        syn::Type::Path(tp) => {
938            *tp = elide_lifetimes(tp);
939        }
940        syn::Type::Reference(r) => {
941            if let Some(lt) = &mut r.lifetime {
942                *lt = Lifetime::new("'_", lt.span());
943            }
944            elide_lifetimes_in_type(&mut r.elem);
945        }
946        syn::Type::Tuple(t) => {
947            for elem in &mut t.elems {
948                elide_lifetimes_in_type(elem);
949            }
950        }
951        _ => {}
952    }
953}
954
955/// Replaces occurrences of a generic type parameter with `Datum<'a>` in a type.
956///
957/// Used to convert types from the user's generic function signature into concrete
958/// types for the generated trait impl's associated types, where `T` is not in scope.
959fn erase_generic_param(ty: &syn::Type, generic_name: &Ident) -> syn::Type {
960    match ty {
961        syn::Type::Path(type_path) => {
962            if type_path.path.is_ident(generic_name) {
963                return syn::parse_quote!(Datum<'a>);
964            }
965            let mut type_path = type_path.clone();
966            for segment in &mut type_path.path.segments {
967                if let syn::PathArguments::AngleBracketed(args) = &mut segment.arguments {
968                    for arg in &mut args.args {
969                        if let syn::GenericArgument::Type(inner) = arg {
970                            *inner = erase_generic_param(inner, generic_name);
971                        }
972                    }
973                }
974            }
975            syn::Type::Path(type_path)
976        }
977        syn::Type::Reference(r) => {
978            let elem = Box::new(erase_generic_param(&r.elem, generic_name));
979            syn::Type::Reference(syn::TypeReference { elem, ..r.clone() })
980        }
981        syn::Type::Tuple(t) => {
982            let elems = t
983                .elems
984                .iter()
985                .map(|e| erase_generic_param(e, generic_name))
986                .collect();
987            syn::Type::Tuple(syn::TypeTuple { elems, ..t.clone() })
988        }
989        _ => ty.clone(),
990    }
991}
992
993/// Erases all generic type parameters from a type, replacing each with `Datum<'a>`.
994fn erase_all_generic_params(ty: &syn::Type, generic_names: &[Ident]) -> syn::Type {
995    let mut ty = ty.clone();
996    for gn in generic_names {
997        ty = erase_generic_param(&ty, gn);
998    }
999    ty
1000}
1001
1002/// Determines the argument type of the nth argument of the function.
1003///
1004/// Adds a lifetime `'a` to the argument type if it is a reference type.
1005///
1006/// Panics if the function has fewer than `nth` arguments. Returns an error if
1007/// the parameter is a `self` receiver.
1008fn arg_type(arg: &syn::ItemFn, nth: usize) -> Result<syn::Type, syn::Error> {
1009    match &arg.sig.inputs[nth] {
1010        syn::FnArg::Typed(pat) => {
1011            // Patch lifetimes to be 'a if reference
1012            if let syn::Type::Reference(r) = &*pat.ty {
1013                if r.lifetime.is_none() {
1014                    let ty = syn::Type::Reference(syn::TypeReference {
1015                        lifetime: Some(Lifetime::new("'a", r.span())),
1016                        ..r.clone()
1017                    });
1018                    return Ok(ty);
1019                }
1020            }
1021            Ok((*pat.ty).clone())
1022        }
1023        syn::FnArg::Receiver(_) => Err(syn::Error::new(
1024            arg.sig.inputs[nth].span(),
1025            "Unsupported argument type",
1026        )),
1027    }
1028}
1029
1030/// Recursively patches lifetimes in a type, adding `'a` to references without a lifetime
1031/// and recursing into generic arguments and tuples.
1032fn patch_lifetimes(ty: &syn::Type) -> syn::Type {
1033    match ty {
1034        syn::Type::Reference(r) => {
1035            let elem = Box::new(patch_lifetimes(&r.elem));
1036            if r.lifetime.is_none() {
1037                syn::Type::Reference(syn::TypeReference {
1038                    lifetime: Some(Lifetime::new("'a", r.span())),
1039                    elem,
1040                    ..r.clone()
1041                })
1042            } else {
1043                syn::Type::Reference(syn::TypeReference { elem, ..r.clone() })
1044            }
1045        }
1046        syn::Type::Tuple(t) => {
1047            let elems = t.elems.iter().map(patch_lifetimes).collect();
1048            syn::Type::Tuple(syn::TypeTuple { elems, ..t.clone() })
1049        }
1050        syn::Type::Path(p) => {
1051            let mut p = p.clone();
1052            for segment in &mut p.path.segments {
1053                if let syn::PathArguments::AngleBracketed(args) = &mut segment.arguments {
1054                    for arg in &mut args.args {
1055                        if let syn::GenericArgument::Type(ty) = arg {
1056                            *ty = patch_lifetimes(ty);
1057                        }
1058                    }
1059                }
1060            }
1061            syn::Type::Path(p)
1062        }
1063        _ => ty.clone(),
1064    }
1065}
1066
1067/// Determine the output type for a function. Returns an error if the function
1068/// does not return a value.
1069fn output_type(arg: &syn::ItemFn) -> Result<&syn::Type, syn::Error> {
1070    match &arg.sig.output {
1071        syn::ReturnType::Type(_, ty) => Ok(&*ty),
1072        syn::ReturnType::Default => Err(syn::Error::new(
1073            arg.sig.output.span(),
1074            "Function needs to return a value",
1075        )),
1076    }
1077}
1078
1079/// Produce a `EagerUnaryFunc` implementation.
1080fn unary_func(
1081    func: &syn::ItemFn,
1082    modifiers: Modifiers,
1083    attr: &TokenStream,
1084) -> darling::Result<TokenStream> {
1085    let fn_name = &func.sig.ident;
1086    let struct_name = camel_case(&func.sig.ident);
1087    let input_ty_raw = arg_type(func, 0)?;
1088    let output_ty_raw = output_type(func)?;
1089    let generic_params = find_generic_type_params(func);
1090    // Erase generic type params → Datum<'a> for use in the trait impl's associated types.
1091    let input_ty = erase_all_generic_params(&input_ty_raw, &generic_params);
1092    let output_ty = erase_all_generic_params(output_ty_raw, &generic_params);
1093    let Modifiers {
1094        is_monotone,
1095        sqlname,
1096        preserves_uniqueness,
1097        inverse,
1098        is_infix_op,
1099        output_type,
1100        mut output_type_expr,
1101        negate,
1102        could_error,
1103        propagates_nulls,
1104        mut introduces_nulls,
1105        is_associative,
1106        is_eliminable_cast,
1107        is_infinity_monotone: _,
1108        test: _,
1109    } = modifiers;
1110
1111    // If generic type parameters are present and no explicit output_type_expr,
1112    // auto-derive one from the structural relationship between input and output types.
1113    // Use raw (pre-erasure) types so we can see the generic parameters.
1114    if !generic_params.is_empty() {
1115        if output_type_expr.is_none() && output_type.is_none() {
1116            if let Some(derived) = derive_output_type_for_generics(
1117                std::slice::from_ref(&input_ty_raw),
1118                output_ty_raw,
1119                &generic_params,
1120                true,
1121            )? {
1122                output_type_expr = Some(syn::parse2(derived)?);
1123                if introduces_nulls.is_none() {
1124                    let nullable = is_option_wrapped(output_ty_raw);
1125                    introduces_nulls = Some(syn::parse_quote!(#nullable));
1126                }
1127            }
1128        }
1129    }
1130
1131    if is_infix_op.is_some() {
1132        return Err(darling::Error::unknown_field(
1133            "is_infix_op not supported for unary functions",
1134        ));
1135    }
1136    if output_type.is_some() && output_type_expr.is_some() {
1137        return Err(darling::Error::unknown_field(
1138            "output_type and output_type_expr cannot be used together",
1139        ));
1140    }
1141    if output_type_expr.is_some() && introduces_nulls.is_none() {
1142        return Err(darling::Error::unknown_field(
1143            "output_type_expr requires introduces_nulls",
1144        ));
1145    }
1146    if negate.is_some() {
1147        return Err(darling::Error::unknown_field(
1148            "negate not supported for unary functions",
1149        ));
1150    }
1151    if propagates_nulls.is_some() {
1152        return Err(darling::Error::unknown_field(
1153            "propagates_nulls not supported for unary functions",
1154        ));
1155    }
1156    if is_associative.is_some() {
1157        return Err(darling::Error::unknown_field(
1158            "is_associative not supported for unary functions",
1159        ));
1160    }
1161
1162    let preserves_uniqueness_fn = preserves_uniqueness.map(|preserves_uniqueness| {
1163        quote! {
1164            fn preserves_uniqueness(&self) -> bool {
1165                #preserves_uniqueness
1166            }
1167        }
1168    });
1169
1170    let inverse_fn = inverse.as_ref().map(|inverse| {
1171        quote! {
1172            fn inverse(&self) -> Option<crate::UnaryFunc> {
1173                #inverse
1174            }
1175        }
1176    });
1177
1178    let is_monotone_fn = is_monotone.map(|is_monotone| {
1179        quote! {
1180            fn is_monotone(&self) -> bool {
1181                #is_monotone
1182            }
1183        }
1184    });
1185
1186    let name = sqlname
1187        .as_ref()
1188        .map_or_else(|| quote! { stringify!(#fn_name) }, |name| quote! { #name });
1189
1190    let (mut output_type, mut introduces_nulls_fn) = if let Some(output_type) = output_type {
1191        let introduces_nulls_fn = quote! {
1192            fn introduces_nulls(&self) -> bool {
1193                <#output_type as ::mz_repr::OutputDatumType<'_, ()>>::nullable()
1194            }
1195        };
1196        let output_type = quote! { <#output_type>::as_column_type() };
1197        (output_type, Some(introduces_nulls_fn))
1198    } else {
1199        (quote! { Self::Output::as_column_type() }, None)
1200    };
1201
1202    if let Some(output_type_expr) = output_type_expr {
1203        output_type = quote! { #output_type_expr };
1204    }
1205
1206    if let Some(introduces_nulls) = introduces_nulls {
1207        introduces_nulls_fn = Some(quote! {
1208            fn introduces_nulls(&self) -> bool {
1209                #introduces_nulls
1210            }
1211        });
1212    }
1213
1214    let could_error_fn = could_error.map(|could_error| {
1215        quote! {
1216            fn could_error(&self) -> bool {
1217                #could_error
1218            }
1219        }
1220    });
1221
1222    let is_eliminable_cast_fn = is_eliminable_cast.map(|is_eliminable_cast| {
1223        quote! {
1224            fn is_eliminable_cast(&self) -> bool {
1225                #is_eliminable_cast
1226            }
1227        }
1228    });
1229
1230    let source = sqlfunc_source(
1231        attr,
1232        func,
1233        std::slice::from_ref(&input_ty_raw),
1234        output_ty_raw,
1235        std::slice::from_ref(&input_ty),
1236    );
1237
1238    let result = quote! {
1239        #[derive(
1240            Ord, PartialOrd, Clone,
1241            Debug, Eq, PartialEq, serde::Serialize,
1242            serde::Deserialize, Hash,
1243        )]
1244        #[cfg_attr(any(test, feature = "proptest"), derive(proptest_derive::Arbitrary))]
1245        pub struct #struct_name;
1246
1247        impl crate::func::EagerUnaryFunc for #struct_name {
1248            type Input<'a> = #input_ty;
1249            type Output<'a> = #output_ty;
1250
1251            fn call<'a>(&self, a: Self::Input<'a>) -> Self::Output<'a> {
1252                #fn_name(a)
1253            }
1254
1255            fn output_sql_type(
1256                &self,
1257                input_type: mz_repr::SqlColumnType
1258            ) -> mz_repr::SqlColumnType {
1259                use mz_repr::AsColumnType;
1260                let output = #output_type;
1261                let propagates_nulls = crate::func::EagerUnaryFunc::propagates_nulls(self);
1262                let nullable = output.nullable;
1263                // The output is nullable if it is nullable by itself or the input is nullable
1264                // and this function propagates nulls
1265                output.nullable(nullable || (propagates_nulls && input_type.nullable))
1266            }
1267
1268            #could_error_fn
1269            #introduces_nulls_fn
1270            #inverse_fn
1271            #is_monotone_fn
1272            #preserves_uniqueness_fn
1273            #is_eliminable_cast_fn
1274        }
1275
1276        impl std::fmt::Display for #struct_name {
1277            fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
1278                f.write_str(#name)
1279            }
1280        }
1281
1282        impl crate::func::FuncName for #struct_name {
1283            const NAME: &'static str = stringify!(#fn_name);
1284            #source
1285        }
1286
1287        #func
1288    };
1289    Ok(result)
1290}
1291
1292/// Produce a `EagerBinaryFunc` implementation.
1293fn binary_func(
1294    func: &syn::ItemFn,
1295    modifiers: Modifiers,
1296    arena: bool,
1297    attr: &TokenStream,
1298) -> darling::Result<TokenStream> {
1299    let fn_name = &func.sig.ident;
1300    let struct_name = camel_case(&func.sig.ident);
1301    let input1_ty_raw = arg_type(func, 0)?;
1302    let input2_ty_raw = arg_type(func, 1)?;
1303    let output_ty_raw = output_type(func)?;
1304    let generic_params = find_generic_type_params(func);
1305    // Erase generic type params → Datum<'a> for use in the trait impl's associated types.
1306    let input1_ty = erase_all_generic_params(&input1_ty_raw, &generic_params);
1307    let input2_ty = erase_all_generic_params(&input2_ty_raw, &generic_params);
1308    let output_ty = erase_all_generic_params(output_ty_raw, &generic_params);
1309
1310    let Modifiers {
1311        is_monotone,
1312        sqlname,
1313        preserves_uniqueness,
1314        inverse,
1315        is_infix_op,
1316        output_type,
1317        mut output_type_expr,
1318        negate,
1319        could_error,
1320        propagates_nulls,
1321        mut introduces_nulls,
1322        is_associative,
1323        is_eliminable_cast,
1324        is_infinity_monotone,
1325        test: _,
1326    } = modifiers;
1327
1328    // Auto-derive output_type_expr from generic parameters, if applicable.
1329    // Use raw (pre-erasure) types so we can see the generic parameters.
1330    if !generic_params.is_empty() {
1331        if output_type_expr.is_none() && output_type.is_none() {
1332            if let Some(derived) = derive_output_type_for_generics(
1333                &[input1_ty_raw.clone(), input2_ty_raw.clone()],
1334                output_ty_raw,
1335                &generic_params,
1336                false,
1337            )? {
1338                output_type_expr = Some(syn::parse2(derived)?);
1339                if introduces_nulls.is_none() {
1340                    let nullable = is_option_wrapped(output_ty_raw);
1341                    introduces_nulls = Some(syn::parse_quote!(#nullable));
1342                }
1343            }
1344        }
1345    }
1346
1347    if preserves_uniqueness.is_some() {
1348        return Err(darling::Error::unknown_field(
1349            "preserves_uniqueness not supported for binary functions",
1350        ));
1351    }
1352    if inverse.is_some() {
1353        return Err(darling::Error::unknown_field(
1354            "inverse not supported for binary functions",
1355        ));
1356    }
1357    if output_type.is_some() && output_type_expr.is_some() {
1358        return Err(darling::Error::unknown_field(
1359            "output_type and output_type_expr cannot be used together",
1360        ));
1361    }
1362    if output_type_expr.is_some() && introduces_nulls.is_none() {
1363        return Err(darling::Error::unknown_field(
1364            "output_type_expr requires introduces_nulls",
1365        ));
1366    }
1367    if is_associative.is_some() {
1368        return Err(darling::Error::unknown_field(
1369            "is_associative not supported for binary functions",
1370        ));
1371    }
1372    if is_eliminable_cast.is_some() {
1373        return Err(darling::Error::unknown_field(
1374            "is_eliminable_cast not supported for binary functions",
1375        ));
1376    }
1377
1378    let negate_fn = negate.map(|negate| {
1379        quote! {
1380            fn negate(&self) -> Option<crate::BinaryFunc> {
1381                #negate
1382            }
1383        }
1384    });
1385
1386    let is_monotone_fn = is_monotone.map(|is_monotone| {
1387        quote! {
1388            fn is_monotone(&self) -> (bool, bool) {
1389                #is_monotone
1390            }
1391        }
1392    });
1393
1394    let is_infinity_monotone_fn = is_infinity_monotone.map(|is_infinity_monotone| {
1395        quote! {
1396            fn is_infinity_monotone(&self) -> bool {
1397                #is_infinity_monotone
1398            }
1399        }
1400    });
1401
1402    let name = sqlname
1403        .as_ref()
1404        .map_or_else(|| quote! { stringify!(#fn_name) }, |name| quote! { #name });
1405
1406    let (mut output_type, mut introduces_nulls_fn) = if let Some(output_type) = output_type {
1407        let introduces_nulls_fn = quote! {
1408            fn introduces_nulls(&self) -> bool {
1409                <#output_type as ::mz_repr::OutputDatumType<'_, ()>>::nullable()
1410            }
1411        };
1412        let output_type = quote! { <#output_type>::as_column_type() };
1413        (output_type, Some(introduces_nulls_fn))
1414    } else {
1415        (quote! { Self::Output::as_column_type() }, None)
1416    };
1417
1418    if let Some(output_type_expr) = output_type_expr {
1419        output_type = quote! { #output_type_expr };
1420    }
1421
1422    if let Some(introduces_nulls) = introduces_nulls {
1423        introduces_nulls_fn = Some(quote! {
1424            fn introduces_nulls(&self) -> bool {
1425                #introduces_nulls
1426            }
1427        });
1428    }
1429
1430    let arena = if arena {
1431        quote! { , temp_storage }
1432    } else {
1433        quote! {}
1434    };
1435
1436    let could_error_fn = could_error.map(|could_error| {
1437        quote! {
1438            fn could_error(&self) -> bool {
1439                #could_error
1440            }
1441        }
1442    });
1443
1444    let is_infix_op_fn = is_infix_op.map(|is_infix_op| {
1445        quote! {
1446            fn is_infix_op(&self) -> bool {
1447                #is_infix_op
1448            }
1449        }
1450    });
1451
1452    let propagates_nulls_fn = propagates_nulls.map(|propagates_nulls| {
1453        quote! {
1454            fn propagates_nulls(&self) -> bool {
1455                #propagates_nulls
1456            }
1457        }
1458    });
1459
1460    // Per-position checks: for each non-nullable parameter, check if
1461    // the corresponding input column is nullable.
1462    let binary_non_nullable_checks =
1463        non_nullable_position_checks(&[input1_ty.clone(), input2_ty.clone()]);
1464
1465    let source = sqlfunc_source(
1466        attr,
1467        func,
1468        &[input1_ty_raw.clone(), input2_ty_raw.clone()],
1469        output_ty_raw,
1470        &[input1_ty.clone(), input2_ty.clone()],
1471    );
1472
1473    let result = quote! {
1474        #[derive(
1475            Ord, PartialOrd, Clone,
1476            Debug, Eq, PartialEq, serde::Serialize,
1477            serde::Deserialize, Hash,
1478        )]
1479        #[cfg_attr(any(test, feature = "proptest"), derive(proptest_derive::Arbitrary))]
1480        pub struct #struct_name;
1481
1482        impl crate::func::binary::EagerBinaryFunc for #struct_name {
1483            type Input<'a> = (#input1_ty, #input2_ty);
1484            type Output<'a> = #output_ty;
1485
1486            fn call<'a>(
1487                &self,
1488                (a, b): Self::Input<'a>,
1489                temp_storage: &'a mz_repr::RowArena
1490            ) -> Self::Output<'a> {
1491                #fn_name(a, b #arena)
1492            }
1493
1494            fn output_sql_type(
1495                &self,
1496                input_types: &[mz_repr::SqlColumnType],
1497            ) -> mz_repr::SqlColumnType {
1498                use mz_repr::AsColumnType;
1499                let output = #output_type;
1500                let propagates_nulls =
1501                    crate::func::binary::EagerBinaryFunc::propagates_nulls(self);
1502                let nullable = output.nullable;
1503                // The output is nullable if:
1504                // 1. The function introduces nulls (output.nullable), or
1505                // 2. A non-nullable parameter's input is nullable (will reject
1506                //    NULL at runtime via try_from_iter), or
1507                // 3. propagates_nulls is true and any input is nullable
1508                //    (optimizer short-circuits all-NULL inputs)
1509                let non_nullable_input_is_nullable =
1510                    false #(#binary_non_nullable_checks)*;
1511                let inputs_nullable = input_types.iter().any(|it| it.nullable);
1512                let is_null = nullable
1513                    || non_nullable_input_is_nullable
1514                    || (propagates_nulls && inputs_nullable);
1515                output.nullable(is_null)
1516            }
1517
1518            #could_error_fn
1519            #introduces_nulls_fn
1520            #is_infix_op_fn
1521            #is_monotone_fn
1522            #is_infinity_monotone_fn
1523            #negate_fn
1524            #propagates_nulls_fn
1525        }
1526
1527        impl std::fmt::Display for #struct_name {
1528            fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
1529                f.write_str(#name)
1530            }
1531        }
1532
1533        impl crate::func::FuncName for #struct_name {
1534            const NAME: &'static str = stringify!(#fn_name);
1535            #source
1536        }
1537
1538        #func
1539
1540    };
1541    Ok(result)
1542}
1543
1544/// Produce an `EagerVariadicFunc` implementation.
1545///
1546/// Two modes based on whether the function has a `&self` receiver:
1547/// * `&self` present: struct defined externally, generates method impl + trait impl + Display
1548/// * No `&self`: generates unit struct + trait impl + Display + preserves original function
1549fn variadic_func(
1550    func: &syn::ItemFn,
1551    modifiers: Modifiers,
1552    struct_ty: Option<syn::Path>,
1553    arena: bool,
1554    has_self: bool,
1555    attr: &TokenStream,
1556) -> darling::Result<TokenStream> {
1557    let fn_name = &func.sig.ident;
1558    let output_ty_raw = output_type(func)?;
1559    let generic_params = find_generic_type_params(func);
1560    let output_ty = erase_all_generic_params(output_ty_raw, &generic_params);
1561    let struct_name = struct_ty
1562        .as_ref()
1563        .and_then(|ty| ty.segments.last())
1564        .map_or_else(|| camel_case(fn_name), |seg| seg.ident.clone());
1565
1566    let Modifiers {
1567        is_monotone,
1568        sqlname,
1569        preserves_uniqueness,
1570        inverse,
1571        is_infix_op,
1572        output_type,
1573        mut output_type_expr,
1574        negate,
1575        could_error,
1576        propagates_nulls,
1577        mut introduces_nulls,
1578        is_associative,
1579        is_eliminable_cast,
1580        is_infinity_monotone: _,
1581        test: _,
1582    } = modifiers;
1583
1584    // Reject modifiers that don't apply to variadic functions.
1585    if preserves_uniqueness.is_some() {
1586        return Err(darling::Error::unknown_field(
1587            "preserves_uniqueness not supported for variadic functions",
1588        ));
1589    }
1590    if inverse.is_some() {
1591        return Err(darling::Error::unknown_field(
1592            "inverse not supported for variadic functions",
1593        ));
1594    }
1595    if negate.is_some() {
1596        return Err(darling::Error::unknown_field(
1597            "negate not supported for variadic functions",
1598        ));
1599    }
1600    if is_eliminable_cast.is_some() {
1601        return Err(darling::Error::unknown_field(
1602            "is_eliminable_cast not supported for variadic functions",
1603        ));
1604    }
1605    if output_type.is_some() && output_type_expr.is_some() {
1606        return Err(darling::Error::unknown_field(
1607            "output_type and output_type_expr cannot be used together",
1608        ));
1609    }
1610    if output_type_expr.is_some() && introduces_nulls.is_none() {
1611        return Err(darling::Error::unknown_field(
1612            "output_type_expr requires introduces_nulls",
1613        ));
1614    }
1615
1616    // Collect input parameters (skip &self, skip &RowArena).
1617    let start = if has_self { 1 } else { 0 };
1618    let end = if arena {
1619        func.sig.inputs.len() - 1
1620    } else {
1621        func.sig.inputs.len()
1622    };
1623    let input_params: Vec<&syn::FnArg> = func
1624        .sig
1625        .inputs
1626        .iter()
1627        .skip(start)
1628        .take(end - start)
1629        .collect();
1630
1631    if input_params.is_empty() {
1632        return Err(darling::Error::custom(
1633            "variadic function must have at least one input parameter",
1634        ));
1635    }
1636
1637    // Extract parameter names and types.
1638    let mut param_names = Vec::new();
1639    let mut param_types = Vec::new();
1640    for param in &input_params {
1641        match param {
1642            syn::FnArg::Typed(pat) => {
1643                if let syn::Pat::Ident(ident) = &*pat.pat {
1644                    param_names.push(ident.ident.clone());
1645                } else {
1646                    return Err(
1647                        darling::Error::custom("unsupported parameter pattern").with_span(&pat.pat)
1648                    );
1649                }
1650                param_types.push(patch_lifetimes(&pat.ty));
1651            }
1652            syn::FnArg::Receiver(_) => {
1653                return Err(darling::Error::custom("unexpected self parameter"));
1654            }
1655        }
1656    }
1657
1658    // Auto-derive output_type_expr from generic parameters, if applicable.
1659    // Use raw (pre-erasure) types so we can see the generic parameters.
1660    if !generic_params.is_empty() {
1661        if output_type_expr.is_none() && output_type.is_none() {
1662            if let Some(derived) = derive_output_type_for_generics(
1663                &param_types,
1664                output_ty_raw,
1665                &generic_params,
1666                false,
1667            )? {
1668                output_type_expr = Some(syn::parse2(derived)?);
1669                if introduces_nulls.is_none() {
1670                    let nullable = is_option_wrapped(output_ty_raw);
1671                    introduces_nulls = Some(syn::parse_quote!(#nullable));
1672                }
1673            }
1674        }
1675    }
1676
1677    let param_types_raw = param_types.clone();
1678    // Erase generic type params → Datum<'a> in param types for the trait impl's associated types.
1679    for ty in &mut param_types {
1680        *ty = erase_all_generic_params(ty, &generic_params);
1681    }
1682
1683    // Build input type: single param = bare type, multiple = tuple.
1684    let input_type: syn::Type = if param_types.len() == 1 {
1685        param_types[0].clone()
1686    } else {
1687        syn::parse_quote! { (#(#param_types),*) }
1688    };
1689
1690    // Build destructure pattern for call.
1691    let destructure = if param_names.len() == 1 {
1692        let name = &param_names[0];
1693        quote! { #name }
1694    } else {
1695        quote! { (#(#param_names),*) }
1696    };
1697
1698    let arena_arg = if arena {
1699        quote! { , temp_storage }
1700    } else {
1701        quote! {}
1702    };
1703
1704    let call_expr = if has_self {
1705        quote! { self.#fn_name(#(#param_names),* #arena_arg) }
1706    } else {
1707        quote! { #fn_name(#(#param_names),* #arena_arg) }
1708    };
1709
1710    // Build modifier functions.
1711    let name = sqlname
1712        .as_ref()
1713        .map_or_else(|| quote! { stringify!(#fn_name) }, |name| quote! { #name });
1714
1715    let (mut output_type_code, mut introduces_nulls_fn) = if let Some(output_type) = output_type {
1716        let introduces_nulls_fn = quote! {
1717            fn introduces_nulls(&self) -> bool {
1718                <#output_type as ::mz_repr::OutputDatumType<'_, ()>>::nullable()
1719            }
1720        };
1721        let output_type_code = quote! { <#output_type>::as_column_type() };
1722        (output_type_code, Some(introduces_nulls_fn))
1723    } else {
1724        (quote! { Self::Output::as_column_type() }, None)
1725    };
1726
1727    if let Some(output_type_expr) = output_type_expr {
1728        output_type_code = quote! { #output_type_expr };
1729    }
1730
1731    if let Some(introduces_nulls) = introduces_nulls {
1732        introduces_nulls_fn = Some(quote! {
1733            fn introduces_nulls(&self) -> bool {
1734                #introduces_nulls
1735            }
1736        });
1737    }
1738
1739    let could_error_fn = could_error.map(|could_error| {
1740        quote! {
1741            fn could_error(&self) -> bool {
1742                #could_error
1743            }
1744        }
1745    });
1746
1747    let is_monotone_fn = is_monotone.map(|is_monotone| {
1748        quote! {
1749            fn is_monotone(&self) -> bool {
1750                #is_monotone
1751            }
1752        }
1753    });
1754
1755    let is_associative_fn = is_associative.map(|is_associative| {
1756        quote! {
1757            fn is_associative(&self) -> bool {
1758                #is_associative
1759            }
1760        }
1761    });
1762
1763    let is_infix_op_fn = is_infix_op.map(|is_infix_op| {
1764        quote! {
1765            fn is_infix_op(&self) -> bool {
1766                #is_infix_op
1767            }
1768        }
1769    });
1770
1771    let propagates_nulls_fn = propagates_nulls.map(|propagates_nulls| {
1772        quote! {
1773            fn propagates_nulls(&self) -> bool {
1774                #propagates_nulls
1775            }
1776        }
1777    });
1778
1779    // Per-position checks: for each non-nullable parameter, check if
1780    // the corresponding input column is nullable.
1781    let non_nullable_checks = non_nullable_position_checks(&param_types);
1782
1783    let trait_impl = quote! {
1784        impl crate::func::variadic::EagerVariadicFunc for #struct_name {
1785            type Input<'a> = #input_type;
1786            type Output<'a> = #output_ty;
1787
1788            fn call<'a>(
1789                &self,
1790                #destructure: Self::Input<'a>,
1791                temp_storage: &'a mz_repr::RowArena,
1792            ) -> Self::Output<'a> {
1793                #call_expr
1794            }
1795
1796            fn output_type(
1797                &self,
1798                input_types: &[mz_repr::SqlColumnType],
1799            ) -> mz_repr::SqlColumnType {
1800                use mz_repr::AsColumnType;
1801                let output = #output_type_code;
1802                let propagates_nulls =
1803                    crate::func::variadic::EagerVariadicFunc::propagates_nulls(self);
1804                let nullable = output.nullable;
1805                // The output is nullable if:
1806                // 1. The function introduces nulls (output.nullable), or
1807                // 2. A non-nullable parameter's input is nullable (will reject
1808                //    NULL at runtime via try_from_iter), or
1809                // 3. propagates_nulls is true and any input is nullable
1810                //    (optimizer short-circuits all-NULL inputs)
1811                let non_nullable_input_is_nullable =
1812                    false #(#non_nullable_checks)*;
1813                let inputs_nullable = input_types.iter().any(|it| it.nullable);
1814                output.nullable(
1815                    nullable
1816                    || non_nullable_input_is_nullable
1817                    || (propagates_nulls && inputs_nullable)
1818                )
1819            }
1820
1821            #could_error_fn
1822            #introduces_nulls_fn
1823            #is_infix_op_fn
1824            #is_monotone_fn
1825            #is_associative_fn
1826            #propagates_nulls_fn
1827        }
1828    };
1829
1830    let display_impl = quote! {
1831        impl std::fmt::Display for #struct_name {
1832            fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
1833                f.write_str(#name)
1834            }
1835        }
1836    };
1837
1838    let source = sqlfunc_source(attr, func, &param_types_raw, output_ty_raw, &param_types);
1839    let funcname_impl = quote! {
1840        impl crate::func::FuncName for #struct_name {
1841            const NAME: &'static str = stringify!(#fn_name);
1842            #source
1843        }
1844    };
1845
1846    let result = if has_self {
1847        // External struct: generate method impl + trait impl + Display.
1848        quote! {
1849            impl #struct_name {
1850                #func
1851            }
1852            #trait_impl
1853            #display_impl
1854            #funcname_impl
1855        }
1856    } else {
1857        // Unit struct: generate struct + trait impl + Display + original function.
1858        quote! {
1859            #[derive(
1860                Ord, PartialOrd, Clone,
1861                Debug, Eq, PartialEq, serde::Serialize,
1862                serde::Deserialize, Hash,
1863            )]
1864            #[cfg_attr(any(test, feature = "proptest"), derive(proptest_derive::Arbitrary))]
1865            pub struct #struct_name;
1866
1867            #trait_impl
1868            #display_impl
1869            #funcname_impl
1870
1871            #func
1872        }
1873    };
1874
1875    Ok(result)
1876}
1877
1878#[cfg(test)]
1879mod render_tests {
1880    use super::render_tokens;
1881
1882    #[mz_ore::test]
1883    fn signature_reads_like_formatted_code() {
1884        let sig: proc_macro2::TokenStream = syn::parse_quote! {
1885            fn f<'a>(mut a: &'a str, b: Option<i32>) -> Result<Cow<'a, str>, E>
1886        };
1887        assert_eq!(
1888            render_tokens(&sig),
1889            "fn f<'a>(mut a: &'a str, b: Option<i32>) -> Result<Cow<'a, str>, E>"
1890        );
1891        let attr: proc_macro2::TokenStream = syn::parse_quote! {
1892            is_monotone = "(true, true)", sqlname = "+", could_error = false
1893        };
1894        assert_eq!(
1895            render_tokens(&attr),
1896            "is_monotone = \"(true, true)\", sqlname = \"+\", could_error = false"
1897        );
1898    }
1899
1900    #[mz_ore::test]
1901    fn literal_contents_are_verbatim() {
1902        let a: proc_macro2::TokenStream = syn::parse_quote! { format!("({x})") };
1903        let b: proc_macro2::TokenStream = syn::parse_quote! { format!("( {x})") };
1904        assert_ne!(render_tokens(&a), render_tokens(&b));
1905        assert_eq!(render_tokens(&a), "format!(\"({x})\")");
1906    }
1907
1908    #[mz_ore::test]
1909    fn formatting_is_invisible() {
1910        let a: proc_macro2::TokenStream = "fn f ( a : i32 ) -> i32 { a + 1 }".parse().unwrap();
1911        let b: proc_macro2::TokenStream = "fn f(a:i32)->i32{a+1}".parse().unwrap();
1912        assert_eq!(render_tokens(&a), render_tokens(&b));
1913    }
1914}