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}