use proc_macro::TokenStream;
use quote::quote;
use syn::{
parse_macro_input, punctuated::Punctuated, spanned::Spanned, token::Comma, FnArg, Ident,
LitByteStr, LitStr,
};
#[cfg(not(target_os = "linux"))]
compile_error!("upgrayedd currently only supports Linux; consider sending a patch.");
fn transform_params(params: Punctuated<FnArg, Comma>) -> Punctuated<Ident, Comma> {
let idents = params.iter().filter_map(|param| {
if let syn::FnArg::Typed(pat_type) = param {
if let syn::Pat::Ident(pat_ident) = *pat_type.pat.clone() {
return Some(pat_ident.ident);
}
}
None
});
let mut punctuated: Punctuated<syn::Ident, Comma> = Punctuated::new();
idents.for_each(|ident| punctuated.push(ident));
punctuated
}
#[proc_macro_attribute]
pub fn upgrayedd(attr: TokenStream, item: TokenStream) -> TokenStream {
let func = parse_macro_input!(item as syn::ItemFn);
let syn::ItemFn {
attrs,
vis,
sig,
block,
} = func;
if !matches!(vis, syn::Visibility::Inherited) {
return syn::Error::new(vis.span(), "upgrayedd-ed functions much be private")
.to_compile_error()
.into();
}
let stmts = &block.stmts;
let syn::Signature {
constness: _,
asyncness: _,
unsafety: _,
abi: _,
fn_token: _,
ident,
generics: _,
paren_token: _,
inputs,
variadic: _,
output,
} = sig.clone();
let inner_var =
parse_macro_input!(attr as Option<Ident>).unwrap_or(Ident::new("upgrayedd", ident.span()));
let real_c_name_lit = LitStr::new(&ident.to_string(), ident.span());
let real_c_name_bytes = {
let real_c_name_lit_bytes = ident.to_string().into_bytes();
LitByteStr::new(&real_c_name_lit_bytes, ident.span())
};
let real_c_name_bytes_nulled = {
let mut real_c_name_lit_bytes = ident.to_string().into_bytes();
real_c_name_lit_bytes.push(0);
LitByteStr::new(&real_c_name_lit_bytes, ident.span())
};
let inner_wrapper = Ident::new(&format!("__upgrayedd_inner_wrapper_{ident}"), ident.span());
let target = Ident::new(&format!("__upgrayedd_target_{ident}"), ident.span());
let args = transform_params(inputs.clone());
let gen = quote! {
static mut #target: std::sync::atomic::AtomicPtr<unsafe extern "C" fn(#inputs) #output> = std::sync::atomic::AtomicPtr::new(std::ptr::null_mut());
#[no_mangle]
#[doc(hidden)]
#[allow(non_snake_case)]
#[export_name = #real_c_name_lit]
pub unsafe extern "C" fn #inner_wrapper(#inputs) #output {
let mut target_ptr = #target.get_mut();
if target_ptr.is_null() {
*target_ptr = &mut std::mem::transmute(
::libc::dlsym(::libc::RTLD_NEXT, std::mem::transmute(#real_c_name_bytes_nulled.as_ptr()))
);
}
if target_ptr.is_null() {
let msg = b"barf: upgrayedd tried to hook something that broke rust's runtime: ";
::libc::write(::libc::STDERR_FILENO, msg.as_ptr() as *const ::libc::c_void, msg.len());
::libc::write(::libc::STDERR_FILENO, #real_c_name_bytes.as_ptr() as *const ::libc::c_void, #real_c_name_bytes.len());
::libc::write(::libc::STDERR_FILENO, b"\n".as_ptr() as *const ::libc::c_void, 1);
std::process::abort();
}
#ident(#args)
}
#[allow(non_snake_case)]
#(#attrs)* #vis #sig {
#[allow(unused_variables)]
let #inner_var = unsafe { *#target.load(std::sync::atomic::Ordering::Relaxed) };
#(#stmts)*
}
};
gen.into()
}