1use darling::{Error, FromField};
17use proc_macro::TokenStream;
18use proc_macro2::{Span, TokenStream as TokenStream2};
19use proc_macro_crate::{crate_name, FoundCrate};
20use quote::quote;
21use syn::{parse_macro_input, Data, DeriveInput, Field, Fields, Ident, Index, Member, Variant};
22
23#[proc_macro_derive(TreeDisplay, attributes(tree))]
34pub fn derive_tree_display(tokens: TokenStream) -> TokenStream {
35 derive(tokens)
36}
37
38fn field_to_member(index: usize, field: &Field) -> Member {
42 match &field.ident {
43 Some(ident) => Member::Named(ident.clone()),
44 None => Member::Unnamed(Index::from(index)),
45 }
46}
47
48fn string_to_ident(name: impl AsRef<str>) -> Ident {
50 Ident::new(name.as_ref(), Span::call_site())
51}
52
53#[derive(Debug, Default, FromField)]
55#[darling(attributes(tree), default, and_then = Self::validate)]
56struct FieldAttributes {
57 map: bool,
59 ignore: bool,
61 unlabeled: bool,
63 label: Option<String>,
65}
66
67impl FieldAttributes {
68 fn validate(self) -> darling::Result<Self> {
70 if self.unlabeled && self.label.is_some() {
71 return Err(darling::Error::custom(
72 "`unlabeled` and `label` cannot be used together",
73 ));
74 }
75 if self.ignore && (self.map || self.unlabeled || self.label.is_some()) {
76 return Err(darling::Error::custom(
77 "`ignore` cannot be combined with any other attribute",
78 ));
79 }
80 Ok(self)
81 }
82}
83
84fn derive(tokens: TokenStream) -> TokenStream {
86 let input = parse_macro_input!(tokens as DeriveInput);
87 let (impl_generics, type_generics, where_clause) = input.generics.split_for_impl();
88 let type_ident = input.ident;
89
90 let is_doctest = std::env::var_os("UNSTABLE_RUSTDOC_TEST_PATH").is_some();
99
100 let ccrate = match crate_name("tree-display") {
101 Ok(FoundCrate::Itself) if !is_doctest => quote!(crate),
102 Ok(FoundCrate::Itself) => quote!(::tree_display),
103 Ok(FoundCrate::Name(name)) => {
104 let ident = Ident::new(&name, Span::call_site());
105 quote!(::#ident)
106 }
107 Err(_) => quote!(::tree_display),
108 };
109
110 let body = match input.data {
111 Data::Struct(data) => derive_struct(&type_ident, data.fields),
112 Data::Enum(data) => derive_enum(data.variants.into_iter().collect()),
113 _ => {
114 return Error::custom("TreeDisplay cannot be derived for unions")
115 .write_errors()
116 .into();
117 }
118 };
119
120 quote! {
121 impl #impl_generics #ccrate::TreeDisplay for #type_ident #type_generics #where_clause {
122 fn tree(&self, context: &#ccrate::context::Context) -> #ccrate::Tree {
123 use #ccrate::{format::{Member, TypeName}, Tree};
124 use ::std::any::TypeId;
125
126 #[allow(dead_code)] const fn type_of<T: ?Sized + 'static>(_: &T) -> TypeId {
128 TypeId::of::<T>()
129 }
130
131 #body
132 }
133 }
134 }
135 .into()
136}
137
138fn derive_struct(type_ident: &Ident, fields: Fields) -> TokenStream2 {
140 let type_name = type_ident.to_string();
141
142 let is_empty_type = matches!(&fields, Fields::Unit)
143 || matches!(&fields, Fields::Unnamed(unnamed) if unnamed.unnamed.is_empty())
144 || matches!(&fields, Fields::Named(named) if named.named.is_empty());
145
146 if is_empty_type {
147 return quote!(Tree::leaf(TypeName::new(#type_name)));
148 }
149
150 let is_newtype = matches!(fields, Fields::Unnamed(_)) && fields.len() == 1;
151 if is_newtype {
152 let field = fields.iter().next().expect("there is exactly one element");
153 let member = field_to_member(0, field);
154 let attrs = match FieldAttributes::from_field(field) {
155 Ok(attrs) => attrs,
156 Err(err) => return err.write_errors().into(),
157 };
158
159 if attrs.map {
160 return quote! {
161 if let Some(mapper) = context.mappers.get(&type_of(&self.#member)) {
162 Tree::leaf(mapper(&self.#member))
163 }
164 else
165 {
166 self.#member.tree(context)
167 }
168 };
169 }
170 return quote!(self.#member.tree(context));
171 }
172
173 let fields = match fields {
174 Fields::Unit => unreachable!(),
175 Fields::Named(named_fields) => named_fields.named,
176 Fields::Unnamed(unnamed_fields) => unnamed_fields.unnamed,
177 };
178
179 let field_handlers = match fields
180 .iter()
181 .enumerate()
182 .map(|(idx, field)| {
183 let member = field_to_member(idx, field);
184 let attrs = FieldAttributes::from_field(field)?;
185 Ok(process_field(quote!(self.#member), member, attrs))
186 })
187 .collect::<darling::Result<Vec<_>>>()
188 {
189 Ok(handlers) => handlers,
190 Err(err) => return err.write_errors().into(),
191 };
192
193 quote! {
194 let mut subtrees = ::std::vec::Vec::<Tree>::new();
195 #(#field_handlers)*
196 Tree::new(TypeName::new(#type_name), subtrees)
197 }
198}
199
200fn derive_enum(variants: Vec<Variant>) -> TokenStream2 {
202 let mut arms = Vec::new();
203 for variant in variants {
204 let variant_ident = variant.ident;
205 let variant_name = variant_ident.to_string();
206
207 match variant.fields {
208 Fields::Unit => {
209 arms.push(quote! {
210 Self::#variant_ident => Tree::leaf(TypeName::new(#variant_name))
211 });
212 }
213
214 Fields::Named(fields) => {
215 let bindings = fields
216 .named
217 .iter()
218 .map(|f| f.ident.clone().unwrap())
219 .collect::<Vec<_>>();
220
221 let handlers = match fields
222 .named
223 .iter()
224 .enumerate()
225 .map(|(idx, field)| {
226 let attrs = FieldAttributes::from_field(field)?;
227 let ident = if attrs.ignore {
228 &string_to_ident("_")
229 } else {
230 &bindings[idx]
231 };
232 Ok(process_field(
233 quote!(#ident),
234 Member::Named(ident.clone()),
235 attrs,
236 ))
237 })
238 .collect::<darling::Result<Vec<_>>>()
239 {
240 Ok(x) => x,
241 Err(err) => return err.write_errors(),
242 };
243
244 arms.push(quote! {
245 Self::#variant_ident{ #( #bindings ),* } => {
246 let mut subtrees = ::std::vec::Vec::<Tree>::new();
247
248 #(#handlers)*
249
250 Tree::new(TypeName::new(#variant_name), subtrees)
251 }
252 });
253 }
254
255 Fields::Unnamed(fields) => {
256 let is_newtype = fields.unnamed.len() == 1;
257 if is_newtype {
258 let ident = string_to_ident("__field0");
260 let field = fields
261 .unnamed
262 .first()
263 .expect("there is exactly one element");
264
265 let attrs = match FieldAttributes::from_field(field) {
266 Ok(attrs) => attrs,
267 Err(err) => return err.write_errors().into(),
268 };
269
270 if attrs.map {
271 arms.push(quote! {
272 Self::#variant_ident( #ident ) => {
273 if let Some(mapper) = context.mappers.get(&type_of(&#ident)) {
274 Tree::leaf(mapper(&#ident))
275 }
276 else
277 {
278 #ident.tree(context)
279 }
280 }
281 });
282 continue;
283 }
284
285 arms.push(quote! {
286 Self::#variant_ident( #ident ) => {
287 #ident.tree(context)
288 }
289 });
290 continue;
291 }
292
293 let bindings = (0..fields.unnamed.len())
294 .map(|i| string_to_ident(format!("__field{i}")))
295 .collect::<Vec<_>>();
296
297 let handlers = match fields
298 .unnamed
299 .iter()
300 .enumerate()
301 .map(|(idx, field)| {
302 let attrs = FieldAttributes::from_field(field)?;
303 let ident = &bindings[idx];
304 Ok(process_field(
305 quote! {#ident},
306 Member::Unnamed(Index::from(idx)),
307 attrs,
308 ))
309 })
310 .collect::<darling::Result<Vec<_>>>()
311 {
312 Ok(x) => x,
313 Err(err) => return err.write_errors(),
314 };
315
316 arms.push(quote! {
317 Self::#variant_ident( #( #bindings ),* ) => {
318 let mut subtrees = ::std::vec::Vec::<Tree>::new();
319
320 #(#handlers)*
321
322 Tree::new(TypeName::new(#variant_name), subtrees)
323 }
324 });
325 }
326 }
327 }
328
329 quote! {
330 match self {
331 #(#arms),*
332 }
333 }
334}
335
336fn process_field(
338 access: TokenStream2,
339 member: Member,
340 attributes: FieldAttributes,
341) -> TokenStream2 {
342 if attributes.ignore {
343 return quote!();
344 }
345
346 let label = match (attributes.unlabeled, attributes.label, &member) {
347 (true, _, _) => None,
348 (false, Some(label), _) => Some(label),
349 (false, _, Member::Named(ident)) => Some(ident.to_string()),
350 (false, _, Member::Unnamed(index)) => Some(format!(".{}", index.index.to_string())),
351 };
352
353 let labeled = match label {
354 Some(l) => quote! { .labeled(Member::new(#l)) },
355 None => quote! {},
356 };
357
358 if attributes.map {
359 return quote! {
360 if let Some(mapper) = context.mappers.get(&type_of(&#access)) {
361 let mapped = mapper(&#access);
362 subtrees.push(Tree::leaf(mapped)#labeled);
363 }
364 else
365 {
366 subtrees.push(#access.tree(context)#labeled);
367 }
368 };
369 }
370
371 quote! {
372 subtrees.push(#access.tree(context)#labeled);
373 }
374}