use std::sync::OnceLock;
use proc_macro::TokenStream;
use proc_macro2::Span;
use proc_macro_error::proc_macro_error;
use quote::quote_spanned;
use syn::{
parse::{Parse, ParseStream}, parse_macro_input, punctuated::Punctuated, spanned::Spanned, FnArg, ItemFn, Pat, ReturnType, Token, Type
};
#[proc_macro_attribute]
#[proc_macro_error]
pub fn stage(attr: TokenStream, item: TokenStream) -> TokenStream {
let input_fn = parse_macro_input!(item as ItemFn);
let meta_args = parse_macro_input!(attr as StageArgs);
generate_stage_impl(StageConfig::from_args(&input_fn, &meta_args).unwrap()).into()
}
#[derive(Clone)]
struct StageConfig {
original_fn: ItemFn,
stage_name: syn::Ident,
is_lazy: (bool, Span),
cache_strategy: (CacheStrategy, Span),
outputs: Vec<(String, Type, Span)>,
inputs: Vec<InputParam>,
state_type: proc_macro2::TokenStream
}
#[derive(Clone)]
enum RefType {
Owned,
Borrowed,
BorrowedMut,
}
impl RefType {
fn quoted(&self) -> proc_macro2::TokenStream {
match self {
RefType::Owned => quote_spanned! {Span::call_site()=> directed::RefType::Owned },
RefType::Borrowed => quote_spanned! {Span::call_site()=> directed::RefType::Borrowed },
RefType::BorrowedMut => quote_spanned! {Span::call_site()=> directed::RefType::BorrowedMut },
}
}
}
#[derive(Clone)]
struct InputParam {
name: syn::Ident,
type_: Type,
ref_type: RefType,
clean_name: String,
span: Span
}
#[derive(Clone)]
struct Outputs(Punctuated<Output, Token![,]>);
impl Parse for Outputs {
fn parse(input: ParseStream) -> syn::Result<Self> {
Punctuated::parse_terminated(input).map(Self)
}
}
#[derive(Clone)]
struct Output {
name: syn::Ident,
ty: Type,
span: Span,
}
impl Parse for Output {
fn parse(input: ParseStream) -> syn::Result<Self> {
let name: syn::Ident = input.parse()?;
let _colon_token: Token![:] = input.parse()?;
let ty: syn::Type = input.parse()?;
Ok(Output { name, ty, span: Span::call_site() })
}
}
enum StageArg {
Flag(syn::Ident),
Output(Outputs),
State(syn::Type)
}
impl Parse for StageArg {
fn parse(input: ParseStream) -> syn::Result<Self> {
let lookahead = input.lookahead1();
if lookahead.peek(syn::Ident) {
let ident: syn::Ident = input.parse()?;
if ident == "out" {
let content;
let _paren_token = syn::parenthesized!(content in input);
return Ok(StageArg::Output(content.parse()?));
} else if ident == "state" {
let content;
let _paren_token = syn::parenthesized!(content in input);
return Ok(StageArg::State(content.parse()?));
} else {
return Ok(StageArg::Flag(ident));
}
}
Err(lookahead.error())
}
}
struct StageArgs {
args: Punctuated<StageArg, Token![,]>,
}
impl Parse for StageArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
Ok(StageArgs {
args: Punctuated::parse_terminated(input)?,
})
}
}
#[derive(PartialEq, Eq, Clone, Copy)]
enum CacheStrategy {
None,
Last,
All,
}
impl StageConfig {
fn from_args(input_fn: &ItemFn, meta_args: &StageArgs) -> syn::Result<Self> {
let stage_name = input_fn.sig.ident.clone();
let mut is_lazy = (false, Span::call_site());
let mut cache_strategy = (CacheStrategy::None, Span::call_site());
let mut outputs = Vec::new();
let mut state_type = quote_spanned!(Span::call_site()=>());
for arg in meta_args.args.iter() {
match arg {
StageArg::Flag(ident) => match ident.to_string().as_str() {
"lazy" => is_lazy = (true, ident.span()),
"cache_last" => cache_strategy = (CacheStrategy::Last, ident.span()),
"cache_all" => cache_strategy = (CacheStrategy::All, ident.span()),
unknown => {
return Err(syn::Error::new(
ident.span(),
format!("Unrecognized attribute: {}", unknown),
));
}
},
StageArg::Output(output_defs) => {
for output in &output_defs.0 {
outputs.push((output.name.to_string(), output.ty.clone(), output.span));
}
},
StageArg::State(ty) => {
state_type = quote_spanned!(ty.span()=>#ty);
}
}
}
let inputs = Self::extract_input_params(&input_fn.sig.inputs)?;
if outputs.is_empty() {
outputs = Self::extract_outputs_from_return_type(&input_fn.sig.output)?;
}
Ok(StageConfig {
original_fn: input_fn.clone(),
stage_name,
is_lazy,
cache_strategy,
outputs,
inputs,
state_type
})
}
fn extract_input_params(
inputs: &syn::punctuated::Punctuated<FnArg, Token![,]>,
) -> syn::Result<Vec<InputParam>> {
let mut result = Vec::new();
for arg in inputs.iter() {
if let FnArg::Typed(pat_type) = arg {
if let Pat::Ident(pat_ident) = &*pat_type.pat {
let arg_name = &pat_ident.ident;
let arg_type = &pat_type.ty;
let arg_name_str = arg_name.to_string();
let is_unused = arg_name_str.starts_with('_');
let clean_name = if is_unused {
arg_name_str[1..].to_string()
} else {
arg_name_str.clone()
};
let ref_type = match &**arg_type {
Type::Reference(type_reference) if type_reference.mutability.is_some() => {
RefType::BorrowedMut
}
Type::Reference(_) => RefType::Borrowed,
_ => RefType::Owned,
};
result.push(InputParam {
name: arg_name.clone(),
type_: *arg_type.clone(),
ref_type,
clean_name,
span: arg.span()
});
}
}
}
Ok(result)
}
fn extract_outputs_from_return_type(
return_type: &ReturnType,
) -> syn::Result<Vec<(String, Type, Span)>> {
match return_type {
ReturnType::Type(_, ty) => {
if let Type::Path(type_path) = &**ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "NodeOutput" {
return Ok(Vec::new());
}
}
}
Ok(vec![("_".to_string(), (**ty).clone(), ty.span())])
}
ReturnType::Default => {
Ok(vec![(
"_".to_string(),
Type::Tuple(syn::TypeTuple {
paren_token: syn::token::Paren::default(),
elems: Punctuated::new(),
}),
Span::mixed_site(),
)])
}
}
}
fn is_multi_output(&self) -> bool {
if let ReturnType::Type(_, ty) = &self.original_fn.sig.output {
if let Type::Path(type_path) = &**ty {
if let Some(segment) = type_path.path.segments.last() {
return segment.ident == "NodeOutput";
}
}
}
false
}
}
fn generate_input_registrations(inputs: &[InputParam]) -> Vec<proc_macro2::TokenStream> {
inputs.iter().map(|input| {
let arg_name = &input.clean_name;
let arg_type = &input.type_;
let ref_type = input.ref_type.quoted();
let span = input.span.clone();
quote_spanned! {span=>
inputs.insert(directed::DataLabel::new_with_type_name(#arg_name, stringify!(#arg_type)), (std::any::TypeId::of::<#arg_type>(), #ref_type));
}
}).collect()
}
fn generate_output_registrations(outputs: &[(String, Type, Span)]) -> Vec<proc_macro2::TokenStream> {
outputs
.iter()
.map(|(name, ty, _span)| {
quote_spanned! {Span::call_site()=>
outputs.insert(directed::DataLabel::new_with_type_name(#name, stringify!(#ty)), std::any::TypeId::of::<#ty>());
}
})
.collect()
}
fn true_type(ty: &syn::Type) -> &syn::Type {
if let syn::Type::Reference(ty) = ty {
&*ty.elem
} else {
ty
}
}
fn generate_extraction_code(inputs: &[InputParam], cache_strategy: (CacheStrategy, Span)) -> Vec<proc_macro2::TokenStream> {
inputs.iter().map(|input| {
let arg_name = &input.name;
let arg_type = true_type(&input.type_);
let clean_arg_name = &input.clean_name;
let reeval_name = quote::format_ident!("{}_reevaluation_rule", clean_arg_name);
let input_span = input.span.clone();
match cache_strategy {
(CacheStrategy::None, _span) => quote_spanned! {input_span=>
let #reeval_name: directed::ReevaluationRule = inputs.get(&directed::DataLabel::new(#clean_arg_name))
.map(|(_, reeval_rule)| *reeval_rule).ok_or_else(|| directed::InjectionError::InputNotFound(#clean_arg_name.into()))?;
let #arg_name: std::sync::Arc<#arg_type> = if #reeval_name == directed::ReevaluationRule::Move {
if let Some((input, _)) = inputs.remove(&directed::DataLabel::new(#clean_arg_name)) {
let dc = std::sync::Arc::downcast::<#arg_type>(input);
match dc {
Ok(val) => val,
#[allow(unused_variables)]
Err(e) => return Err(directed::InjectionError::InputTypeMismatchDetails{ name: #clean_arg_name, expected: stringify!(#arg_type)})
}
} else {
return Err(directed::InjectionError::InputNotFound(#clean_arg_name.into()));
}
} else {
if let Some((input, _)) = inputs.get(&directed::DataLabel::new(#clean_arg_name)) {
match std::sync::Arc::downcast::<#arg_type>(input.clone()) {
Ok(val) => val,
Err(_) => return Err(directed::InjectionError::InputTypeMismatchDetails{ name: #clean_arg_name, expected: stringify!(#arg_type)})
}
} else {
return Err(directed::InjectionError::InputNotFound(#clean_arg_name.into()));
}
};
},
(CacheStrategy::Last, _span) | (CacheStrategy::All, _span) => quote_spanned! {input_span=>
let (#arg_name, #reeval_name): (std::sync::Arc<#arg_type>, directed::ReevaluationRule) = if let Some((input, reeval_rule)) = inputs.get(&directed::DataLabel::new(#clean_arg_name)) {
match std::sync::Arc::downcast::<#arg_type>(input.clone()) {
Ok(val) => (val, *reeval_rule),
Err(_) => return Err(directed::InjectionError::InputTypeMismatchDetails{ name: #clean_arg_name, expected: stringify!(#arg_type)})
}
} else {
return Err(directed::InjectionError::InputNotFound(#clean_arg_name.into()));
};
},
}
}).collect()
}
fn input_injection(inputs: &[InputParam]) -> proc_macro2::TokenStream {
let mut inject_opaque_out_code = Vec::new();
let mut inject_transparent_out_to_owned_in_code = Vec::new();
let mut inject_transparent_out_to_opaque_ref_in_code = Vec::new();
for input in inputs.iter() {
let clean_arg_name = &input.clean_name;
let arg_type = true_type(&input.type_);
let span = input.span.clone();
inject_opaque_out_code.push(quote_spanned! {span=>
#clean_arg_name => {
#[allow(unused_variables)]
let input_changed = node.input_changed();
let output_val = parent.outputs_mut()
.remove(&output)
.ok_or_else(|| directed::InjectionError::OutputNotFound(output.clone()))?;
let output_val = std::sync::Arc::downcast::<#arg_type>(output_val)
.map_err(|_| directed::InjectionError::OutputTypeMismatch(output.clone()))?;
node.inputs_mut().insert(input, (output_val, directed::ReevaluationRule::Move));
Ok(())
}
});
inject_transparent_out_to_owned_in_code.push(quote_spanned! {span=>
#clean_arg_name => {
#[allow(unused_variables)]
let input_changed = node.input_changed();
let output_val = parent.outputs_mut()
.get(&output)
.ok_or_else(|| directed::InjectionError::OutputNotFound(output.clone()))?
.clone(); let output_val = std::sync::Arc::downcast::<#arg_type>(output_val)
.map_err(|_| directed::InjectionError::OutputTypeMismatch(output.clone()))?;
match node.inputs_mut().get(&input) {
Some((input_val, _)) => {
let input_val = input_val
.downcast_ref::<#arg_type>()
.ok_or_else(|| directed::InjectionError::InputTypeMismatch(input.clone()))?;
if !input_changed && output_val.as_ref() != input_val {
node.set_input_changed(true);
}
},
None => {
node.set_input_changed(true);
}
}
node.inputs_mut().insert(input, (output_val, directed::ReevaluationRule::CacheLast));
Ok(())
}
});
inject_transparent_out_to_opaque_ref_in_code.push(quote_spanned! {span=>
#clean_arg_name => {
#[allow(unused_variables)]
let input_changed = node.input_changed();
let output_val_arc = parent.outputs_mut()
.get(&output)
.ok_or_else(|| directed::InjectionError::OutputNotFound(output.clone()))?;
let output_val_ref = std::sync::Arc::downcast::<#arg_type>(output_val_arc.clone())
.map_err(|_| directed::InjectionError::InputTypeMismatch(input.clone()))?;
match node.inputs_mut().get(&input) {
Some((input_val, _)) => {
let input_val = input_val
.downcast_ref::<#arg_type>()
.ok_or_else(|| directed::InjectionError::InputTypeMismatch(input.clone()))?;
if !input_changed && input_val != &*output_val_ref {
node.set_input_changed(true);
}
},
None => {
node.set_input_changed(true);
}
}
node.inputs_mut().insert(input, (output_val_ref, directed::ReevaluationRule::CacheLast));
Ok(())
}
});
}
let default_case = quote_spanned! {Span::call_site()=>
name => Err(directed::InjectionError::InputNotFound(name.into()))
};
inject_opaque_out_code.push(default_case.clone());
inject_transparent_out_to_owned_in_code.push(default_case.clone());
inject_transparent_out_to_opaque_ref_in_code.push(default_case);
quote_spanned! {Span::call_site()=>
fn inject_opaque_out(node: &mut dyn directed::AnyNode, parent: &mut Box<dyn directed::AnyNode>, output: directed::DataLabel, input: directed::DataLabel) -> Result<(), directed::InjectionError> {
match input.inner() {
#(#inject_opaque_out_code)*
}
}
fn inject_transparent_out_to_owned_in(node: &mut dyn directed::AnyNode, parent: &mut Box<dyn directed::AnyNode>, output: directed::DataLabel, input: directed::DataLabel) -> Result<(), directed::InjectionError> {
match input.inner() {
#(#inject_transparent_out_to_owned_in_code)*
}
}
fn inject_transparent_out_to_opaque_ref_in(node: &mut dyn directed::AnyNode, parent: &mut Box<dyn directed::AnyNode>, output: directed::DataLabel, input: directed::DataLabel) -> Result<(), directed::InjectionError> {
match input.inner() {
#(#inject_transparent_out_to_opaque_ref_in_code)*
}
}
if parent.reeval_rule() == directed::ReevaluationRule::Move {
if node.reeval_rule() == directed::ReevaluationRule::Move && node.input_reftype(&input) != Some(directed::RefType::Owned) {
inject_transparent_out_to_opaque_ref_in(node, parent, output, input)
} else {
inject_opaque_out(node, parent, output, input)
}
} else {
inject_transparent_out_to_owned_in(node, parent, output, input)
}
}
}
fn generate_output_handling(config: &StageConfig, cache_strategy: (CacheStrategy, Span)) -> proc_macro2::TokenStream {
let arg_names = config.inputs.iter().map(|input| &input.name).collect::<Vec<_>>();
let clean_names = config.inputs.iter().map(|input| &input.clean_name).collect::<Vec<_>>();
let arg_types = config.inputs.iter().map(|input| &input.type_).collect::<Vec<_>>();
let downcast_ref_calls = clean_names.iter().zip(arg_types.iter()).zip(arg_names.iter()).map(|((clean_name, arg_type), arg_name)| {
let name_span = arg_name.span();
quote_spanned!{name_span=>
if let Some(cached_in) = cached.inputs.get(&#clean_name.into()) {
if let Some(dc) = in_val.0.downcast_ref::<#arg_type>() {
if dc.downcast_eq(&**cached_in) {
return true;
}
}
}
}
}).collect::<Vec<_>>();
let fn_call = if config.is_multi_output() {
quote_spanned! {Span::call_site()=>
Self::get_fn()(state, #(#arg_names),*)
}
} else {
quote_spanned! {Span::call_site()=>
directed::NodeOutput::new_simple(Self::get_fn()(state, #(#arg_names),*))
}
};
if cache_strategy.0 == CacheStrategy::All {
quote_spanned!{cache_strategy.1=>
let hash: u64 = {
#[allow(unused_imports)]
use std::hash::Hash;
#[allow(unused_imports)]
use std::hash::Hasher;
#[allow(unused_mut)]
let mut hasher = std::hash::DefaultHasher::new();
#(#arg_names.hash(&mut hasher);)*
hasher.finish()
};
#[allow(unused_variables)]
let cached = cache.get(&hash).and_then(|cached| {
#[allow(unused_imports)]
use directed::DowncastEq;
cached.iter().find(|cached| {
inputs.iter().all(|(in_name, in_val)| {
#(#downcast_ref_calls)*
false
})
})
});
if let Some(cached) = cached {
if cached.outputs.len() == 1 && cached.outputs.get(&"_".into()).is_some() {
Ok(NodeOutput::new_simple(cached.outputs.get(&"_".into()).unwrap().clone()))
} else {
let mut result = NodeOutput::new();
for (out_name, out_val) in cached.outputs.iter() {
result = result.add(&out_name.name, out_val.clone());
}
Ok(result)
}
} else {
let result = #fn_call;
let cache_entry = {
#[allow(unused_mut)]
let mut cached = directed::Cached {
inputs: std::collections::HashMap::new(),
outputs: std::collections::HashMap::new(),
};
for (in_name, in_val) in inputs.iter() {
cached.inputs.insert(in_name.clone(), in_val.0.clone());
}
match &result {
NodeOutput::Standard(val) => {
cached.outputs.insert("_".into(), val.clone());
},
NodeOutput::Named(vals) => {
for (key, val) in vals {
cached.outputs.insert(key.clone(), val.clone());
}
},
}
cached
};
if let None = cache.get(&hash) {
cache.insert(hash, Vec::new());
}
if let Some(vec) = cache.get_mut(&hash) {
vec.push(cache_entry);
}
Ok(result)
}
}
} else {
quote_spanned!(cache_strategy.1=>Ok(#fn_call))
}
}
fn prepare_input_types(config: &StageConfig) -> Vec<proc_macro2::TokenStream> {
let args = config
.inputs
.iter()
.map(|input| (&input.name, &input.clean_name, &input.ref_type));
let mut output = Vec::new();
for (arg_name, clean_name, ref_type) in args {
let reeval_name = quote::format_ident!("{}_reevaluation_rule", clean_name);
match ref_type {
RefType::Owned => {
output.push(quote_spanned!{arg_name.span()=>
let #arg_name = match #reeval_name {
directed::ReevaluationRule::Move => {
match std::sync::Arc::into_inner(#arg_name) {
Some(arg) => arg,
None => {return Err(directed::InjectionError::TooManyReferences(stringify!(#arg_name)))}
}
},
directed::ReevaluationRule::CacheLast | directed::ReevaluationRule::CacheAll => {
(*#arg_name).clone()
},
};
});
}
RefType::Borrowed => {
output.push(quote_spanned!{arg_name.span()=>
let #arg_name = #arg_name.as_ref();
});
}
RefType::BorrowedMut => panic!("Mutable refs are not yet supported"),
}
}
output
}
fn generate_stage_impl(config: StageConfig) -> proc_macro2::TokenStream {
let original_fn = &config.original_fn;
let stage_name = &config.stage_name;
let state_type = &config.state_type;
let fn_attrs = &original_fn.attrs;
let fn_vis = &original_fn.vis;
let original_args = &original_fn.sig.inputs;
let fn_return_type = &original_fn.sig.output;
let original_body = &original_fn.block;
let input_registrations = generate_input_registrations(&config.inputs);
let output_registrations = generate_output_registrations(&config.outputs);
let extraction_code = generate_extraction_code(&config.inputs, config.cache_strategy);
let injection_code = input_injection(&config.inputs);
let prepare_input_types_code = prepare_input_types(&config);
let output_handling = generate_output_handling(&config, config.cache_strategy);
let eval_strategy = if config.is_lazy.0 {
quote_spanned! {config.is_lazy.1=> directed::EvalStrategy::Lazy }
} else {
quote_spanned! {config.is_lazy.1=> directed::EvalStrategy::Urgent }
};
let reevaluation_rule = match &config.cache_strategy {
(CacheStrategy::None, span) => quote_spanned! {*span=> directed::ReevaluationRule::Move },
(CacheStrategy::Last, span) => quote_spanned! {*span=> directed::ReevaluationRule::CacheLast },
(CacheStrategy::All, span) => quote_spanned! {*span=> directed::ReevaluationRule::CacheAll },
};
quote_spanned! {Span::call_site()=>
#[derive(Clone)]
#fn_vis struct #stage_name {
inputs: std::collections::HashMap<directed::DataLabel, (std::any::TypeId, directed::RefType)>,
outputs: std::collections::HashMap<directed::DataLabel, std::any::TypeId>,
}
impl #stage_name {
pub fn new() -> Self {
let mut inputs = std::collections::HashMap::new();
let mut outputs = std::collections::HashMap::new();
#(#input_registrations)*
#(#output_registrations)*
Self { inputs, outputs }
}
}
impl directed::Stage for #stage_name {
type State = #state_type;
type BaseFn = fn(state: &mut #state_type, #original_args) #fn_return_type;
fn inputs(&self) -> &std::collections::HashMap<directed::DataLabel, (std::any::TypeId, directed::RefType)> {
&self.inputs
}
fn outputs(&self) -> &std::collections::HashMap<directed::DataLabel, std::any::TypeId> {
&self.outputs
}
fn evaluate(
&self,
state: &mut Self::State,
inputs: &mut std::collections::HashMap<directed::DataLabel, (std::sync::Arc<dyn std::any::Any + Send + Sync>, directed::ReevaluationRule)>,
cache: &mut std::collections::HashMap<u64, Vec<directed::Cached>>
) -> Result<directed::NodeOutput, directed::InjectionError> {
#(#extraction_code)*
#(#prepare_input_types_code)*
#output_handling
}
fn eval_strategy(&self) -> directed::EvalStrategy {
#eval_strategy
}
fn reeval_rule(&self) -> directed::ReevaluationRule {
#reevaluation_rule
}
fn inject_input(&self, node: &mut directed::Node<Self>, parent: &mut Box<dyn directed::AnyNode>, output: directed::DataLabel, input: directed::DataLabel) -> Result<(), directed::InjectionError> {
#injection_code
}
fn name(&self) -> &str {
stringify!(#stage_name)
}
fn get_fn() -> Self::BaseFn {
#(#fn_attrs)*
fn original_fn(state: &mut #state_type, #original_args) #fn_return_type #original_body
return original_fn;
}
}
}
}