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#[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}