1#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/README.md"))]
2#![forbid(unsafe_code)]
4#![deny(missing_docs)]
5#![allow(linker_messages)]
8#[cfg(test)]
9mod fuzz;
10use proc_macro2::{TokenStream, TokenTree};
11use quote::quote;
12use syn::{ItemTrait, parse_macro_input};
13
14mod apply;
15mod ast;
16mod batch_trait_entry;
17mod codegen;
18mod diagnostic;
19mod parse;
20mod path_prefix;
21mod preprocess;
22mod scan;
23
24use batch_trait_entry::parse_batch_trait_entry;
25
26use ast::{Op, reset_fresh_counter};
27use diagnostic::compile_error_str;
28use preprocess::where_process;
29use preprocess::{build_from_item, get_trait_item, parse_names_from_tokens};
30use scan::{Cursor, scan_stop};
31
32#[proc_macro_attribute]
67pub fn batch_impl(
68 attr: proc_macro::TokenStream, item: proc_macro::TokenStream,
69) -> proc_macro::TokenStream {
70 let trait_item = parse_macro_input!(item as ItemTrait);
71 expand_attr_macro(attr, trait_item, true).unwrap_or_else(Into::into)
72}
73
74#[proc_macro_attribute]
92pub fn batch_impl_only(
93 attr: proc_macro::TokenStream, item: proc_macro::TokenStream,
94) -> proc_macro::TokenStream {
95 let trait_item = parse_macro_input!(item as ItemTrait);
96 expand_attr_macro(attr, trait_item, false).unwrap_or_else(Into::into)
97}
98
99fn expand_attr_macro(
101 attr: proc_macro::TokenStream, trait_item: ItemTrait, include_trait: bool,
102) -> Result<proc_macro::TokenStream, TokenStream> {
103 reset_fresh_counter();
104 let trait_name = trait_item.ident.clone();
105 let attr_vec = TokenStream::from(attr).into_iter().collect::<Vec<_>>();
106 let attr_vec = preprocess::angle_collect(&attr_vec)?;
109
110 let (trait_full_path, trait_last_ident, rest_tokens) = if !include_trait {
115 match path_prefix::try_parse_path_prefix(&attr_vec) {
116 Some((path, last_ident, rest)) => {
117 match last_ident {
120 Some(id) if id == trait_name => {
121 let path_ts = path.into_iter().collect();
122 (path_ts, trait_name.clone(), rest)
125 }
126 Some(id) => {
127 let msg = format!(
128 "batch-impl: 路径前缀 `#...{}` \
129 的末尾标识符与 trait 名 `{}` \
130 不一致;二者必须相同",
131 id, trait_name,
132 );
133 return Err(compile_error_str(&msg));
134 }
135 None => {
136 let msg = "batch-impl: 路径前缀 `#` 后 \
137 期望至少一个标识符作为 trait 路径";
138 return Err(compile_error_str(msg));
139 }
140 }
141 }
142 None => (quote![#trait_name], trait_name.clone(), attr_vec.clone()),
143 }
144 } else {
145 (quote![#trait_name], trait_name.clone(), attr_vec.clone())
146 };
147
148 let mut cursor = Cursor::new(&rest_tokens);
149 let expanded = preprocess::expand_tokens(&mut cursor, &trait_item)?;
150 let expanded = where_process(&mut Cursor::new(&expanded))?;
153 let is_unsafe = trait_item.unsafety.is_some();
154 let trait_bounds = extract_trait_bounds(&trait_item);
155 let expanded = expand_empty_trait_generics(&expanded, &trait_item)?;
159 cursor = Cursor::new(&expanded);
160 let start_trait = if include_trait { trait_item.into() } else { None };
161 let impls = parse_batch_trait_entry(
162 &mut cursor,
163 Op::Comma,
164 &trait_full_path,
165 &trait_last_ident,
166 is_unsafe,
167 start_trait,
168 &trait_bounds,
169 );
170 Ok(preprocess::render_angles(impls).into())
172}
173
174#[derive(Default)]
176pub(crate) struct TraitParam {
177 pub(crate) name: String,
178 pub(crate) bound: Option<TokenStream>,
179 pub(crate) refs: Vec<String>,
180}
181
182#[derive(Default)]
194pub(crate) struct TraitBounds {
195 pub(crate) params: Vec<TraitParam>,
196}
197
198fn extract_trait_bounds(trait_item: &ItemTrait) -> TraitBounds {
199 let type_const_names: Vec<String> = trait_item
201 .generics
202 .params
203 .iter()
204 .filter_map(|p| match p {
205 syn::GenericParam::Type(tp) => Some(tp.ident.to_string()),
206 syn::GenericParam::Const(cp) => Some(cp.ident.to_string()),
207 _ => None,
208 })
209 .collect();
210 let lt_names: Vec<String> = trait_item
211 .generics
212 .params
213 .iter()
214 .filter_map(|p| match p {
215 syn::GenericParam::Lifetime(ld) => {
216 Some(format!("'{}", ld.lifetime.ident))
217 }
218 _ => None,
219 })
220 .collect();
221 let mut params = vec![];
222 for p in &trait_item.generics.params {
223 match p {
224 syn::GenericParam::Type(tp) => {
225 let bound = if tp.bounds.is_empty() {
226 None
227 } else {
228 let b = &tp.bounds;
231 Some(quote!(#b))
232 };
233 let refs = bound
234 .as_ref()
235 .map(|b| bound_refs(b, &type_const_names, <_names))
236 .unwrap_or_default();
237 params.push(TraitParam { name: tp.ident.to_string(), bound, refs });
238 }
239 syn::GenericParam::Lifetime(ld) => params.push(TraitParam {
240 name: format!("'{}", ld.lifetime.ident),
241 bound: None,
242 refs: vec![],
243 }),
244 syn::GenericParam::Const(cp) => params.push(TraitParam {
245 name: cp.ident.to_string(),
246 bound: None,
247 refs: vec![],
248 }),
249 }
250 }
251 TraitBounds { params }
252}
253
254fn bound_refs(
258 bound: &TokenStream, type_const_names: &[String], lt_names: &[String],
259) -> Vec<String> {
260 let mut refs = vec![];
261 let mut iter = bound.clone().into_iter().peekable();
262 while let Some(tt) = iter.next() {
263 match tt {
264 TokenTree::Ident(id) if type_const_names.contains(&id.to_string()) => {
265 refs.push(id.to_string())
266 }
267 TokenTree::Punct(p) if p.as_char() == '\'' => {
268 if let Some(TokenTree::Ident(id)) = iter.peek() {
269 let name = format!("'{}", id);
270 if lt_names.contains(&name) {
271 refs.push(name);
272 }
273 }
274 }
275 _ => {}
276 }
277 }
278 refs
279}
280
281fn args_all_bindings(args: &[TokenTree]) -> bool {
285 let mut rest = args;
286 while let Some(idx) = scan_stop(rest, &[',']) {
287 if scan_stop(&rest[..idx], &['=']).is_none() {
289 return false;
290 }
291 rest = &rest[idx + 1..];
292 }
293 scan_stop(rest, &['=']).is_some()
294}
295
296fn expand_empty_trait_generics(
307 tokens: &[TokenTree], trait_def: &ItemTrait,
308) -> Result<Vec<TokenTree>, TokenStream> {
309 if trait_def.generics.params.is_empty() {
310 return Ok(tokens.to_vec());
311 }
312 let mut arg_names: Vec<TokenStream> = vec![];
314 for p in &trait_def.generics.params {
315 match p {
316 syn::GenericParam::Lifetime(ld) => arg_names.push(quote!(#ld)),
317 syn::GenericParam::Type(tp) => {
318 let id = &tp.ident;
319 arg_names.push(quote!(#id));
320 }
321 syn::GenericParam::Const(cp) => {
322 let id = &cp.ident;
323 arg_names.push(quote!(#id));
324 }
325 }
326 }
327 let mut out = vec![];
328 let mut i = 0;
329 while i < tokens.len() {
330 match &tokens[i] {
331 TokenTree::Ident(id) => {
336 let group = match tokens.get(i + 1) {
337 Some(TokenTree::Group(g))
338 if g.delimiter() == proc_macro2::Delimiter::None =>
339 {
340 g
341 }
342 _ => {
343 out.push(tokens[i].clone());
344 i += 1;
345 continue;
346 }
347 };
348 let args: Vec<TokenTree> = group.stream().into_iter().collect();
349 let bindings_only = !args.is_empty() && args_all_bindings(&args);
350 if args.is_empty() || bindings_only {
351 let all_params: Vec<_> =
354 trait_def.generics.params.iter().collect();
355 out.push(
356 proc_macro2::Group::new(
357 proc_macro2::Delimiter::None,
358 quote!(#(#all_params),*),
359 )
360 .into(),
361 );
362 out.extend(quote!(#id));
363 let args_ts: TokenStream = if args.is_empty() {
364 quote!(#(#arg_names),*)
365 } else {
366 let bind_ts: TokenStream = args.iter().cloned().collect();
367 quote!(#(#arg_names),* , #bind_ts)
368 };
369 out.push(
370 proc_macro2::Group::new(
371 proc_macro2::Delimiter::None,
372 args_ts,
373 )
374 .into(),
375 );
376 i += 2;
377 } else {
378 out.push(tokens[i].clone());
379 i += 1;
380 }
381 }
382 _ => {
383 out.push(tokens[i].clone());
384 i += 1;
385 }
386 }
387 }
388 Ok(out)
389}
390
391#[proc_macro]
413pub fn batch_trait(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
414 expand_batch_trait(input).unwrap_or_else(Into::into)
415}
416
417fn expand_batch_trait(
419 input: proc_macro::TokenStream,
420) -> Result<proc_macro::TokenStream, TokenStream> {
421 reset_fresh_counter();
422 let tokens = TokenStream::from(input).into_iter().collect::<Vec<_>>();
423 let tokens = preprocess::angle_collect(&tokens)?;
424 let tokens = where_process(&mut Cursor::new(&tokens))?;
425 let mut cursor = Cursor::new(&tokens);
426 let mut result = quote![];
427 loop {
428 while cursor.is_punct(';') {
430 cursor.bump();
431 }
432 if cursor.at_end() {
433 break;
434 }
435
436 let is_unsafe = if matches!(cursor.peek(), Some(TokenTree::Ident(id)) if *id == "unsafe")
438 {
439 cursor.bump();
440 true
441 } else {
442 false
443 };
444
445 let path_start = cursor.pos();
448 while let Some(token) = cursor.peek() {
449 match token {
450 TokenTree::Punct(p) if p.as_char() == ':' => {
451 if cursor.is_single_colon() {
452 break;
453 } else {
454 cursor.bump();
455 cursor.bump();
456 }
457 }
458 _ => cursor.bump(),
459 }
460 }
461 let trait_path = cursor.slice_since(path_start);
462 if trait_path.is_empty() {
463 result.extend(compile_error_str("batch_trait! 中期望 trait 名称"));
464 break;
465 }
466 let trait_full_path = trait_path.iter().cloned().collect();
468 let trait_last_ident =
470 match trait_path
471 .iter()
472 .filter_map(|tt| {
473 if let TokenTree::Ident(id) = tt { id.into() } else { None }
474 })
475 .next_back()
476 {
477 Some(ident) => ident,
478 None => {
479 result.extend(compile_error_str(
480 "batch_trait! 中期望标识符作为 trait 名称",
481 ));
482 break;
483 }
484 };
485 if !cursor.is_punct(':') {
486 result.extend(compile_error_str(
487 "batch_trait! 中期望 ':' 分隔 trait 名称和 impl-specs",
488 ));
489 break;
490 }
491 cursor.bump();
492 let impl_code = parse_batch_trait_entry(
493 &mut cursor,
494 Op::Semi,
495 &trait_full_path,
496 trait_last_ident,
497 is_unsafe,
498 None,
499 &Default::default(),
501 );
502 result.extend(impl_code);
503 }
504 Ok(preprocess::render_angles(result).into())
505}
506
507#[doc(hidden)]
520#[proc_macro]
521pub fn batch_preprocess_test(
522 input: proc_macro::TokenStream,
523) -> proc_macro::TokenStream {
524 let tokens = TokenStream::from(input).into_iter().collect::<Vec<_>>();
525 let tokens = match preprocess::angle_collect(&tokens) {
526 Ok(v) => v,
527 Err(e) => return e.into(),
528 };
529 let Some(TokenTree::Group(names_group)) = tokens.first() else {
531 return compile_error_str(
532 "batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
533 )
534 .into();
535 };
536 if names_group.delimiter() != proc_macro2::Delimiter::Parenthesis {
537 return compile_error_str(
538 "batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
539 )
540 .into();
541 }
542 let Some(TokenTree::Group(body_group)) = tokens.get(1) else {
543 return compile_error_str(
544 "batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
545 )
546 .into();
547 };
548 if body_group.delimiter() != proc_macro2::Delimiter::Brace {
549 return compile_error_str(
550 "batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
551 )
552 .into();
553 }
554 let trait_ts = tokens[2..].iter().cloned().collect();
555 let trait_item = match syn::parse2(trait_ts) {
556 Ok(t) => t,
557 Err(_) => {
558 return compile_error_str(
559 "batch-impl: batch_preprocess_test 无法解析 trait 定义",
560 )
561 .into();
562 }
563 };
564 let names = match parse_names_from_tokens(
565 &names_group.stream().into_iter().collect::<Vec<_>>(),
566 &trait_item,
567 ) {
568 Ok(names) => names,
569 Err(e) => return e.into(),
570 };
571 let body = body_group.stream();
572 let mut methods = TokenStream::new();
573 for name in &names {
574 let item = match get_trait_item(&trait_item, name) {
575 Ok(item) => item,
576 Err(e) => return e.into(),
577 };
578 methods.extend(build_from_item(item, &body));
579 }
580 preprocess::render_angles(methods).into()
581}
582
583#[cfg(test)]
584mod angle_tests {
585 use super::*;
586 use proc_macro2::TokenStream as TS2;
587 use std::str::FromStr;
588
589 fn roundtrip(s: &str) -> String {
591 let ts: TS2 = FromStr::from_str(s).unwrap();
592 let v: Vec<_> = ts.into_iter().collect();
593 let collected = preprocess::angle_collect(&v).unwrap();
594 preprocess::render_angles(collected.into_iter().collect()).to_string()
595 }
596
597 #[test]
598 fn angle_roundtrip() {
599 assert_eq!(roundtrip("Vec<T>"), "Vec < T >");
600 assert_eq!(roundtrip("A<B<C>>"), "A < B < C > >");
601 assert_eq!(
602 roundtrip("Box<dyn Fn() + Send>"),
603 "Box < dyn Fn () + Send >"
604 );
605 assert_eq!(roundtrip("<T: Clone> A<T>"), "< T : Clone > A < T >");
606 assert_eq!(roundtrip("A<Item=T>"), "A < Item = T >");
607 assert_eq!(roundtrip("fn(A) -> B"), "fn (A) -> B");
609 }
610
611 #[test]
612 fn angle_unmatched_errors() {
613 let ts: TS2 = FromStr::from_str("A <").unwrap();
615 assert!(
616 preprocess::angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_err()
617 );
618 let ts: TS2 = FromStr::from_str("A >").unwrap();
619 assert!(
620 preprocess::angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_err()
621 );
622 let ts: TS2 = FromStr::from_str("m![a < b]").unwrap();
624 assert!(
625 preprocess::angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_ok()
626 );
627 }
628
629 #[test]
630 fn none_group_flattened() {
631 let inner: TS2 = FromStr::from_str("Vec<T>").unwrap();
633 let none = proc_macro2::Group::new(proc_macro2::Delimiter::None, inner);
634 let collected = preprocess::angle_collect(&[none.into()]).unwrap();
635 let rendered = preprocess::render_angles(collected.into_iter().collect());
636 assert_eq!(rendered.to_string(), "Vec < T >");
637 }
638}