use itertools::Itertools;
use proc_macro2::TokenStream;
use quote::{ToTokens, format_ident, quote};
use syn::{Data, DeriveInput, Result};
use crate::{codegen_utils, types};
pub fn generate(tokens: TokenStream) -> Result<TokenStream> {
let input: DeriveInput = syn::parse2(tokens)?;
let struct_name = &input.ident;
let generics = &input.generics;
let Data::Struct(struct_) = &input.data else {
return Err(syn::Error::new_spanned(input, "expect struct"));
};
let fields = struct_
.fields
.iter()
.map(Field::parse)
.collect::<Result<Vec<Field>>>()?;
let fields0 = fields
.iter()
.map(|f| codegen_utils::field(&f.name, &f.type_));
let append_values = fields.iter().enumerate().map(|(i, f)| {
let field = &f.ident;
let append_value = codegen_utils::gen_append_value(&f.type_);
let append_null = codegen_utils::gen_append_null(&f.type_);
let builder_type = codegen_utils::builder_type(&f.type_);
match f.option {
false => quote! {{
let builder = builder.field_builder::<#builder_type>(#i).unwrap();
let v = self.#field;
#append_value
}},
true => quote! {{
let builder = builder.field_builder::<#builder_type>(#i).unwrap();
match self.#field {
Some(v) => #append_value,
None => #append_null,
}
}},
}
});
let append_nulls = fields.iter().enumerate().map(|(i, f)| {
let builder_type = codegen_utils::builder_type(&f.type_);
let append_null = codegen_utils::gen_append_null(&f.type_);
quote! {{
let builder = builder.field_builder::<#builder_type>(#i).unwrap();
#append_null
}}
});
let static_name = format_ident!("{}_METADATA", struct_name.to_string().to_uppercase());
let export_name = format!(
"arrowudt_{}",
codegen_utils::base64_encode(&format!(
"{}={}",
struct_name,
fields
.iter()
.map(|f| format!("{}:{}", f.name, f.type_))
.join(",")
))
);
Ok(quote! {
#[unsafe(export_name = #export_name)]
static #static_name: () = ();
impl #generics ::arrow_udf::types::StructType for #struct_name #generics {
fn fields() -> ::arrow_udf::codegen::arrow_schema::Fields {
use ::arrow_udf::codegen::arrow_schema::{self, Field, TimeUnit, IntervalUnit};
vec![#(#fields0),*].into()
}
fn append_to(self, builder: &mut ::arrow_udf::codegen::arrow_array::builder::StructBuilder) {
use ::arrow_udf::codegen::arrow_array::builder::*;
#(#append_values)*
builder.append(true);
}
fn append_null(builder: &mut ::arrow_udf::codegen::arrow_array::builder::StructBuilder) {
use ::arrow_udf::codegen::arrow_array::builder::*;
#(#append_nulls)*
builder.append_null();
}
}
})
}
#[derive(Debug)]
struct Field {
ident: syn::Ident,
name: String,
type_: String,
option: bool,
}
impl Field {
fn parse(field: &syn::Field) -> Result<Self> {
let ident = field
.ident
.clone()
.ok_or_else(|| syn::Error::new_spanned(field, "expect field name"))?;
let mut name = ident.to_string();
if name.starts_with("r#") {
name = name[2..].to_string();
}
let ty = &field.ty;
let (option, ty) = match strip_outer_type(ty, "Option") {
Some(ty) => (true, ty),
None => (false, ty),
};
let (list, ty) = match strip_outer_type(ty, "Vec") {
Some(ty) if ty.to_token_stream().to_string() != "u8" => (true, ty),
_ => (false, ty),
};
let mut type_ =
types::type_of(&ty.to_token_stream().to_string().replace(' ', "")).to_string();
if list {
type_ += "[]";
}
Ok(Self {
ident,
name,
type_,
option,
})
}
}
fn strip_outer_type<'a>(ty: &'a syn::Type, type_: &str) -> Option<&'a syn::Type> {
let syn::Type::Path(path) = ty else {
return None;
};
let seg = path.path.segments.last()?;
if seg.ident != type_ {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
let Some(syn::GenericArgument::Type(ty)) = args.args.first() else {
return None;
};
Some(ty)
}
#[cfg(test)]
mod tests {
use proc_macro2::TokenStream;
use syn::File;
fn pretty_print(output: TokenStream) -> String {
let output: File = syn::parse2(output).unwrap();
prettyplease::unparse(&output)
}
#[test]
fn test_struct_type() {
let code = include_str!("testdata/struct.input.rs");
let input: TokenStream = str::parse(code).unwrap();
let output = super::generate(input).unwrap();
let output = pretty_print(output);
let expected = expect_test::expect_file!["testdata/struct.output.rs"];
expected.assert_eq(&output);
}
}