1#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/README.md"))]
2use proc_macro2::{Spacing, TokenStream, TokenTree};
3use quote::quote;
4use syn::{ItemTrait, parse_macro_input};
5
6mod apply;
7mod apply_tuple;
8mod batch_trait_entry;
9mod codegen;
10mod diagnostic;
11mod generic;
12mod parse;
13mod parse_atom;
14mod path_prefix;
15mod preprocess;
16mod preprocess_helpers;
17mod scan;
18mod types;
19mod types_render;
20
21use batch_trait_entry::parse_batch_trait_entry;
22
23use diagnostic::compile_error_str;
24use scan::Cursor;
25use types::{Op, reset_fresh_counter};
26
27#[proc_macro_attribute]
61pub fn batch_impl(
62 attr: proc_macro::TokenStream,
63 item: proc_macro::TokenStream,
64) -> proc_macro::TokenStream {
65 expand_attr_macro(attr, item, true)
66}
67
68#[proc_macro_attribute]
83pub fn batch_impl_only(
84 attr: proc_macro::TokenStream,
85 item: proc_macro::TokenStream,
86) -> proc_macro::TokenStream {
87 expand_attr_macro(attr, item, false)
88}
89
90fn expand_attr_macro(
92 attr: proc_macro::TokenStream,
93 item: proc_macro::TokenStream,
94 include_trait: bool,
95) -> proc_macro::TokenStream {
96 reset_fresh_counter();
97 let trait_item = parse_macro_input!(item as ItemTrait);
98 let trait_name = trait_item.ident.clone();
99 let attr_vec = TokenStream::from(attr).into_iter().collect::<Vec<_>>();
100
101 let (trait_full_path, trait_last_ident, rest_tokens) =
106 if !include_trait {
107 match path_prefix::try_parse_path_prefix(&attr_vec) {
108 Some((path, last_ident, rest)) => {
109 match last_ident {
112 Some(id) if id == trait_name => {
113 let path_ts: TokenStream =
114 path.into_iter().collect();
115 (path_ts, trait_name.clone(), rest)
118 },
119 Some(id) => {
120 let msg = format!(
121 "batch-impl: 路径前缀 `#...{}` \
122 的末尾标识符与 trait 名 `{}` \
123 不一致;二者必须相同",
124 id, trait_name,
125 );
126 return compile_error_str(&msg).into();
127 },
128 None => {
129 let msg = "batch-impl: 路径前缀 `#` 后 \
130 期望至少一个标识符作为 trait 路径";
131 return compile_error_str(msg).into();
132 },
133 }
134 },
135 None => {
136 let ts = quote![#trait_name];
137 (ts, trait_name.clone(), attr_vec.clone())
138 },
139 }
140 } else {
141 let ts = quote![#trait_name];
142 (ts, trait_name.clone(), attr_vec.clone())
143 };
144
145 let mut cursor = Cursor::new(&rest_tokens);
146 let expanded =
147 match preprocess::expand_tokens(&mut cursor, &trait_item) {
148 Ok(tokens) => tokens,
149 Err(err) => return err.into(),
150 };
151 cursor = Cursor::new(&expanded);
152 let is_unsafe = trait_item.unsafety.is_some();
153 let start_trait = if include_trait {
154 Some(trait_item)
155 } else {
156 None
157 };
158 let impls = parse_batch_trait_entry(
159 &mut cursor,
160 Op::Comma,
161 &trait_full_path,
162 &trait_last_ident,
163 is_unsafe,
164 start_trait,
165 );
166 impls.into()
167}
168
169#[proc_macro]
190pub fn batch_trait(
191 input: proc_macro::TokenStream,
192) -> proc_macro::TokenStream {
193 reset_fresh_counter();
194 let tokens = TokenStream::from(input).into_iter().collect::<Vec<_>>();
195 let mut cursor = Cursor::new(&tokens);
196 let mut result = quote![];
197 loop {
198 while cursor.is_punct(';') {
200 cursor.bump();
201 }
202 if cursor.at_end() {
203 break;
204 }
205
206 let is_unsafe = if matches!(cursor.peek(), Some(TokenTree::Ident(id)) if *id == "unsafe")
208 {
209 cursor.bump();
210 true
211 } else {
212 false
213 };
214
215 let path_start = cursor.pos();
217 let mut depth = 0i32;
218 while let Some(token) = cursor.peek() {
219 match token {
220 TokenTree::Punct(p) if p.as_char() == '<' => {
221 depth += 1;
222 cursor.bump();
223 },
224 TokenTree::Punct(p) if p.as_char() == '>' => {
225 depth -= 1;
226 cursor.bump();
227 },
228 TokenTree::Punct(p)
229 if p.as_char() == ':' && depth == 0 =>
230 {
231 if matches!(cursor.peek_at(1), Some(TokenTree::Punct(p2)) if p.spacing()==Spacing::Joint && p2.as_char() == ':')
232 {
233 cursor.bump();
234 cursor.bump();
235 } else {
236 break;
237 }
238 },
239 _ => cursor.bump(),
240 }
241 }
242 let trait_path = cursor.slice_since(path_start);
243 if trait_path.is_empty() {
244 result.extend(compile_error_str(
245 "batch_trait! 中期望 trait 名称",
246 ));
247 break;
248 }
249 let trait_full_path = trait_path.iter().cloned().collect();
251 let trait_last_ident = match trait_path
253 .iter()
254 .filter_map(|tt| {
255 if let TokenTree::Ident(id) = tt {
256 Some(id)
257 } else {
258 None
259 }
260 })
261 .next_back()
262 {
263 Some(ident) => ident,
264 None => {
265 result.extend(compile_error_str(
266 "batch_trait! 中期望标识符作为 trait 名称",
267 ));
268 break;
269 },
270 };
271 if !cursor.is_punct(':') {
272 result.extend(compile_error_str(
273 "batch_trait! 中期望 ':' 分隔 trait 名称和 impl-specs",
274 ));
275 break;
276 }
277 cursor.bump();
278 let impl_code = parse_batch_trait_entry(
279 &mut cursor,
280 Op::Semi,
281 &trait_full_path,
282 trait_last_ident,
283 is_unsafe,
284 None,
285 );
286 result.extend(impl_code);
287 }
288 result.into()
289}
290
291