use proc_macro::TokenStream;
use quote::quote;
use syn::{ItemFn, parse_macro_input};
#[proc_macro_attribute]
pub fn trace_function(_attr: TokenStream, item: TokenStream) -> TokenStream {
let input_fn = parse_macro_input!(item as ItemFn);
let fn_name = &input_fn.sig.ident;
let fn_name_str = fn_name.to_string();
let fn_vis = &input_fn.vis;
let fn_sig = &input_fn.sig;
let fn_block = &input_fn.block;
let fn_attrs = &input_fn.attrs;
let fn_output = &input_fn.sig.output;
let entry_msg = format!("> {}", fn_name_str);
let _exit_msg = format!("< {}", fn_name_str);
let expanded = match fn_output {
syn::ReturnType::Default => {
#[cfg(feature = "trace")]
quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
::hyperlight_guest_tracing::create_trace_record(#entry_msg);
#fn_block
::hyperlight_guest_tracing::create_trace_record(#_exit_msg);
}
}
#[cfg(not(feature = "trace"))]
quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
#fn_block
}
}
}
syn::ReturnType::Type(_, _) => {
#[cfg(feature = "trace")]
quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
::hyperlight_guest_tracing::create_trace_record(#entry_msg);
let __trace_result = (|| #fn_block )();
::hyperlight_guest_tracing::create_trace_record(#_exit_msg);
__trace_result
}
}
#[cfg(not(feature = "trace"))]
quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
#fn_block
}
}
}
};
TokenStream::from(expanded)
}
struct TraceMacroInput {
message: syn::Lit,
statement: Option<proc_macro2::TokenStream>,
}
impl syn::parse::Parse for TraceMacroInput {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let message: syn::Lit = input.parse()?;
if !matches!(message, syn::Lit::Str(_)) {
return Err(input.error("first argument to trace! must be a string literal"));
}
if let syn::Lit::Str(ref lit_str) = message {
if lit_str.value().is_empty() {
return Err(input.error("trace message must not be empty"));
}
}
let statement = if input.peek(syn::Token![,]) {
let _: syn::Token![,] = input.parse()?;
Some(input.parse()?)
} else {
None
};
Ok(TraceMacroInput { message, statement })
}
}
#[proc_macro]
pub fn trace(input: TokenStream) -> TokenStream {
let parsed = syn::parse_macro_input!(input as TraceMacroInput);
let trace_message = match parsed.message {
syn::Lit::Str(ref lit_str) => lit_str.value(),
_ => unreachable!(),
};
if let Some(statement) = parsed.statement {
let entry_msg = format!("+ {}", trace_message);
let _exit_msg = format!("- {}", trace_message);
#[cfg(feature = "trace")]
let expanded = quote! {
{
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
::hyperlight_guest_tracing::create_trace_record(#entry_msg);
let __trace_result = #statement;
::hyperlight_guest_tracing::create_trace_record(#_exit_msg);
__trace_result
}
};
#[cfg(not(feature = "trace"))]
let expanded = quote! {
{
const _: () = assert!(
#entry_msg.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
#statement
}
};
TokenStream::from(expanded)
} else {
#[cfg(feature = "trace")]
let expanded = quote! {
{
const _: () = assert!(
#trace_message.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
::hyperlight_guest_tracing::create_trace_record(#trace_message);
}
};
#[cfg(not(feature = "trace"))]
let expanded = quote! {
{
const _: () = assert!(
#trace_message.len() <= hyperlight_guest_tracing::MAX_TRACE_MSG_LEN,
"Trace message exceeds the maximum bytes length",
);
}
};
TokenStream::from(expanded)
}
}
#[proc_macro]
pub fn flush(_input: TokenStream) -> TokenStream {
#[cfg(feature = "trace")]
let expanded = quote! {
{
::hyperlight_guest_tracing::flush_trace_buffer();
}
};
#[cfg(not(feature = "trace"))]
let expanded = quote! {};
TokenStream::from(expanded)
}