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