1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{Data, DeriveInput, Field, Type, parse_quote};
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6enum CfgDisplayAttr {
7    DebugFormat,
8    IndentNested,
9    Skip,
10    Path,
11}
12
13fn parse_cfg_display_attrs(field: &Field) -> Vec<CfgDisplayAttr> {
14    let Some(cfg_display_attr) =
15        field.attrs.iter().find(|attr| attr.path().is_ident("cfg_display"))
16    else {
17        return vec![];
18    };
19
20    let mut attrs = Vec::new();
21    cfg_display_attr
22        .parse_nested_meta(|meta| {
23            if meta.path.is_ident("debug_fmt") {
24                attrs.push(CfgDisplayAttr::DebugFormat);
25            } else if meta.path.is_ident("indent_nested") {
26                attrs.push(CfgDisplayAttr::IndentNested);
27            } else if meta.path.is_ident("skip") {
28                attrs.push(CfgDisplayAttr::Skip);
29            } else if meta.path.is_ident("path") {
30                attrs.push(CfgDisplayAttr::Path);
31            } else {
32                return Err(meta.error("Invalid cfg_display meta"));
33            }
34
35            Ok(())
36        })
37        .expect("Failed to parse cfg_display field attribute");
38
39    attrs
40}
41
42pub fn config_display(input: TokenStream) -> TokenStream {
43    let input: DeriveInput = syn::parse(input).expect("Unable to parse input");
44
45    let Data::Struct(struct_data) = input.data else {
46        panic!("ConfigDisplay derive macro only applies to structs");
47    };
48
49    let fields: Vec<_> = struct_data
50        .fields
51        .iter()
52        .map(|field| (field, parse_cfg_display_attrs(field)))
53        .filter(|(_, attrs)| !attrs.contains(&CfgDisplayAttr::Skip))
54        .collect();
55
56    assert!(!fields.is_empty(), "ConfigDisplay derive macro only applies to structs with fields");
57
58    let writeln_statements: Vec<_> = fields
59        .iter()
60        .enumerate()
61        .map(|(i, (field, attrs))| {
62            let Some(field_ident) = &field.ident else {
63                panic!("ConfigDisplay derive macro only supports structs with named fields");
64            };
65
66            let debug_fmt = attrs.contains(&CfgDisplayAttr::DebugFormat);
67            let fmt_string = if debug_fmt {
68                format!("  {field_ident}: {{:?}}")
69            } else {
70                format!("  {field_ident}: {{}}")
71            };
72
73            let is_option = match &field.ty {
74                Type::Path(path) => {
75                    let first_segment = path.path.segments.iter().next();
76                    first_segment.is_some_and(|segment| segment.ident == "Option")
77                }
78                _ => false,
79            };
80            let is_path = attrs.contains(&CfgDisplayAttr::Path);
81
82            let format_invocation = if is_option {
83                let none_str = format!("  {field_ident}: <None>");
84
85                let field_display = if is_path {
86                    quote! { f.display() }
87                } else {
88                    quote! { f }
89                };
90
91                quote! {
92                    self.#field_ident.as_ref().map(|f| format!(#fmt_string, #field_display))
93                        .unwrap_or_else(|| #none_str.into())
94                }
95            } else if is_path {
96                quote! {
97                    format!(#fmt_string, self.#field_ident.display())
98                }
99            } else {
100                quote! {
101                    format!(#fmt_string, self.#field_ident)
102                }
103            };
104
105            let indent_nested = attrs.contains(&CfgDisplayAttr::IndentNested);
106            let format_invocation = if indent_nested {
107                quote! {
108                    #format_invocation.replace("\n  ", "\n    ")
109                }
110            } else {
111                format_invocation
112            };
113
114            if i == fields.len() - 1 {
115                quote! {
116                    ::std::write!(f, "{}", #format_invocation)
117                }
118            } else {
119                quote! {
120                    ::std::writeln!(f, "{}", #format_invocation)?;
121                }
122            }
123        })
124        .collect();
125
126    let mut generics = input.generics.clone();
127    for type_param in generics.type_params_mut() {
128        type_param.bounds.push(parse_quote!(::std::fmt::Display));
129        type_param.bounds.push(parse_quote!(::std::fmt::Debug));
130    }
131    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
132
133    let struct_ident = &input.ident;
134    let expanded = quote! {
135        impl #impl_generics ::std::fmt::Display for #struct_ident #type_generics #where_clause {
136            fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
137                ::std::writeln!(f)?;
138                #(#writeln_statements)*
139            }
140        }
141    };
142
143    expanded.into()
144}