Skip to main content

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                .filter(|_| {
100                    field.peek_is(|tt| tt.is_punch(':'))
101                        && (!field.peek_i_is(1, |tt| tt.is_punch(':'))
102                            || field.peek_i_puncts(1, "::").is_some())
103                })
104        })
105        .collect()
106}
107
108struct Star {
109    attrs: Vec<TokenStream>,
110    ref_tt: Option<TokenTree>,
111    mut_tt: Option<TokenTree>,
112    star: TokenTree,
113}
114
115fn parse_fields_use(
116    name: Ident,
117    mut iter: ParseIter<impl Iterator<Item = TokenTree>>,
118    decl_fields: &[String],
119) -> Result<TokenStream, TokenStream> {
120    iter.next_outer_attributes();
121    let mut star_info = None;
122    let mut used_field = HashSet::new();
123
124    let mut body = iter
125        .split_puncts_all(",")
126        .map(|field| {
127            let mut iter = field.parse_iter();
128            let attrs = iter.next_attributes();
129            let ref_tt = iter.next_if(|tt| tt.is_keyword("ref"));
130            let mut_tt = iter.next_if(|tt| tt.is_keyword("mut"));
131            if let Some(star) = iter.next_if(|tt| tt.is_punch('*')) {
132                star_info = Some(Star { attrs, ref_tt, mut_tt, star });
133                Ok(None)
134            } else {
135                iter.next()
136                    .and_then(|it| it.into_ident().ok())
137                    .map(|ident| {
138                        used_field.insert(ident.to_string());
139                        streams(attrs.into_iter().chain([stream(
140                            flat([ref_tt, mut_tt])
141                                .chain([ident.tt()])
142                                .chain(iter),
143                        )]))
144                        .into()
145                    })
146                    .ok_or_else(|| err!("cannot find field"))
147            }
148        })
149        .filter_map(Result::transpose)
150        .try_join(puncts(", "))?;
151
152    if let Some(Star {
153        attrs,
154        ref_tt,
155        mut_tt,
156        star,
157    }) = star_info {
158        for field in decl_fields.iter()
159            .filter(|field| !used_field.contains(*field))
160        {
161            if !body.is_empty() {
162                body.push(','.alone().tt());
163            }
164            body.extend(attrs.iter().cloned());
165            body.extend(ref_tt.clone());
166            body.extend(mut_tt.clone());
167            body.push(field.ident(star.span()).tt());
168        }
169    }
170
171    Ok(stream([name.tt(), body.grouped_brace().tt()]))
172}
173
174fn flat<I>(iter: I) -> impl Iterator<Item = <I::Item as IntoIterator>::Item>
175where I: IntoIterator,
176      I::Item: IntoIterator,
177{
178    iter.into_iter().flatten()
179}