1use 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#[derive(Debug, Default, darling::FromMeta)]
18pub(crate) struct Modifiers {
19 is_monotone: Option<Expr>,
22 is_infinity_monotone: Option<Expr>,
26 sqlname: Option<SqlName>,
28 preserves_uniqueness: Option<Expr>,
30 inverse: Option<Expr>,
32 negate: Option<Expr>,
34 is_infix_op: Option<Expr>,
37 output_type: Option<syn::Path>,
39 output_type_expr: Option<Expr>,
41 could_error: Option<Expr>,
44 propagates_nulls: Option<Expr>,
46 introduces_nulls: Option<Expr>,
48 is_associative: Option<Expr>,
50 is_eliminable_cast: Option<Expr>,
52 test: Option<bool>,
54}
55
56#[derive(Debug)]
59enum SqlName {
60 Literal(syn::Lit),
62 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 Expr::Group(mac) => Self::from_expr(&mac.expr),
87 _ => Err(darling::Error::unexpected_expr_type(expr)),
88 }
89 }
90}
91
92pub 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 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
138fn 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 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 TokenTree::Punct(p) if "<!".contains(p.as_char()) => matches!(last, Last::Ident),
173 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
219fn fnv1a64(bytes: &[u8]) -> u64 {
222 bytes.iter().fold(0xcbf29ce484222325, |hash, byte| {
223 (hash ^ u64::from(*byte)).wrapping_mul(0x100000001b3)
224 })
225}
226
227fn 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
283fn 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
296fn 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
307fn 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
328fn 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)] #[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
384fn 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
398enum Arity {
400 Nullary,
401 Unary { arena: bool },
402 Binary { arena: bool },
403 Variadic { arena: bool, has_self: bool },
404}
405
406fn 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
421fn 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 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
465fn 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 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
492fn 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
502fn 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
516fn 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
551fn 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#[derive(Debug, Clone)]
570enum GenericUsage {
571 Absent,
573 Bare,
575 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
595fn 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
606fn 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 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 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 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 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 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 return GenericUsage::Bare;
674 }
675 }
676 }
677 }
678 best
679 }
680 _ => GenericUsage::Absent,
681 }
682}
683
684fn 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
710fn 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
720fn 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 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
748fn 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 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 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 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 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 let expr = match (&output_usage, &source_usage) {
839 (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 (GenericUsage::Bare, GenericUsage::Bare) => {
853 quote! {
854 {
855 #consistency_checks
856 #input_access.scalar_type.clone().nullable(#nullable)
857 }
858 }
859 }
860 (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 _ => {
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
889fn 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
908fn 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
934fn 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
955fn 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
993fn 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
1002fn arg_type(arg: &syn::ItemFn, nth: usize) -> Result<syn::Type, syn::Error> {
1009 match &arg.sig.inputs[nth] {
1010 syn::FnArg::Typed(pat) => {
1011 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
1030fn 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
1067fn 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
1079fn 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 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_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 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
1292fn 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 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 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 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 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
1544fn 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 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 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 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 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 ¶m_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 for ty in &mut param_types {
1680 *ty = erase_all_generic_params(ty, &generic_params);
1681 }
1682
1683 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 let destructure = if param_names.len() == 1 {
1692 let name = ¶m_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 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 let non_nullable_checks = non_nullable_position_checks(¶m_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 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, ¶m_types_raw, output_ty_raw, ¶m_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 quote! {
1849 impl #struct_name {
1850 #func
1851 }
1852 #trait_impl
1853 #display_impl
1854 #funcname_impl
1855 }
1856 } else {
1857 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}