fields_glob/
lib.rs

1#![doc = include_str!("../README.md")]
2use std::{collections::HashSet, convert::identity, fmt::Write};
3
4use proc_macro::*;
5use proc_macro_tool::{
6    err, puncts, stream, streams, ParseIter, ParseIterExt, PunctExt, SetSpan,
7    StrExt, StreamIterExt, TokenStreamExt, TokenTreeExt, WalkExt,
8};
9
10#[doc = include_str!("../README.md")]
11#[proc_macro_derive(fields_glob, attributes(fields_glob_export_macro))]
12pub fn fields_glob_derive(adt: TokenStream) -> TokenStream {
13    let mut iter = adt.parse_iter();
14    let export_macro = iter.next_attributes().into_iter().any(|attr| {
15        attr.into_iter().nth(1).is_some_and(|tt| {
16            tt.as_group()
17                .and_then(|group| group.stream().into_iter().next())
18                .is_some_and(|tt| tt.is_keyword("fields_glob_export_macro"))
19        })
20    });
21    iter.next_vis();
22
23    let strukt = iter.next().unwrap();
24    if !strukt.is_keyword("struct") {
25        return err("fields_glob only support struct", strukt);
26    }
27
28    let name = iter.next().unwrap().into_ident().unwrap();
29
30    let mut prev = None;
31    let mut last = iter.next().unwrap();
32
33    for next in iter {
34        prev = last.into();
35        last = next;
36    }
37
38    let body = if !last.is_punch(';') {
39        last
40    } else if let Some(prev) = prev {
41        prev
42    } else {
43        return err("fields_glob cannot support unit-like struct", last);
44    }.into_group().unwrap();
45
46    if !body.is_delimiter_brace() {
47        return err("fields_glob only support named field struct", body);
48    }
49
50    let mut template = export_macro
51        .then_some("#[macro_export]")
52        .unwrap_or_default()
53        .to_owned();
54    writeln!(template, "/// `{name}! {{}}` fields_glob support macro").unwrap();
55    template += r#"
56    macro_rules! % {
57        ($($t:tt)*) => {
58            ::fields_glob::fields_glob_impl! {
59                %
60                @
61                $($t)*
62            }
63        };
64    }"#;
65    template.parse::<TokenStream>().unwrap()
66        .walk(|tt| match tt.set_spaned(name.span()) {
67            tt if tt.is_punch('%') => name.clone().tt(),
68            tt if tt.is_punch('@') => body.clone().tt(),
69            tt => tt,
70        })
71}
72
73#[proc_macro]
74pub fn fields_glob_impl(input: TokenStream) -> TokenStream {
75    let mut iter = input.parse_iter();
76    let [name, fields] = iter.next_tts();
77    let TokenTree::Ident(name) = name else {
78        return err("invalid input, expected struct name", name);
79    };
80    let TokenTree::Group(fields) = fields else {
81        return err("invalid input, expected struct body", fields);
82    };
83    let fields = parse_fields_declare(fields);
84    parse_fields_use(name, iter, &fields)
85        .unwrap_or_else(identity)
86}
87
88fn parse_fields_declare(fields: Group) -> Vec<String> {
89    let mut parse_iter = fields.stream().parse_iter();
90    parse_iter.next_outer_attributes();
91    parse_iter
92        .split_puncts_all(",")
93        .filter_map(|field| {
94            let mut field = field.parse_iter();
95            field.next_attributes();
96            field.next_vis();
97            field.next()
98                .and_then(|tt| tt.as_ident().map(ToString::to_string))
99        })
100        .collect()
101}
102
103struct Star {
104    attrs: Vec<TokenStream>,
105    ref_tt: Option<TokenTree>,
106    mut_tt: Option<TokenTree>,
107    star: TokenTree,
108}
109
110fn parse_fields_use(
111    name: Ident,
112    mut iter: ParseIter<impl Iterator<Item = TokenTree>>,
113    decl_fields: &[String],
114) -> Result<TokenStream, TokenStream> {
115    iter.next_outer_attributes();
116    let mut star_info = None;
117    let mut used_field = HashSet::new();
118
119    let mut body = iter
120        .split_puncts_all(",")
121        .map(|field| {
122            let mut iter = field.parse_iter();
123            let attrs = iter.next_attributes();
124            let ref_tt = iter.next_if(|tt| tt.is_keyword("ref"));
125            let mut_tt = iter.next_if(|tt| tt.is_keyword("mut"));
126            if let Some(star) = iter.next_if(|tt| tt.is_punch('*')) {
127                star_info = Some(Star { attrs, ref_tt, mut_tt, star });
128                Ok(None)
129            } else {
130                iter.next()
131                    .and_then(|it| it.into_ident().ok())
132                    .map(|ident| {
133                        used_field.insert(ident.to_string());
134                        streams(attrs.into_iter().chain([stream(
135                            flat([ref_tt, mut_tt])
136                                .chain([ident.tt()])
137                                .chain(iter),
138                        )]))
139                        .into()
140                    })
141                    .ok_or_else(|| err!("cannot find field"))
142            }
143        })
144        .filter_map(Result::transpose)
145        .try_join(puncts(", "))?;
146
147    if let Some(Star {
148        attrs,
149        ref_tt,
150        mut_tt,
151        star,
152    }) = star_info {
153        for field in decl_fields.iter()
154            .filter(|field| !used_field.contains(*field))
155        {
156            if !body.is_empty() {
157                body.push(','.alone().tt());
158            }
159            body.extend(attrs.iter().cloned());
160            body.extend(ref_tt.clone());
161            body.extend(mut_tt.clone());
162            body.push(field.ident(star.span()).tt());
163        }
164    }
165
166    Ok(stream([name.tt(), body.grouped_brace().tt()]))
167}
168
169fn flat<I>(iter: I) -> impl Iterator<Item = <I::Item as IntoIterator>::Item>
170where I: IntoIterator,
171      I::Item: IntoIterator,
172{
173    iter.into_iter().flatten()
174}