From e5f61421fbd3acb9b6a5c0f82d6f52d8eec6c348 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 12 Apr 2026 17:31:44 +0000 Subject: [PATCH 1/2] Support default method bodies in #[devirt] trait definitions Trait methods with default bodies (e.g. `fn is_large(&self) -> bool { self.area() > 100.0 }`) are now accepted by the proc-macro attribute. The default body is placed on the inner trait's `__spec_*` provided method, with `self.method()` calls rewritten to `self.__spec_method()` via `syn::visit_mut` so sibling calls resolve within the inner trait. Macro invocations (format!, write!, etc.) are handled via token-level rewriting since syn doesn't parse macro token streams. https://claude.ai/code/session_01BJen3TbKCFohcjNFeYwBiA --- Cargo.toml | 2 +- crates/core/tests/equivalence.rs | 61 +++++++ crates/core/tests/ui_attr.rs | 3 + .../core/tests/ui_attr/attr_default_body.rs | 60 +++++++ .../tests/ui_attr/attr_default_override.rs | 27 +++ .../core/tests/ui_attr/attr_default_send.rs | 30 ++++ crates/macros/src/lib.rs | 161 +++++++++++++++--- 7 files changed, 323 insertions(+), 21 deletions(-) create mode 100644 crates/core/tests/ui_attr/attr_default_body.rs create mode 100644 crates/core/tests/ui_attr/attr_default_override.rs create mode 100644 crates/core/tests/ui_attr/attr_default_send.rs diff --git a/Cargo.toml b/Cargo.toml index f9444e5..d20ff8c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -16,7 +16,7 @@ libfuzzer-sys = "0.4" arbitrary = { version = "1", features = ["derive"] } devirt = { path = "crates/core" } devirt-macros = { path = "crates/macros", version = "0.2.0" } -syn = { version = "2", features = ["full"] } +syn = { version = "2", features = ["full", "visit-mut"] } quote = "1" proc-macro2 = "1" vstd = { version = "=0.0.0-2026-04-12-0118", default-features = false } diff --git a/crates/core/tests/equivalence.rs b/crates/core/tests/equivalence.rs index d274437..d8c6037 100644 --- a/crates/core/tests/equivalence.rs +++ b/crates/core/tests/equivalence.rs @@ -144,6 +144,67 @@ fn attr_auto_trait_dispatch() { assert_eq!(boxed.get(), 7); } +// ── Default method bodies ────────────────────────────────────────────────── + +#[cfg(feature = "macros")] +mod attr_defaults { + pub struct DefHot { + pub val: u64, + } + + pub struct DefCold { + pub val: u64, + } + + #[devirt::devirt(DefHot)] + pub trait Defaulted { + fn get(&self) -> u64; + fn is_big(&self) -> bool { + self.get() > 100 + } + } + + #[devirt::devirt] + impl Defaulted for DefHot { + fn get(&self) -> u64 { + self.val + } + } + + #[devirt::devirt] + impl Defaulted for DefCold { + fn get(&self) -> u64 { + self.val + 1 + } + } +} + +#[cfg(feature = "macros")] +#[test] +fn attr_default_body_dispatch() { + use attr_defaults::{DefCold, DefHot, Defaulted}; + + // Hot type via &dyn Trait + let h = DefHot { val: 200 }; + assert!((&h as &dyn Defaulted).is_big()); + let h2 = DefHot { val: 50 }; + assert!(!(&h2 as &dyn Defaulted).is_big()); + + // Cold type via &dyn Trait + let c = DefCold { val: 200 }; + assert!((&c as &dyn Defaulted).is_big()); + + // Via &(dyn Trait + Send) + let h3 = DefHot { val: 200 }; + assert!((&h3 as &(dyn Defaulted + Send)).is_big()); + let c2 = DefCold { val: 50 }; + assert!(!(&c2 as &(dyn Defaulted + Send)).is_big()); + + // Via &(dyn Trait + Send + Sync) + let h4 = DefHot { val: 200 }; + assert!((&h4 as &(dyn Defaulted + Send + Sync)).is_big()); +} + // ── Extended proc-macro tests: supertraits, method lifetimes, #[must_use] ── #[cfg(feature = "macros")] diff --git a/crates/core/tests/ui_attr.rs b/crates/core/tests/ui_attr.rs index 5301ac6..2913e4a 100644 --- a/crates/core/tests/ui_attr.rs +++ b/crates/core/tests/ui_attr.rs @@ -13,6 +13,9 @@ fn ui_attr() { t.pass("tests/ui_attr/attr_supertraits.rs"); t.pass("tests/ui_attr/attr_must_use.rs"); t.pass("tests/ui_attr/attr_dyn_send.rs"); + t.pass("tests/ui_attr/attr_default_body.rs"); + t.pass("tests/ui_attr/attr_default_override.rs"); + t.pass("tests/ui_attr/attr_default_send.rs"); t.compile_fail("tests/ui_attr/attr_must_use_unused.rs"); t.compile_fail("tests/ui_attr/attr_missing_args.rs"); t.compile_fail("tests/ui_attr/attr_unsafe_missing_on_impl.rs"); diff --git a/crates/core/tests/ui_attr/attr_default_body.rs b/crates/core/tests/ui_attr/attr_default_body.rs new file mode 100644 index 0000000..befe207 --- /dev/null +++ b/crates/core/tests/ui_attr/attr_default_body.rs @@ -0,0 +1,60 @@ +use std::fmt::Write; + +struct Hot { + val: f64, +} + +struct Cold { + val: f64, +} + +#[devirt::devirt(Hot)] +pub trait Shape { + fn area(&self) -> f64; + fn is_large(&self) -> bool { + self.area() > 100.0 + } + fn describe(&self) -> String { + let mut s = String::new(); + if self.is_large() { + write!(s, "large (area={})", self.area()).ok(); + } else { + write!(s, "small (area={})", self.area()).ok(); + } + s + } +} + +#[devirt::devirt] +impl Shape for Hot { + fn area(&self) -> f64 { + self.val + } +} + +#[devirt::devirt] +impl Shape for Cold { + fn area(&self) -> f64 { + self.val + 1.0 + } +} + +fn main() { + // Hot type, small + let h = Hot { val: 50.0 }; + let d: &dyn Shape = &h; + assert!(!d.is_large()); + assert!(d.describe().contains("small")); + + // Hot type, large + let big = Hot { val: 200.0 }; + let d2: &dyn Shape = &big; + assert!(d2.is_large()); + assert!(d2.describe().contains("large")); + + // Cold type + let c = Cold { val: 150.0 }; + let d3: &dyn Shape = &c; + assert!(d3.is_large()); + assert!(d3.describe().contains("large")); +} diff --git a/crates/core/tests/ui_attr/attr_default_override.rs b/crates/core/tests/ui_attr/attr_default_override.rs new file mode 100644 index 0000000..11f7c2a --- /dev/null +++ b/crates/core/tests/ui_attr/attr_default_override.rs @@ -0,0 +1,27 @@ +struct Hot { + val: f64, +} + +#[devirt::devirt(Hot)] +pub trait Shape { + fn area(&self) -> f64; + fn is_large(&self) -> bool { + self.area() > 100.0 + } +} + +#[devirt::devirt] +impl Shape for Hot { + fn area(&self) -> f64 { + self.val + } + fn is_large(&self) -> bool { + false // override default + } +} + +fn main() { + let h = Hot { val: 200.0 }; + let d: &dyn Shape = &h; + assert!(!d.is_large()); // overridden to always false +} diff --git a/crates/core/tests/ui_attr/attr_default_send.rs b/crates/core/tests/ui_attr/attr_default_send.rs new file mode 100644 index 0000000..6bb6432 --- /dev/null +++ b/crates/core/tests/ui_attr/attr_default_send.rs @@ -0,0 +1,30 @@ +struct Hot { + val: f64, +} + +#[devirt::devirt(Hot)] +pub trait Shape { + fn area(&self) -> f64; + fn is_large(&self) -> bool { + self.area() > 100.0 + } +} + +#[devirt::devirt] +impl Shape for Hot { + fn area(&self) -> f64 { + self.val + } +} + +fn check(s: &(dyn Shape + Send)) -> bool { + s.is_large() +} + +fn main() { + let h = Hot { val: 200.0 }; + assert!(check(&h)); + + let small = Hot { val: 50.0 }; + assert!(!check(&small)); +} diff --git a/crates/macros/src/lib.rs b/crates/macros/src/lib.rs index 92c5e19..0080eb9 100644 --- a/crates/macros/src/lib.rs +++ b/crates/macros/src/lib.rs @@ -5,9 +5,12 @@ //! is an implementation detail of `devirt` and should not be used //! directly. +use std::collections::HashSet; + use proc_macro::TokenStream; use quote::{format_ident, quote}; use syn::punctuated::Punctuated; +use syn::visit_mut::VisitMut; use syn::{Token, parse_macro_input}; /// Proc-macro attribute for transparent devirtualization. @@ -107,12 +110,6 @@ fn validate_trait(trait_item: &syn::ItemTrait) -> Result<(), syn::Error> { } fn validate_trait_method(f: &syn::TraitItemFn) -> Result<(), syn::Error> { - if f.default.is_some() { - return Err(syn::Error::new_spanned( - f, - "#[devirt] does not support default method bodies", - )); - } if f.sig.asyncness.is_some() { return Err(syn::Error::new_spanned( &f.sig, @@ -164,20 +161,9 @@ fn emit_trait_expansion( let supertraits = &trait_item.supertraits; let inner_name = format_ident!("__{name}Impl"); - // __spec_* method declarations for the inner trait. - let spec_decls: Vec<_> = trait_item - .items - .iter() - .filter_map(|item| { - let syn::TraitItem::Fn(m) = item else { - return None; - }; - let mut spec_sig = m.sig.clone(); - spec_sig.ident = format_ident!("__spec_{}", spec_sig.ident); - let attrs = &m.attrs; - Some(quote! { #(#attrs)* #spec_sig; }) - }) - .collect(); + // __spec_* method declarations for the inner trait (with default + // bodies rewritten so `self.method()` → `self.__spec_method()`). + let spec_decls = generate_spec_decls(trait_item); // Dispatch methods for the inherent impl on `dyn Trait`. let dispatch_methods: Vec<_> = trait_item @@ -297,6 +283,55 @@ fn emit_trait_expansion( .into() } +// ── Default-body spec declarations ───────────────────────────────────────── + +/// Generate `__spec_*` method declarations for the inner trait. +/// +/// Methods without a default body become required (`__spec_foo(...);`). +/// Methods with a default body get the body rewritten so that +/// `self.method()` calls become `self.__spec_method()`, then emitted +/// as provided methods on the inner trait. +fn generate_spec_decls( + trait_item: &syn::ItemTrait, +) -> Vec { + let method_names: HashSet = trait_item + .items + .iter() + .filter_map(|item| { + if let syn::TraitItem::Fn(m) = item { + Some(m.sig.ident.to_string()) + } else { + None + } + }) + .collect(); + + trait_item + .items + .iter() + .filter_map(|item| { + let syn::TraitItem::Fn(m) = item else { + return None; + }; + let mut spec_sig = m.sig.clone(); + spec_sig.ident = format_ident!("__spec_{}", spec_sig.ident); + let attrs = &m.attrs; + + m.default.as_ref().map_or_else( + || Some(quote! { #(#attrs)* #spec_sig; }), + |default_body| { + let mut body = default_body.clone(); + let mut rewriter = RewriteSelfCalls { + method_names: method_names.clone(), + }; + rewriter.visit_block_mut(&mut body); + Some(quote! { #(#attrs)* #spec_sig #body }) + }, + ) + }) + .collect() +} + // ── Shared helpers ────────────────────────────────────────────────────────── /// Clone a method signature, replacing wildcard `_` patterns with @@ -593,3 +628,89 @@ fn expand_impl(attr: &TokenStream, impl_item: &syn::ItemImpl) -> TokenStream { } .into() } + +// ── AST rewriting for default method bodies ──────────────────────────────── + +/// Rewrites `self.method()` calls in default method bodies to +/// `self.__spec_method()` so they resolve within the inner trait. +struct RewriteSelfCalls { + /// Original method names defined on this trait. + method_names: HashSet, +} + +impl VisitMut for RewriteSelfCalls { + fn visit_expr_method_call_mut(&mut self, i: &mut syn::ExprMethodCall) { + // Recurse into sub-expressions first. + syn::visit_mut::visit_expr_method_call_mut(self, i); + + // Only rewrite direct `self.method()` calls. + if is_self_expr(&i.receiver) { + let name = i.method.to_string(); + if self.method_names.contains(&name) { + i.method = format_ident!("__spec_{}", i.method); + } + } + } + + fn visit_macro_mut(&mut self, i: &mut syn::Macro) { + // `syn::visit_mut` does not descend into macro token streams, + // so we do a token-level rewrite for `self.method()` patterns + // inside macro invocations (e.g. `format!`, `write!`). + i.tokens = + rewrite_self_calls_in_tokens(&self.method_names, i.tokens.clone()); + } +} + +fn is_self_expr(expr: &syn::Expr) -> bool { + match expr { + syn::Expr::Path(p) => p.path.is_ident("self"), + // Handle `(self)` — parenthesized. + syn::Expr::Paren(p) => is_self_expr(&p.expr), + _ => false, + } +} + +/// Token-level rewrite of `self . method_name` → `self . __spec_method_name` +/// inside macro invocations, where `syn::visit_mut` cannot descend. +fn rewrite_self_calls_in_tokens( + method_names: &HashSet, + tokens: proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + let tts: Vec = tokens.into_iter().collect(); + let mut out = Vec::with_capacity(tts.len()); + let mut i = 0; + while i < tts.len() { + // Match pattern: Ident("self") Punct('.') Ident(method_name) + if i + 2 < tts.len() + && let proc_macro2::TokenTree::Ident(ref id) = tts[i] + && *id == "self" + && let proc_macro2::TokenTree::Punct(ref dot) = tts[i + 1] + && dot.as_char() == '.' + && let proc_macro2::TokenTree::Ident(ref method) = tts[i + 2] + && method_names.contains(&method.to_string()) + { + out.push(tts[i].clone()); + out.push(tts[i + 1].clone()); + out.push(proc_macro2::TokenTree::Ident( + proc_macro2::Ident::new( + &format!("__spec_{method}"), + method.span(), + ), + )); + i += 3; + continue; + } + // Recurse into groups (parenthesized, braced, bracketed). + if let proc_macro2::TokenTree::Group(ref g) = tts[i] { + let inner = + rewrite_self_calls_in_tokens(method_names, g.stream()); + let mut ng = proc_macro2::Group::new(g.delimiter(), inner); + ng.set_span(g.span()); + out.push(proc_macro2::TokenTree::Group(ng)); + } else { + out.push(tts[i].clone()); + } + i += 1; + } + out.into_iter().collect() +} From d73ec0354affe006e725ca32167c2c163e836c93 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 12 Apr 2026 18:03:43 +0000 Subject: [PATCH 2/2] Rewrite sibling calls in impl bodies; extend equivalence tests Apply the same RewriteSelfCalls transformation to impl method bodies in expand_impl, so self.method() inside a #[devirt] impl block is rewritten to self.__spec_method(). Without this, sibling calls in impl bodies fail to compile because the public trait is an empty marker. Extend attr_default_body_dispatch with three additional runtime cases: - DefOver: overrides the default is_big with a sibling call (self.get() > 1000), exercising impl-body rewriting - DefFmt: relies on the default describe() that uses format!() with self.get(), exercising macro-invocation token rewriting - Assertions via &dyn Defaulted and &(dyn Defaulted + Send) https://claude.ai/code/session_01BJen3TbKCFohcjNFeYwBiA --- crates/core/tests/equivalence.rs | 50 +++++++++++++++++++++++++++++++- crates/macros/src/lib.rs | 20 ++++++++++++- 2 files changed, 68 insertions(+), 2 deletions(-) diff --git a/crates/core/tests/equivalence.rs b/crates/core/tests/equivalence.rs index d8c6037..2bbc63d 100644 --- a/crates/core/tests/equivalence.rs +++ b/crates/core/tests/equivalence.rs @@ -156,12 +156,25 @@ mod attr_defaults { pub val: u64, } + /// Overrides the default `is_big`. + pub struct DefOver { + pub val: u64, + } + + /// Relies on the `describe` default that uses `format!`. + pub struct DefFmt { + pub val: u64, + } + #[devirt::devirt(DefHot)] pub trait Defaulted { fn get(&self) -> u64; fn is_big(&self) -> bool { self.get() > 100 } + fn describe(&self) -> String { + format!("val={}", self.get()) + } } #[devirt::devirt] @@ -177,12 +190,31 @@ mod attr_defaults { self.val + 1 } } + + #[devirt::devirt] + impl Defaulted for DefOver { + fn get(&self) -> u64 { + self.val + } + fn is_big(&self) -> bool { + // Exercises sibling-call rewriting inside impl bodies: + // self.get() must be rewritten to self.__spec_get(). + self.get() > 1000 + } + } + + #[devirt::devirt] + impl Defaulted for DefFmt { + fn get(&self) -> u64 { + self.val + } + } } #[cfg(feature = "macros")] #[test] fn attr_default_body_dispatch() { - use attr_defaults::{DefCold, DefHot, Defaulted}; + use attr_defaults::{DefCold, DefFmt, DefHot, DefOver, Defaulted}; // Hot type via &dyn Trait let h = DefHot { val: 200 }; @@ -203,6 +235,22 @@ fn attr_default_body_dispatch() { // Via &(dyn Trait + Send + Sync) let h4 = DefHot { val: 200 }; assert!((&h4 as &(dyn Defaulted + Send + Sync)).is_big()); + + // Overridden default: DefOver uses `self.get() > 1000` (tests + // sibling-call rewriting in impl bodies). + let o = DefOver { val: 200 }; + assert!(!(&o as &dyn Defaulted).is_big()); + assert!(!(&o as &(dyn Defaulted + Send)).is_big()); + let o2 = DefOver { val: 2000 }; + assert!((&o2 as &dyn Defaulted).is_big()); + + // Default body with write! macro (exercises token-level rewriting) + let f = DefFmt { val: 42 }; + assert_eq!((&f as &dyn Defaulted).describe(), "val=42"); + assert_eq!((&f as &(dyn Defaulted + Send)).describe(), "val=42"); + // Hot type's describe (also via default body) + let h5 = DefHot { val: 7 }; + assert_eq!((&h5 as &dyn Defaulted).describe(), "val=7"); } // ── Extended proc-macro tests: supertraits, method lifetimes, #[must_use] ── diff --git a/crates/macros/src/lib.rs b/crates/macros/src/lib.rs index 0080eb9..175a9ef 100644 --- a/crates/macros/src/lib.rs +++ b/crates/macros/src/lib.rs @@ -601,6 +601,20 @@ fn expand_impl(attr: &TokenStream, impl_item: &syn::ItemImpl) -> TokenStream { let inner_name = format_ident!("__{trait_name}Impl"); let ty = &impl_item.self_ty; + // Collect method names so sibling calls in impl bodies + // (e.g. `self.area()`) are rewritten to `self.__spec_area()`. + let method_names: HashSet = impl_item + .items + .iter() + .filter_map(|item| { + if let syn::ImplItem::Fn(m) = item { + Some(m.sig.ident.to_string()) + } else { + None + } + }) + .collect(); + let spec_methods: Vec<_> = impl_item .items .iter() @@ -611,7 +625,11 @@ fn expand_impl(attr: &TokenStream, impl_item: &syn::ItemImpl) -> TokenStream { let mut spec_sig = m.sig.clone(); spec_sig.ident = format_ident!("__spec_{}", spec_sig.ident); let attrs = &m.attrs; - let block = &m.block; + let mut block = m.block.clone(); + let mut rewriter = RewriteSelfCalls { + method_names: method_names.clone(), + }; + rewriter.visit_block_mut(&mut block); Some(quote! { #(#attrs)* #[inline]