1use std::collections::HashMap;
16use std::collections::hash_map::Entry;
17use std::iter::IntoIterator;
18
19use proc_macro2::Span;
20use proc_macro2::TokenStream;
21use quote::ToTokens;
22use quote::quote;
23use syn::Attribute;
24use syn::Data;
25use syn::DataEnum;
26use syn::DataStruct;
27use syn::DeriveInput;
28use syn::Error;
29use syn::Expr;
30use syn::Field;
31use syn::Fields;
32use syn::Ident;
33use syn::Lit;
34use syn::LitStr;
35use syn::Member;
36use syn::Meta;
37use syn::MetaList;
38use syn::Path;
39use syn::Result;
40use syn::Token;
41use syn::Variant;
42use syn::parse_macro_input;
43use syn::parse_quote;
44use syn::punctuated::Punctuated;
45use syn::spanned::Spanned;
46use syn::token::Mut;
47
48#[proc_macro_derive(Traversable, attributes(traverse))]
49pub fn derive_traversable(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
50 expand_with(input, |stream| impl_traversable(stream, false))
51}
52
53#[proc_macro_derive(TraversableMut, attributes(traverse))]
54pub fn derive_traversable_mut(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
55 expand_with(input, |stream| impl_traversable(stream, true))
56}
57
58fn expand_with(
59 input: proc_macro::TokenStream,
60 handler: impl Fn(DeriveInput) -> Result<TokenStream>,
61) -> proc_macro::TokenStream {
62 let input = parse_macro_input!(input as DeriveInput);
63 handler(input)
64 .unwrap_or_else(|error| error.to_compile_error())
65 .into()
66}
67
68fn extract_meta(attrs: Vec<Attribute>, attr_name: &str) -> Result<Option<Meta>> {
69 let macro_attrs = attrs
70 .into_iter()
71 .filter(|attr| attr.path().is_ident(attr_name))
72 .collect::<Vec<Attribute>>();
73
74 if let Some(second) = macro_attrs.get(2) {
75 return Err(Error::new_spanned(second, "duplicate attribute"));
76 }
77
78 macro_attrs
79 .first()
80 .map(|attr| Ok(attr.meta.clone()))
81 .transpose()
82}
83
84#[derive(Default)]
85struct Params(HashMap<Path, Meta>);
86
87impl Params {
88 fn from_attrs(attrs: Vec<Attribute>, attr_name: &str) -> Result<Self> {
89 Ok(extract_meta(attrs, attr_name)?
90 .map(|meta| {
91 if let Meta::List(meta_list) = meta {
92 Self::from_meta_list(meta_list)
93 } else {
94 Err(Error::new_spanned(meta, "invalid attribute"))
95 }
96 })
97 .transpose()?
98 .unwrap_or_default())
99 }
100
101 fn from_meta_list(meta_list: MetaList) -> Result<Self> {
102 let mut params = HashMap::new();
103 let nested = meta_list.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)?;
104 for meta in nested {
105 let path = meta.path();
106 let entry = params.entry(path.clone());
107 if matches!(entry, Entry::Occupied(_)) {
108 return Err(Error::new_spanned(path, "duplicate parameter"));
109 }
110 entry.or_insert(meta);
111 }
112 Ok(Self(params))
113 }
114
115 fn validate(&self, allowed_params: &[&str]) -> Result<()> {
116 for path in self.0.keys() {
117 if !allowed_params
118 .iter()
119 .any(|allowed_param| path.is_ident(allowed_param))
120 {
121 return Err(Error::new_spanned(
122 path,
123 format!(
124 "unknown parameter, supported: {}",
125 allowed_params.join(", ")
126 ),
127 ));
128 }
129 }
130 Ok(())
131 }
132
133 fn param(&mut self, name: &str) -> Result<Option<Param>> {
134 self.0
135 .remove(&Ident::new(name, Span::call_site()).into())
136 .map(Param::from_meta)
137 .transpose()
138 }
139}
140
141impl Iterator for Params {
142 type Item = Result<Param>;
143 fn next(&mut self) -> Option<Self::Item> {
144 self.0
145 .keys()
146 .next()
147 .cloned()
148 .map(|path| Param::from_meta(self.0.remove(&path).unwrap()))
149 }
150}
151
152enum Param {
153 Unit(Span),
154 StringLiteral(Span, LitStr),
155 NestedParams(Span),
156}
157
158impl Param {
159 fn from_meta(meta: Meta) -> Result<Self> {
160 let span = meta.span();
161 match meta {
162 Meta::Path(_) => Ok(Param::Unit(span)),
163 Meta::List(_) => Ok(Param::NestedParams(span)),
164 Meta::NameValue(name_value) => {
165 if let Expr::Lit(expr_lit) = &name_value.value {
166 if let Lit::Str(lit_str) = &expr_lit.lit {
167 Ok(Param::StringLiteral(span, lit_str.clone()))
168 } else {
169 Err(Error::new_spanned(name_value, "invalid parameter"))
170 }
171 } else {
172 Err(Error::new_spanned(name_value, "invalid parameter"))
173 }
174 }
175 }
176 }
177
178 fn span(&self) -> Span {
179 match self {
180 Self::Unit(span) | Self::StringLiteral(span, _) | Self::NestedParams(span) => *span,
181 }
182 }
183
184 fn unit(self) -> Result<()> {
185 if let Self::Unit(_) = self {
186 Ok(())
187 } else {
188 Err(Error::new(self.span(), "invalid parameter"))
189 }
190 }
191
192 fn string_literal(self) -> Result<LitStr> {
193 if let Self::StringLiteral(_, lit_str) = self {
194 Ok(lit_str)
195 } else {
196 Err(Error::new(self.span(), "invalid parameter"))
197 }
198 }
199}
200
201#[inline(always)]
202fn resolve_crate_name() -> Path {
203 parse_quote!(::traversable)
204}
205
206fn impl_traversable(input: DeriveInput, mutable: bool) -> Result<TokenStream> {
207 let mut params = Params::from_attrs(input.attrs, "traverse")?;
208 params.validate(&["skip_self", "skip_children"])?;
209
210 let skip_visit_self = params
211 .param("skip_self")?
212 .map(Param::unit)
213 .transpose()?
214 .is_some();
215 let skip_children = params
216 .param("skip_children")?
217 .map(Param::unit)
218 .transpose()?
219 .is_some();
220
221 let name = input.ident;
222 let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
223
224 let visitor = Ident::new(
225 if mutable { "VisitorMut" } else { "Visitor" },
226 Span::call_site(),
227 );
228
229 let enter_method = Ident::new(
230 if mutable { "enter_mut" } else { "enter" },
231 Span::call_site(),
232 );
233
234 let leave_method = Ident::new(
235 if mutable { "leave_mut" } else { "leave" },
236 Span::call_site(),
237 );
238
239 let crate_name = resolve_crate_name();
240
241 let enter_self = if skip_visit_self {
242 None
243 } else {
244 Some(quote! {
245 #crate_name::#visitor::#enter_method(visitor, self)?;
246 })
247 };
248
249 let leave_self = if skip_visit_self {
250 None
251 } else {
252 Some(quote! {
253 #crate_name::#visitor::#leave_method(visitor, self)?;
254 })
255 };
256
257 let traverse_fields = match input.data {
258 Data::Struct(struct_) => {
259 if skip_children {
260 Ok(TokenStream::new())
261 } else {
262 traverse_struct(struct_, mutable)
263 }
264 }
265 Data::Enum(enum_) => {
266 if skip_children {
267 Ok(TokenStream::new())
268 } else {
269 traverse_enum(enum_, mutable)
270 }
271 }
272 Data::Union(union_) => {
273 return Err(Error::new_spanned(
274 union_.union_token,
275 "unions are not supported",
276 ));
277 }
278 }?;
279
280 let impl_trait = Ident::new(
281 if mutable {
282 "TraversableMut"
283 } else {
284 "Traversable"
285 },
286 Span::call_site(),
287 );
288
289 let method = Ident::new(
290 if mutable { "traverse_mut" } else { "traverse" },
291 Span::call_site(),
292 );
293
294 let mut_modifier = if mutable {
295 Some(Mut(Span::call_site()))
296 } else {
297 None
298 };
299
300 Ok(quote! {
301 impl #impl_generics #crate_name::#impl_trait for #name #ty_generics #where_clause {
302 fn #method<V: #crate_name::#visitor>(
303 & #mut_modifier self,
304 visitor: &mut V
305 ) -> ::core::ops::ControlFlow<V::Break> {
306 #enter_self
307 #traverse_fields
308 #leave_self
309 ::core::ops::ControlFlow::Continue(())
310 }
311 }
312 })
313}
314
315fn traverse_struct(s: DataStruct, mutable: bool) -> Result<TokenStream> {
316 s.fields
317 .into_iter()
318 .enumerate()
319 .map(|(index, field)| {
320 let member = field.ident.as_ref().map_or_else(
321 || Member::Unnamed(index.into()),
322 |ident| Member::Named(ident.clone()),
323 );
324 let mut_modifier = if mutable {
325 Some(Mut(Span::call_site()))
326 } else {
327 None
328 };
329 traverse_field("e! { & #mut_modifier self.#member }, field, mutable)
330 })
331 .collect()
332}
333
334fn traverse_enum(e: DataEnum, mutable: bool) -> Result<TokenStream> {
335 let variants = e
336 .variants
337 .into_iter()
338 .map(|x| traverse_variant(x, mutable))
339 .collect::<Result<TokenStream>>()?;
340 Ok(quote! {
341 match self {
342 #variants
343 _ => {}
344 }
345 })
346}
347
348fn traverse_variant(v: Variant, mutable: bool) -> Result<TokenStream> {
349 let mut params = Params::from_attrs(v.attrs, "traverse")?;
350 params.validate(&["skip"])?;
351 if params.param("skip")?.map(Param::unit).is_some() {
352 return Ok(TokenStream::new());
353 }
354 let name = v.ident;
355 let destructuring = destructure_fields(v.fields.clone())?;
356 let fields = v
357 .fields
358 .into_iter()
359 .enumerate()
360 .map(|(index, field)| {
361 traverse_field(
362 &field
363 .ident
364 .clone()
365 .unwrap_or_else(|| Ident::new(&format!("i{}", index), Span::call_site()))
366 .to_token_stream(),
367 field,
368 mutable,
369 )
370 })
371 .collect::<Result<TokenStream>>()?;
372 Ok(quote! {
373 Self::#name #destructuring => {
374 #fields
375 }
376 })
377}
378
379fn destructure_fields(fields: Fields) -> Result<TokenStream> {
380 Ok(match fields {
381 Fields::Named(fields) => {
382 let field_list = fields
383 .named
384 .into_iter()
385 .map(|field| {
386 let mut params = Params::from_attrs(field.attrs, "traverse")?;
387 let field_name = field.ident.unwrap();
388 Ok(if params.param("skip")?.map(Param::unit).is_some() {
389 quote! { #field_name: _ }
390 } else {
391 field_name.into_token_stream()
392 })
393 })
394 .collect::<Result<Vec<TokenStream>>>()?;
395 quote! {
396 { #( #field_list ),* }
397 }
398 }
399 Fields::Unnamed(fields) => {
400 let field_list = fields
401 .unnamed
402 .into_iter()
403 .enumerate()
404 .map(|(index, field)| {
405 let mut params = Params::from_attrs(field.attrs, "traverse")?;
406 Ok(if params.param("skip")?.map(Param::unit).is_some() {
407 quote! { _ }
408 } else {
409 Ident::new(&format!("i{index}",), Span::call_site()).into_token_stream()
410 })
411 })
412 .collect::<Result<Vec<TokenStream>>>()?;
413 quote! {
414 ( #( #field_list ),* )
415 }
416 }
417 Fields::Unit => TokenStream::new(),
418 })
419}
420
421fn traverse_field(value: &TokenStream, field: Field, mutable: bool) -> Result<TokenStream> {
422 let mut params = Params::from_attrs(field.attrs, "traverse")?;
423 params.validate(&["skip", "with"])?;
424
425 if params.param("skip")?.map(Param::unit).is_some() {
426 return Ok(TokenStream::new());
427 }
428
429 let crate_name = resolve_crate_name();
430
431 match params.param("with")? {
432 None => Ok(if mutable {
433 quote! { #crate_name::TraversableMut::traverse_mut(#value, visitor)?; }
434 } else {
435 quote! { #crate_name::Traversable::traverse(#value, visitor)?; }
436 }),
437 Some(traverse_fn) => {
438 let traverse_fn = traverse_fn.string_literal()?.parse::<Path>()?;
439 Ok(quote! {
440 #traverse_fn(#value, visitor)?;
441 })
442 }
443 }
444}