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}