use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::quote;
use syn::{
Attribute, Ident, ImplItem, ImplItemFn, ItemImpl, LitStr, Token, Type, Visibility, braced,
bracketed,
parse::{Parse, ParseStream},
parse_macro_input,
punctuated::Punctuated,
};
struct NodeArgs {
active: Vec<Ident>,
passive: Vec<Ident>,
output: Option<(Ident, Type)>,
}
impl Parse for NodeArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
let mut active: Option<Vec<Ident>> = None;
let mut passive: Option<Vec<Ident>> = None;
let mut output: Option<(Ident, Type)> = None;
while !input.is_empty() {
let key: Ident = input.parse()?;
input.parse::<Token![=]>()?;
match key.to_string().as_str() {
"active" => {
if active.is_some() {
return Err(syn::Error::new(key.span(), "duplicate key `active`"));
}
let content;
bracketed!(content in input);
let list = Punctuated::<Ident, Token![,]>::parse_terminated(&content)?;
active = Some(list.into_iter().collect());
}
"passive" => {
if passive.is_some() {
return Err(syn::Error::new(key.span(), "duplicate key `passive`"));
}
let content;
bracketed!(content in input);
let list = Punctuated::<Ident, Token![,]>::parse_terminated(&content)?;
passive = Some(list.into_iter().collect());
}
"output" => {
if output.is_some() {
return Err(syn::Error::new(key.span(), "duplicate key `output`"));
}
let field: Ident = input.parse()?;
input.parse::<Token![:]>()?;
let ty: Type = input.parse()?;
output = Some((field, ty));
}
_ => {
return Err(syn::Error::new(
key.span(),
format!("unknown key `{key}`; expected `active`, `passive`, or `output`"),
));
}
}
if input.peek(Token![,]) {
input.parse::<Token![,]>()?;
}
}
Ok(NodeArgs {
active: active.unwrap_or_default(),
passive: passive.unwrap_or_default(),
output,
})
}
}
#[proc_macro_attribute]
pub fn node(attr: TokenStream, item: TokenStream) -> TokenStream {
let args = parse_macro_input!(attr as NodeArgs);
let mut impl_block = parse_macro_input!(item as ItemImpl);
let self_ty = impl_block.self_ty.clone();
let (impl_generics, _, where_clause) = impl_block.generics.split_for_impl();
if !args.active.is_empty() || !args.passive.is_empty() {
let active_fields = &args.active;
let passive_fields = &args.passive;
let upstreams_fn: ImplItemFn = syn::parse_quote! {
fn upstreams(&self) -> ::wingfoil::UpStreams {
let mut active: ::std::vec::Vec<::std::rc::Rc<dyn ::wingfoil::Node>> = ::std::vec::Vec::new();
let mut passive: ::std::vec::Vec<::std::rc::Rc<dyn ::wingfoil::Node>> = ::std::vec::Vec::new();
#(active.extend(::wingfoil::AsUpstreamNodes::as_upstream_nodes(&self.#active_fields));)*
#(passive.extend(::wingfoil::AsUpstreamNodes::as_upstream_nodes(&self.#passive_fields));)*
::wingfoil::UpStreams::new(active, passive)
}
};
impl_block.items.push(ImplItem::Fn(upstreams_fn));
}
let peek_ref_impl = args.output.map(|(field, ty)| {
quote! {
impl #impl_generics ::wingfoil::StreamPeekRef<#ty> for #self_ty #where_clause {
fn peek_ref(&self) -> &#ty {
&self.#field
}
}
}
});
quote! {
#impl_block
#peek_ref_impl
}
.into()
}
struct LatencyStagesInput {
visibility: Visibility,
name: Ident,
stages: Vec<Ident>,
type_name_override: Option<LitStr>,
}
impl Parse for LatencyStagesInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let attrs = input.call(Attribute::parse_outer)?;
let mut type_name_override: Option<LitStr> = None;
for attr in &attrs {
if attr.path().is_ident("type_name") {
if type_name_override.is_some() {
return Err(syn::Error::new_spanned(
attr,
"duplicate #[type_name(...)] attribute",
));
}
let lit: LitStr = attr.parse_args().map_err(|_| {
syn::Error::new_spanned(
attr,
"expected #[type_name(\"...\")] with a single string literal",
)
})?;
type_name_override = Some(lit);
} else {
return Err(syn::Error::new_spanned(
attr,
"unrecognized attribute on latency_stages!; only #[type_name(\"...\")] is supported",
));
}
}
let visibility: Visibility = input.parse()?;
let name: Ident = input.parse()?;
let content;
braced!(content in input);
let list = Punctuated::<Ident, Token![,]>::parse_terminated(&content)?;
let stages: Vec<Ident> = list.into_iter().collect();
if stages.is_empty() {
return Err(syn::Error::new(
name.span(),
"latency_stages! requires at least one stage",
));
}
Ok(LatencyStagesInput {
visibility,
name,
stages,
type_name_override,
})
}
}
fn pascal_to_snake(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 4);
for (i, ch) in s.chars().enumerate() {
if ch.is_ascii_uppercase() {
if i != 0 {
out.push('_');
}
out.push(ch.to_ascii_lowercase());
} else {
out.push(ch);
}
}
out
}
#[proc_macro]
pub fn latency_stages(item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as LatencyStagesInput);
let LatencyStagesInput {
visibility,
name,
stages,
type_name_override,
} = input;
let n = stages.len();
let module_name = Ident::new(&pascal_to_snake(&name.to_string()), Span::call_site());
let stage_strs: Vec<String> = stages.iter().map(|i| i.to_string()).collect();
let stage_indices: Vec<usize> = (0..n).collect();
let field_names = &stages;
let marker_names = &stages;
let zero_copy_send_body = match type_name_override {
Some(lit) => quote! {
unsafe fn type_name() -> &'static str { #lit }
},
None => quote! {},
};
let expanded = quote! {
#[repr(C)]
#[derive(
::std::clone::Clone, ::std::marker::Copy,
::std::fmt::Debug, ::std::default::Default,
::std::cmp::PartialEq, ::std::cmp::Eq,
::std::hash::Hash,
::serde::Serialize, ::serde::Deserialize,
)]
#visibility struct #name {
#( pub #field_names: u64, )*
}
impl Latency for #name {
const N: usize = #n;
fn stage_names() -> &'static [&'static str] {
&[ #( #stage_strs ),* ]
}
#[inline]
fn stamps(&self) -> &[u64] {
unsafe {
::std::slice::from_raw_parts(
self as *const Self as *const u64,
<Self as Latency>::N,
)
}
}
#[inline]
fn stamp_mut(&mut self, idx: usize) -> &mut u64 {
assert!(idx < <Self as Latency>::N, "stage index out of bounds");
unsafe { &mut *((self as *mut Self as *mut u64).add(idx)) }
}
}
#[cfg(feature = "iceoryx2")]
unsafe impl ::iceoryx2::prelude::ZeroCopySend for #name {
#zero_copy_send_body
}
#[allow(non_snake_case, non_camel_case_types)]
#visibility mod #module_name {
use super::*;
#(
pub struct #marker_names;
impl Stage<super::#name> for #marker_names {
const NAME: &'static str = #stage_strs;
const INDEX: usize = #stage_indices;
}
)*
}
};
expanded.into()
}