use crate::parse::ReductionArgs;
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::{quote, quote_spanned};
use syn::parse_macro_input;
use syn::spanned::Spanned;
fn create_reduction(
typeident: syn::Ident,
reduction: String,
op: proc_macro2::TokenStream,
array_types: &Vec<syn::Ident>,
) -> proc_macro2::TokenStream {
let lamellar = quote::format_ident!("__lamellar");
let am_data: syn::Path = syn::parse("lamellar::AmData".parse().unwrap()).unwrap();
let am: syn::Path = syn::parse("lamellar::am".parse().unwrap()).unwrap();
let reduction = quote::format_ident!("{:}", reduction);
let reduction_gen = quote::format_ident!("{:}_{:}_reduction_gen", typeident, reduction);
let reduction_id_gen = quote::format_ident!("{:}_{:}_reduction_id", typeident, reduction);
let reduction_name = quote::format_ident!("{:}_{:}_reduction", typeident, reduction);
let array_impls = quote! {
#[allow(non_camel_case_types)]
#[#am_data(Clone,Debug,AmGroup(false))]
struct #reduction_name {
data: #lamellar::array::LamellarByteArray,
start_pe: usize,
end_pe: usize,
}
#[#am(AmGroup(false))]
impl LamellarAM for #reduction_name {
async fn exec(&self) -> Vec<u8> {
if self.start_pe == self.end_pe {
let local = self.data.local_data::<#typeident>().await;
match local.reduce(#op){
None => Vec::new(),
Some(v) => #lamellar::serialize(&v, true).expect("failed to serialize reduction result"),
}
} else {
let mid_pe = (self.start_pe + self.end_pe) / 2;
let op = #op;
let left = __lamellar_team.spawn_am_pe(self.start_pe, #reduction_name {
data: self.data.clone(), start_pe: self.start_pe, end_pe: mid_pe
});
let right = __lamellar_team.spawn_am_pe(mid_pe + 1, #reduction_name {
data: self.data.clone(), start_pe: mid_pe + 1, end_pe: self.end_pe
});
let left_bytes = left.await;
let right_bytes = right.await;
if left_bytes.is_empty() && right_bytes.is_empty() {
Vec::new()
} else if left_bytes.is_empty() {
right_bytes
} else if right_bytes.is_empty() {
left_bytes
} else {
let left_val = #lamellar::deserialize::<#typeident>(&left_bytes, true).expect("merge_scalar des left");
let right_val = #lamellar::deserialize::<#typeident>(&right_bytes, true).expect("merge_scalar des right");
let res = op(left_val, right_val);
#lamellar::serialize(&res, true).expect("failed to serialize reduction result")
}
}
}
}
};
let mut gen_match_stmts = quote! {};
gen_match_stmts.extend(quote! {
#lamellar::array::LamellarByteArray::NativeAtomicArray(_) => panic!("this type is not a native atomic"),
#lamellar::array::LamellarByteArray::NetworkAtomicArray(_) => panic!("this type is not a network atomic"),
});
for array_type in array_types {
gen_match_stmts.extend(quote! {
#lamellar::array::LamellarByteArray::#array_type(_) => std::sync::Arc::new(#reduction_name {
data: data.clone(), start_pe: 0, end_pe: num_pes - 1
}),
});
}
let expanded = quote! {
fn #reduction_gen (data: #lamellar::array::LamellarByteArray, num_pes: usize)
-> std::sync::Arc<dyn #lamellar::active_messaging::RemoteActiveMessage + Sync + Send>{
match data{
#gen_match_stmts
}
}
fn #reduction_id_gen () -> std::any::TypeId{
std::any::TypeId::of::<#typeident>()
}
#lamellar::inventory::submit! {
#lamellar::array::ReduceKey{
id: #reduction_id_gen,
name: stringify!(#reduction), gen: #reduction_gen
}
}
#array_impls
};
let user_expanded = quote_spanned! {expanded.span()=>
const _: () = {
extern crate lamellar as __lamellar;
use __lamellar::active_messaging::prelude::*;
#expanded
};
};
if lamellar == "crate" {
expanded
} else {
user_expanded
}
}
pub(crate) fn __register_reduction(item: TokenStream) -> TokenStream {
let args = parse_macro_input!(item as ReductionArgs);
let mut output = quote! {};
let array_types: Vec<syn::Ident> = vec![
quote::format_ident!("LocalLockArray"),
quote::format_ident!("GlobalLockArray"),
quote::format_ident!("AtomicArray"),
quote::format_ident!("GenericAtomicArray"),
quote::format_ident!("UnsafeArray"),
quote::format_ident!("ReadOnlyArray"),
];
for ty in args.tys {
let mut closure = args.closure.clone();
let tyc = ty.clone();
if let syn::Pat::Ident(a) = &closure.inputs[0] {
let pat: syn::PatType = syn::PatType {
attrs: vec![],
pat: Box::new(syn::Pat::Ident(a.clone())),
colon_token: syn::Token),
ty: Box::new(syn::Type::Path(tyc.clone())),
};
closure.inputs[0] = syn::Pat::Type(pat);
}
if let syn::Pat::Ident(b) = &closure.inputs[1] {
let tyc = ty.clone();
let pat: syn::PatType = syn::PatType {
attrs: vec![],
pat: Box::new(syn::Pat::Ident(b.clone())),
colon_token: syn::Token),
ty: Box::new(syn::Type::Path(tyc.clone())),
};
closure.inputs[1] = syn::Pat::Type(pat);
}
output.extend(create_reduction(
ty.path.segments[0].ident.clone(),
args.name.to_string(),
quote! {#closure},
&array_types,
));
}
TokenStream::from(output)
}