use proc_macro2::TokenStream;
use quote::{quote, ToTokens};
use syn::{
parse_quote, visit::Visit, visit_mut::VisitMut, Expr, ExprField, Ident,
Member,
};
use crate::names;
pub fn soa_impl_transform(
item: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let Ok(item_impl) = syn::parse::<syn::ItemImpl>(item) else {
panic!("soa_impl can only be applied to impl blocks");
};
let struct_ident = extract_struct_ident(&item_impl);
let field_names = collect_self_field_names(&item_impl);
let original_tokens = item_impl.to_token_stream();
let ref_impl = generate_ref_impl(&item_impl, &struct_ident, &field_names);
let ref_mut_impl =
generate_ref_mut_impl(&item_impl, &struct_ident, &field_names);
quote! {
#original_tokens
#ref_impl
#ref_mut_impl
}
.into()
}
fn extract_struct_ident(item_impl: &syn::ItemImpl) -> syn::Ident {
if let syn::Type::Path(type_path) = &*item_impl.self_ty {
type_path
.path
.segments
.first()
.expect("expected a type name")
.ident
.clone()
} else {
panic!(
"soa_impl can only be applied to impl blocks for a named struct"
);
}
}
struct FieldNameCollector {
field_names: Vec<Ident>,
}
impl FieldNameCollector {
fn new() -> Self {
FieldNameCollector {
field_names: Vec::new(),
}
}
}
impl<'ast> Visit<'ast> for FieldNameCollector {
fn visit_expr_field(&mut self, expr: &'ast ExprField) {
if is_self_field_expr(&expr.base) {
if let Member::Named(ident) = &expr.member {
if !self.field_names.contains(ident) {
self.field_names.push(ident.clone());
}
}
}
syn::visit::visit_expr_field(self, expr);
}
}
fn is_self_field_expr(expr: &Expr) -> bool {
if let Expr::Path(path_expr) = expr {
path_expr.path.is_ident("self")
} else {
false
}
}
fn collect_self_field_names(item_impl: &syn::ItemImpl) -> Vec<Ident> {
let mut collector = FieldNameCollector::new();
collector.visit_item_impl(item_impl);
collector.field_names
}
fn is_compound_assign_op(op: &syn::BinOp) -> bool {
matches!(
op,
syn::BinOp::AddAssign(_)
| syn::BinOp::SubAssign(_)
| syn::BinOp::MulAssign(_)
| syn::BinOp::DivAssign(_)
| syn::BinOp::RemAssign(_)
| syn::BinOp::BitAndAssign(_)
| syn::BinOp::BitOrAssign(_)
| syn::BinOp::BitXorAssign(_)
| syn::BinOp::ShlAssign(_)
| syn::BinOp::ShrAssign(_)
)
}
struct SelfFieldTransformer<'a> {
field_names: &'a [Ident],
is_ref_mut: bool,
suppress: bool,
}
impl SelfFieldTransformer<'_> {
fn is_known_field(&self, expr: &Expr) -> bool {
if let Expr::Field(field_expr) = expr {
if is_self_field_expr(&field_expr.base) {
if let Member::Named(ident) = &field_expr.member {
return self.field_names.contains(ident);
}
}
}
false
}
fn wrap_deref(expr: &Expr) -> Expr {
let tokens = expr.to_token_stream();
parse_quote! { (*#tokens) }
}
fn prefix_deref(expr: &Expr) -> Expr {
let tokens = expr.to_token_stream();
parse_quote! { *#tokens }
}
}
impl VisitMut for SelfFieldTransformer<'_> {
fn visit_expr_mut(&mut self, expr: &mut Expr) {
match expr {
Expr::Assign(assign) => {
self.visit_expr_assign_mut(assign);
return;
}
Expr::Binary(binary) if is_compound_assign_op(&binary.op) => {
self.visit_expr_binary_mut(binary);
return;
}
Expr::Reference(reference) => {
self.visit_expr_reference_mut(reference);
return;
}
_ => {}
}
if !self.suppress && self.is_known_field(expr) {
*expr = SelfFieldTransformer::wrap_deref(expr);
return;
}
syn::visit_mut::visit_expr_mut(self, expr);
}
fn visit_expr_assign_mut(&mut self, expr: &mut syn::ExprAssign) {
let prev = self.suppress;
self.suppress = true;
self.visit_expr_mut(&mut expr.left);
self.suppress = prev;
self.visit_expr_mut(&mut expr.right);
if self.is_ref_mut && self.is_known_field(&expr.left) {
let prefixed = SelfFieldTransformer::prefix_deref(&expr.left);
*expr.left = prefixed;
}
}
fn visit_expr_binary_mut(&mut self, expr: &mut syn::ExprBinary) {
if is_compound_assign_op(&expr.op) {
let prev = self.suppress;
self.suppress = true;
self.visit_expr_mut(&mut expr.left);
self.suppress = prev;
self.visit_expr_mut(&mut expr.right);
if self.is_ref_mut && self.is_known_field(&expr.left) {
let prefixed = SelfFieldTransformer::prefix_deref(&expr.left);
*expr.left = prefixed;
}
} else {
self.visit_expr_mut(&mut expr.left);
self.visit_expr_mut(&mut expr.right);
}
}
fn visit_expr_reference_mut(&mut self, expr: &mut syn::ExprReference) {
let prev = self.suppress;
self.suppress = true;
self.visit_expr_mut(&mut expr.expr);
self.suppress = prev;
}
fn visit_expr_method_call_mut(&mut self, expr: &mut syn::ExprMethodCall) {
if !self.is_known_field(&expr.receiver) {
self.visit_expr_mut(&mut expr.receiver);
}
for arg in &mut expr.args {
self.visit_expr_mut(arg);
}
}
fn visit_block_mut(&mut self, block: &mut syn::Block) {
syn::visit_mut::visit_block_mut(self, block);
}
}
fn is_ref_self_method(method: &syn::ImplItemFn) -> bool {
if let Some(syn::FnArg::Receiver(receiver)) = method.sig.inputs.first() {
matches!(receiver.kind, syn::ReceiverKind::Reference(_, _, None))
} else {
false
}
}
fn is_mut_self_method(method: &syn::ImplItemFn) -> bool {
if let Some(syn::FnArg::Receiver(receiver)) = method.sig.inputs.first() {
matches!(receiver.kind, syn::ReceiverKind::Reference(_, _, Some(_)))
} else {
false
}
}
fn returns_self(method: &syn::ImplItemFn) -> bool {
if let syn::ReturnType::Type(_, ty) = &method.sig.output {
if let syn::Type::Path(type_path) = &**ty {
return type_path.path.is_ident("Self");
}
}
false
}
fn mentions_self(method: &syn::ImplItemFn) -> bool {
tokens_mention_self(method.to_token_stream())
}
fn tokens_mention_self(tokens: proc_macro2::TokenStream) -> bool {
use proc_macro2::TokenTree;
for tt in tokens {
match tt {
TokenTree::Ident(i) if i == "Self" => return true,
TokenTree::Group(g) if tokens_mention_self(g.stream()) => {
return true
}
_ => {}
}
}
false
}
fn generate_ref_impl(
item_impl: &syn::ItemImpl,
struct_ident: &syn::Ident,
field_names: &[Ident],
) -> TokenStream {
let ref_name = names::ref_name(struct_ident);
let mut methods: Vec<syn::ImplItemFn> = Vec::new();
for item in &item_impl.items {
if let syn::ImplItem::Fn(method) = item {
if is_ref_self_method(method)
&& !returns_self(method)
&& !mentions_self(method)
{
let mut cloned = method.clone();
let mut visitor = SelfFieldTransformer {
field_names,
is_ref_mut: false,
suppress: false,
};
visitor.visit_impl_item_fn_mut(&mut cloned);
methods.push(cloned);
}
}
}
if methods.is_empty() {
return TokenStream::new();
}
quote! {
impl<'a> #ref_name<'a> {
#(#methods)*
}
}
}
fn generate_ref_mut_impl(
item_impl: &syn::ItemImpl,
struct_ident: &syn::Ident,
field_names: &[Ident],
) -> TokenStream {
let ref_mut_name = names::ref_mut_name(struct_ident);
let mut methods: Vec<syn::ImplItemFn> = Vec::new();
for item in &item_impl.items {
if let syn::ImplItem::Fn(method) = item {
if (is_mut_self_method(method) || is_ref_self_method(method))
&& !returns_self(method)
&& !mentions_self(method)
{
let mut cloned = method.clone();
let mut visitor = SelfFieldTransformer {
field_names,
is_ref_mut: true,
suppress: false,
};
visitor.visit_impl_item_fn_mut(&mut cloned);
methods.push(cloned);
}
}
}
if methods.is_empty() {
return TokenStream::new();
}
quote! {
impl<'a> #ref_mut_name<'a> {
#(#methods)*
}
}
}