use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::{parse_macro_input, ItemFn};
pub fn intercept_impl(_attr: TokenStream, item: TokenStream) -> TokenStream {
let input_fn = parse_macro_input!(item as ItemFn);
let output = generate_intercepted_fn(&input_fn);
output.into()
}
fn generate_intercepted_fn(input_fn: &ItemFn) -> TokenStream2 {
let fn_name = &input_fn.sig.ident;
let vis = &input_fn.vis;
let constness = input_fn.sig.constness;
let unsafety = input_fn.sig.unsafety;
let generics = &input_fn.sig.generics;
let output = &input_fn.sig.output;
let body = &input_fn.block;
let is_async = input_fn.sig.asyncness.is_some();
let mut regular_params: Vec<syn::Pat> = Vec::new();
for param in &input_fn.sig.inputs {
if let syn::FnArg::Typed(pat_type) = param {
regular_params.push(pat_type.pat.as_ref().clone());
}
}
let arg_calls: Vec<TokenStream2> = regular_params.iter().map(|pat| {
quote! {
.with_arg(::tx_di_core::aop::ArgValue::Other(
::std::format!("{:?}", &#pat)
))
}
}).collect();
let params = &input_fn.sig.inputs;
let is_result_ret = is_result_return_type(&input_fn.sig.output);
let after_block = if is_result_ret {
quote! {
let mut __cr = match &__result {
Ok(_) => ::tx_di_core::aop::CallResult::Ok,
Err(e) => ::tx_di_core::aop::CallResult::Err(::std::format!("{}", e)),
};
__chain.after_all(&__ctx, &mut __cr);
}
} else {
quote! {
let mut __cr = ::tx_di_core::aop::CallResult::Ok;
__chain.after_all(&__ctx, &mut __cr);
}
};
let async_prefix = if is_async { quote! { async } } else { quote! {} };
quote! {
#vis #constness #unsafety #async_prefix fn #fn_name #generics (#params) #output {
let __ctx = ::tx_di_core::aop::CallContext::new(stringify!(#fn_name))
#(#arg_calls)*;
let __key = self as *const Self as usize;
let __chain = ::tx_di_core::aop::get_interceptor_chain(__key)
.expect("[di] 拦截器链未初始化:请确认组件已通过 #[component(intercept(...))] 声明,且 App 已运行初始化阶段");
__chain.before_all(&__ctx).unwrap_or_else(|e| {
panic!("[di] 拦截器拒绝 method={}: {}", stringify!(#fn_name), e)
});
let __result = #body;
#after_block
__result
}
}
}
fn is_result_return_type(output: &syn::ReturnType) -> bool {
match output {
syn::ReturnType::Type(_, ty) => {
let s = quote! { #ty }.to_string();
s.starts_with("Result ") || s.starts_with("Result<")
|| s.starts_with("::std::result::Result ")
|| s.starts_with("::core::result::Result ")
|| s.starts_with("RIE ") || s.starts_with("RIE<")
|| s.starts_with("AppResult ") || s.starts_with("AppResult<")
|| s.starts_with("::tx_di_core::RIE ")
|| s.starts_with("::tx_di_core::RIE<")
}
_ => false,
}
}