Skip to main content

arroyo_udf_macros/
lib.rs

1use arrow_schema::DataType;
2use arroyo_udf_common::parse::{is_vec_u8, ParsedUdf};
3use proc_macro2::{Span, TokenStream};
4use quote::{format_ident, quote};
5use syn::parse::{Parse, ParseStream};
6use syn::spanned::Spanned;
7use syn::{parse_quote, FnArg, ItemFn};
8
9fn data_type_to_arrow_type_token(data_type: &DataType) -> TokenStream {
10    match data_type {
11        DataType::Utf8 => quote!(GenericStringType<i32>),
12        DataType::Boolean => quote!(BooleanType),
13        DataType::Int16 => quote!(Int16Type),
14        DataType::Int32 => quote!(Int32Type),
15        DataType::Int64 => quote!(Int64Type),
16        DataType::Int8 => quote!(Int8Type),
17        DataType::UInt8 => quote!(UInt8Type),
18        DataType::UInt16 => quote!(UInt16Type),
19        DataType::UInt32 => quote!(UInt32Type),
20        DataType::UInt64 => quote!(UInt64Type),
21        DataType::Float32 => quote!(Float32Type),
22        DataType::Float64 => quote!(Float64Type),
23        DataType::Binary => quote!(GenericBinaryType<i32>),
24        DataType::List(f) => data_type_to_arrow_type_token(f.data_type()),
25        _ => panic!("Unsupported data type: {:?}", data_type),
26    }
27}
28
29struct ParsedFunction(ParsedUdf, ItemFn);
30
31impl Parse for ParsedFunction {
32    fn parse(input: ParseStream) -> syn::Result<Self> {
33        let function: ItemFn = input.parse()?;
34
35        if function.sig.asyncness.is_some() {
36            if let Some(vec) = function.sig.inputs.iter().find_map(|t| match t {
37                FnArg::Receiver(_) => None,
38                FnArg::Typed(t) => {
39                    if ParsedUdf::vec_inner_type(&t.ty).is_some() && !is_vec_u8(&t.ty) {
40                        Some(t.ty.span())
41                    } else {
42                        None
43                    }
44                }
45            }) {
46                return Err(syn::Error::new(
47                    vec.span(),
48                    "Async UDAFs are not supported (hint: remove the Vec<_> args)",
49                ));
50            }
51        }
52
53        Ok(ParsedFunction(
54            ParsedUdf::try_parse(&function)
55                .map_err(|e| syn::Error::new(Span::call_site(), e.to_string()))?,
56            function,
57        ))
58    }
59}
60
61#[proc_macro_attribute]
62pub fn udf(
63    _attr: proc_macro::TokenStream,
64    input: proc_macro::TokenStream,
65) -> proc_macro::TokenStream {
66    let parsed: ParsedFunction = match syn::parse(input) {
67        Ok(parsed) => parsed,
68        Err(e) => {
69            return e.to_compile_error().into();
70        }
71    };
72
73    let mangle = Some(quote! { #[no_mangle] });
74    let tokens = if parsed.0.udf_type.is_async() {
75        async_udf(parsed, mangle)
76    } else {
77        sync_udf(parsed, mangle)
78    };
79
80    (quote! {
81        #tokens
82    })
83    .into()
84}
85
86/// Used to generate a statically-linked UDF for testing
87#[proc_macro_attribute]
88pub fn local_udf(
89    attr: proc_macro::TokenStream,
90    input: proc_macro::TokenStream,
91) -> proc_macro::TokenStream {
92    let input_str = input.to_string();
93    let def = format!("#[udf({})]{}", attr, input_str);
94    let parsed: ParsedFunction = syn::parse(input).unwrap();
95    let name = parsed.0.name.clone();
96
97    let (tokens, interface) = if parsed.0.udf_type.is_async() {
98        let tokens = async_udf(parsed, None);
99        let interface = quote! {
100            arroyo_udf_host::UdfInterface::Async(std::sync::Arc::new(arroyo_udf_host::ContainerOrLocal::Local(
101                arroyo_udf_host::AsyncUdfDylibInterface::new(
102                    __start,
103                    __send,
104                    __drain_results,
105                    __stop_runtime,
106                ))))
107        };
108        (tokens, interface)
109    } else {
110        let tokens = sync_udf(parsed, None);
111        let interface = quote! {
112            arroyo_udf_host::UdfInterface::Sync(std::sync::Arc::new(arroyo_udf_host::ContainerOrLocal::Local(
113                arroyo_udf_host::UdfDylibInterface::new(__run))))
114        };
115        (tokens, interface)
116    };
117
118    (quote!(
119        #tokens
120
121        pub fn __local() -> arroyo_udf_host::LocalUdf {
122            let config = arroyo_udf_host::parse::ParsedUdf::try_parse(&syn::parse_str(#input_str).unwrap()).unwrap();
123
124            arroyo_udf_host::LocalUdf {
125                def: #def,
126                config: arroyo_udf_host::UdfDylib::new(
127                    #name.to_string(),
128                    datafusion::logical_expr::Signature::exact(
129                        config.args.into_iter().map(|a| a.data_type).collect(),
130                        datafusion::logical_expr::Volatility::Volatile),
131                    config.ret_type.data_type,
132                    #interface,
133                ),
134                is_aggregate: config.vec_arguments > 0,
135                is_async: config.udf_type.is_async(),
136            }
137        }
138    )).into()
139}
140
141fn arg_vars(parsed: &ParsedUdf) -> (Vec<TokenStream>, Vec<TokenStream>) {
142    parsed
143        .args
144        .iter()
145        .enumerate()
146        .map(|(i, arg_type)| {
147            let arrow_type = data_type_to_arrow_type_token(&arg_type.data_type);
148            let id = format_ident!("arg_{}", i);
149            let def = match &arg_type.data_type {
150                DataType::Utf8 => {
151                    quote!(let #id = arroyo_udf_plugin::arrow::array::StringArray::from(args.next().unwrap());)
152                }
153                DataType::Binary => {
154                    quote!(let #id = arroyo_udf_plugin::arrow::array::BinaryArray::from(args.next().unwrap());)
155                }
156                DataType::List(field) => {
157                    let filter = if !field.is_nullable() {
158                        quote!(.filter_map(|x| x))
159                    } else {
160                        quote!()
161                    };
162
163                    quote!(let #id = arroyo_udf_plugin::arrow::array::PrimitiveArray::<arroyo_udf_plugin::arrow::datatypes::#arrow_type>::from(
164                        args.next().unwrap()
165                    ).iter()#filter.collect();)
166                }
167                _ => {
168                    quote!(let #id = arroyo_udf_plugin::arrow::array::PrimitiveArray::<arroyo_udf_plugin::arrow::datatypes::#arrow_type>::from(args.next().unwrap());)
169                }
170            };
171
172
173            (def, quote!(#id))
174        })
175        .unzip()
176}
177
178fn sync_udf(parsed: ParsedFunction, mangle: Option<TokenStream>) -> TokenStream {
179    let (parsed, item) = (parsed.0, parsed.1);
180    let udf_name = format_ident!("{}", parsed.name);
181
182    let results_builder = match parsed.ret_type.data_type {
183        DataType::Utf8 => {
184            quote!(let mut results_builder = arroyo_udf_plugin::arrow::array::StringBuilder::with_capacity(batch_size, batch_size * 8);)
185        }
186        DataType::Binary => {
187            quote!(let mut results_builder = arroyo_udf_plugin::arrow::array::GenericByteBuilder::<arroyo_udf_plugin::arrow::array::types::GenericBinaryType<i32>>
188                ::with_capacity(batch_size, batch_size * 8);)
189        }
190        _ => {
191            let return_type = data_type_to_arrow_type_token(&parsed.ret_type.data_type);
192            quote!(let mut results_builder = arroyo_udf_plugin::arrow::array::PrimitiveBuilder::<arroyo_udf_plugin::arrow::datatypes::#return_type>::with_capacity(batch_size);)
193        }
194    };
195
196    let (defs, args) = arg_vars(&parsed);
197
198    let udaf = parsed
199        .args
200        .iter()
201        .any(|arg| matches!(arg.data_type, DataType::List(_)));
202
203    let unwrapping: Vec<_> = parsed
204        .args
205        .iter()
206        .enumerate()
207        .map(|(i, arg_type)| {
208            let id = format_ident!("arg_{}", i);
209
210            let append_none = match parsed.ret_type.data_type {
211                DataType::Utf8 => {
212                    quote!(results_builder.append_option(None::<String>);)
213                }
214                DataType::Binary => {
215                    quote!(results_builder.append_option(None::<Vec<u8>>);)
216                }
217                _ => quote!(results_builder.append_option(None);),
218            };
219
220            if arg_type.nullable {
221                quote!()
222            } else {
223                parse_quote! {
224                    let Some(#id) = #id else {
225                        #append_none
226                        continue;
227                    };
228                }
229            }
230        })
231        .collect();
232
233    let mut arg_destructure = quote!(arg_0);
234    let mut arg_zip = quote!(arg_0.iter());
235    for i in 1..args.len() {
236        let next_arg = format_ident!("arg_{}", i);
237        arg_zip = quote!(#arg_zip.zip(#next_arg.iter()));
238        arg_destructure = quote!((#arg_destructure, #next_arg))
239    }
240
241    let call = if parsed.ret_type.nullable {
242        quote!(results_builder.append_option(#udf_name(#(#args),*));)
243    } else {
244        quote!(results_builder.append_option(Some(#udf_name(#(#args),*)));)
245    };
246
247    let call_loop = if udaf {
248        quote! {
249            #call
250        }
251    } else {
252        quote! {
253            for #arg_destructure in #arg_zip {
254                #(#unwrapping;)*
255                #call
256            }
257        }
258    };
259
260    quote! {
261        #item
262
263        #mangle
264        pub extern "C-unwind" fn __run(args: arroyo_udf_plugin::FfiArrays) -> arroyo_udf_plugin::RunResult {
265            let args = args.into_vec();
266            let batch_size = args[0].len();
267
268            let result = std::panic::catch_unwind(|| {
269                let mut args = args.into_iter();
270                #results_builder
271
272                #(#defs;)*
273
274                #call_loop
275
276                arroyo_udf_plugin::arrow::array::Array::to_data(&results_builder.finish())
277            });
278
279
280            match result {
281                Ok(data) => {
282                    arroyo_udf_plugin::RunResult::Ok(arroyo_udf_plugin::FfiArraySchema::from_data(data))
283                }
284                Err(e) => {
285                    arroyo_udf_plugin::RunResult::Err
286                }
287            }
288        }
289    }
290}
291
292fn async_udf(parsed: ParsedFunction, mangle: Option<TokenStream>) -> TokenStream {
293    let (parsed, item) = (parsed.0, parsed.1);
294
295    let (defs, args) = arg_vars(&parsed);
296
297    let name = format_ident!("{}", parsed.name);
298    let call_args: Vec<_> = args
299        .iter()
300        .zip(parsed.args)
301        .map(|(arg, t)| {
302            if t.nullable {
303                quote!(if arroyo_udf_plugin::arrow::array::Array::is_null(&#arg, 0) { None } else { Some(#arg.value(0))})
304            } else {
305                quote!(#arg.value(0))                
306            }
307        })
308        .collect();
309
310    let datum = match parsed.ret_type.data_type {
311        DataType::Boolean => quote!(Bool),
312        DataType::Int32 => quote!(I32),
313        DataType::Int64 => quote!(I64),
314        DataType::UInt32 => quote!(U32),
315        DataType::UInt64 => quote!(U64),
316        DataType::Float32 => quote!(F32),
317        DataType::Float64 => quote!(F64),
318        DataType::Timestamp(_, _) => quote!(Timestamp),
319        DataType::Binary => quote!(Bytes),
320        DataType::Utf8 => quote!(String),
321        _ => panic!("unsupported return type {}", parsed.ret_type.data_type),
322    };
323
324    let wrap_return = if parsed.ret_type.nullable {
325        quote!(arroyo_udf_plugin::ArrowDatum::#datum(result))
326    } else {
327        quote!(arroyo_udf_plugin::ArrowDatum::#datum(Some(result)))
328    };
329
330    let wrapper = quote! {
331        async fn __wrapper(id: u64, timeout: std::time::Duration, args: Vec<arroyo_udf_plugin::arrow::array::ArrayData>) ->
332           (u64, Result<arroyo_udf_plugin::ArrowDatum, arroyo_udf_plugin::async_udf::tokio::time::error::Elapsed>) {
333            let mut args = args.into_iter();
334
335            #(#defs;)*
336
337            match arroyo_udf_plugin::async_udf::tokio::time::timeout(timeout, #name(#(#call_args, )*)).await {
338                Ok(result) => (id, Ok(#wrap_return)),
339                Err(e) => (id, Err(e)),
340            }
341        }
342    };
343
344    let results_builder = match parsed.ret_type.data_type {
345        DataType::Utf8 => quote!(arroyo_udf_plugin::arrow::array::StringBuilder::new()),
346        DataType::Binary => quote!(arroyo_udf_plugin::arrow::array::GenericByteBuilder::<
347            arroyo_udf_plugin::arrow::array::types::GenericBinaryType<i32>,
348        >::new()),
349        _ => {
350            let return_type = data_type_to_arrow_type_token(&parsed.ret_type.data_type);
351            quote!(arroyo_udf_plugin::arrow::array::PrimitiveBuilder::<arroyo_udf_plugin::arrow::datatypes::#return_type>::new())
352        }
353    };
354
355    let start = quote! {
356        #mangle
357        pub extern "C-unwind" fn __start(ordered: bool, timeout_micros: u64, allowed_in_flight: u32) -> arroyo_udf_plugin::async_udf::SendableFfiAsyncUdfHandle {
358            let (x, handle) = arroyo_udf_plugin::async_udf::AsyncUdf::new(
359                ordered, std::time::Duration::from_micros(timeout_micros), allowed_in_flight, Box::new(#results_builder), __wrapper
360            );
361
362            x.start();
363
364            arroyo_udf_plugin::async_udf::SendableFfiAsyncUdfHandle { ptr: handle.into_ffi() }
365        }
366    };
367
368    quote! {
369        #item
370
371        #wrapper
372
373        #start
374
375        #mangle
376        pub extern "C-unwind" fn __send(handle: arroyo_udf_plugin::async_udf::SendableFfiAsyncUdfHandle,
377            id: u64, arrays: arroyo_udf_plugin::FfiArrays) -> arroyo_udf_plugin::async_udf::async_ffi::FfiFuture<bool> {
378            use arroyo_udf_plugin::async_udf::async_ffi::FutureExt;
379            arroyo_udf_plugin::async_udf::send(handle, id, arrays).into_ffi()
380        }
381
382        #mangle
383        pub extern "C-unwind" fn __drain_results(handle: arroyo_udf_plugin::async_udf::SendableFfiAsyncUdfHandle) -> arroyo_udf_plugin::async_udf::DrainResult {
384            arroyo_udf_plugin::async_udf::drain_results(handle)
385        }
386
387        #mangle
388        pub extern "C-unwind" fn __stop_runtime(handle: arroyo_udf_plugin::async_udf::SendableFfiAsyncUdfHandle) {
389            arroyo_udf_plugin::async_udf::stop_runtime(handle);
390        }
391    }
392}