Skip to main content

batch_impl/
lib.rs

1use proc_macro::TokenStream;
2use proc_macro2::TokenStream as TokenStream2;
3use quote::quote;
4use syn::{parse_macro_input, ItemTrait};
5
6mod core;
7
8use core::types::{ParseResult, err};
9use core::codegen::generate_impl;
10use core::parser::parse_top_level;
11
12// ===========================================================================
13// #[batch_impl(...)]
14// ===========================================================================
15
16/// 为 trait 批量生成 impl 块的属性宏。
17///
18/// # 语法概览
19///
20/// ```text
21/// #[batch_impl( impl-spec [, impl-spec]* [ { body }]? )]
22/// impl-spec = [ <impl-泛型> ] [ Trait名<trait-泛型> ] 目标 [ { body } ]
23/// ```
24///
25/// # `^` 运算符(右结合)
26///
27/// `A^B^C = A^(B^C)`
28///
29/// | 写法 | 展开 |
30/// |------|------|
31/// | `&^T` | `&T` |
32/// | `&mut^T` | `&mut T` |
33/// | `self^T` | `T` |
34/// | `A^B` | `A<B>` |
35/// | `A^<X,Y>` | `A<X,Y>` |
36/// | `[A1,A2]^B` | `A1<B>, A2<B>` |
37/// | `A^[B1,B2]` | `A<B1>, A<B2>` |
38/// | `[A1,A2]^[B1,B2]` | 笛卡尔积 `A1<B1>, A1<B2>, A2<B1>, A2<B2>` |
39/// | `Box^Box^T` | `Box<Box<T>>` |
40/// | `Box^[Box^T]` | `Box<[Box<T>]>` |
41///
42/// # 元组 `^`(追加/生成)
43///
44/// 追加(右侧是类型):
45/// | 写法 | 展开 |
46/// |------|------|
47/// | `()^T` | `(T,)` |
48/// | `(A,B)^T` | `(A, B, T)` |
49///
50/// 生成(右侧是数字/范围):
51/// | 写法 | 展开 |
52/// |------|------|
53/// | `()^N` | `(), (X,), (X,X), ...` |
54/// | `(T)^N` | `(), (T,), (T,T), ...` |
55/// | `(<tr>)^N` | `(), (A:tr,), (A:tr,B:tr), ...` |
56/// | `^M..N` | 长度 M 到 N-1 |
57/// | `^M..=N` | 长度 M 到 N |
58///
59/// 笛卡尔积生成(前缀含逗号):
60/// | 写法 | 展开 |
61/// |------|------|
62/// | `(T1,T2)^N` | 长度 0..N-1 的所有 T1/T2 组合 |
63/// | `(<tr>,T)^N` | 带 bound 的泛型+固定类型组合 |
64///
65/// # `-` 运算符(左结合)
66///
67/// `A-B = A^B`,`A-B-C = (A-B)-C`
68///
69/// | 写法 | 展开 |
70/// |------|------|
71/// | `()-[A,B]` | `(A,), (B,)` |
72/// | `()-[A,B]-[C,D]` | `(A,C), (A,D), (B,C), (B,D)` |
73/// | `()-[A]-[B]-[C]` | `(A,B,C)` |
74///
75/// # 优先级
76///
77/// `^` 高于 `-`,`,` 最低
78///
79/// # 其他规则
80///
81/// - **`[]` 歧义**:有逗号是并列列表,无逗号是切片类型
82/// - **`()` 歧义**:`()=空元组`, `(A,)=单元素元组`, `(A)=分组(非元组)`
83/// - **trait 泛型必须显式**:trait 有泛型时必须写 `Trait名<T>`
84/// - **泛型继承**:嵌套 `[...]` 中子项不写泛型则继承父级,写了则追加并去重
85/// - **body 继承**:列表级 `{...}` 被所有子项共享;子项 `{...}` 覆盖列表级
86/// - **目标类型透传**:不解析,原样透传
87///
88/// # 基本示例
89///
90/// ```
91/// # use batch_impl::batch_impl;
92/// #[batch_impl(usize, isize)]
93/// trait Numeric {}
94/// ```
95///
96/// # 类型标注 / const 泛型 / 生命周期
97///
98/// ```
99/// # use batch_impl::batch_impl;
100/// #[batch_impl(<T: Clone + std::fmt::Debug> Vec<T>)]
101/// trait DebugClone {}
102///
103/// #[batch_impl(<const N: usize> [i32; N])]
104/// trait FixedSize {}
105///
106/// #[batch_impl(<'a, T: 'a> &'a T)]
107/// trait RefTrait {}
108/// ```
109///
110/// # trait 自身带泛型参数
111///
112/// ```
113/// # use batch_impl::batch_impl;
114/// #[batch_impl(<T> FromValue<T> i32 {
115///     fn wrap(_val: T) -> Self { 0 }
116/// })]
117/// trait FromValue<T> { fn wrap(val: T) -> Self; }
118///
119/// #[batch_impl(<T> Wrapper<Vec<T>> Vec<T> {
120///     fn inner(self) -> Vec<T> { self }
121/// })]
122/// trait Wrapper<C> { fn inner(self) -> C; }
123/// ```
124///
125/// > **必须显式写 `Trait名<泛型>`**,否则生成的 impl 缺少 trait 泛型。
126///
127/// # 并列列表
128///
129/// ```
130/// # use batch_impl::batch_impl;
131/// #[batch_impl([usize, isize, f32] {
132///     fn tag(&self) -> &'static str { "number" }
133/// })]
134/// trait Tagged { fn tag(&self) -> &'static str; }
135///
136/// // 嵌套泛型合并(<U> 追加到父级 <T>,去重)
137/// # use std::collections::HashMap;
138/// #[batch_impl(<T> Describe<T> [Vec<T>, <U> HashMap<T, U>] {
139///     fn describe(&self) -> String { format!("len={}", self.len()) }
140/// })]
141/// trait Describe<T> { fn describe(&self) -> String; }
142/// // → impl<T>    Describe<T> for Vec<T>
143/// // → impl<T, U> Describe<T> for HashMap<T, U>
144/// ```
145///
146/// # `^` 运算符示例
147///
148/// ```
149/// # use batch_impl::batch_impl;
150/// #[batch_impl([&, self]^u32)]
151/// trait RefOrOwned {}
152///
153/// # use std::collections::HashMap;
154/// #[batch_impl(HashMap^<u32, i32>)]
155/// trait MapMarker {}
156/// ```
157///
158/// # 多项独立 body
159///
160/// ```
161/// # use batch_impl::batch_impl;
162/// #[batch_impl(
163///     usize  { fn id(&self) -> usize { *self }      },
164///     String { fn id(&self) -> usize { self.len() }  }
165/// )]
166/// trait Identifiable { fn id(&self) -> usize; }
167/// ```
168///
169/// # 复杂类型透传
170///
171/// ```
172/// # use batch_impl::batch_impl;
173/// #[batch_impl((i32, String), &str, Box<dyn std::fmt::Display>, dyn Fn() + Send + Sync)]
174/// trait ComplexMarker {}
175/// ```
176///
177/// # 设计约束
178///
179/// - **where 子句**:不在 DSL 内,复杂 bound 写在 trait 定义自身
180/// - **`for<'a>` 高阶 trait bound**:where 子句式范畴;类型内部走 token 透传
181/// - **`<>` 多参不能在 `[]` 内**:Rust `[<` 解析歧义,拆成独立表达式
182#[proc_macro_attribute]
183pub fn batch_impl(attr: TokenStream, item: TokenStream) -> TokenStream {
184    let trait_item = parse_macro_input!(item as ItemTrait);
185    let trait_name = &trait_item.ident;
186    let name_ts = quote! { #trait_name };
187
188    // 检测是否是 unsafe trait
189    let is_unsafe_trait = trait_item.unsafety.is_some();
190
191    let attr_ts: TokenStream2 = attr.into();
192    let specs = parse_top_level(attr_ts, trait_name, trait_name.span());
193
194    let mut output = TokenStream2::new();
195    match specs {
196        ParseResult::Ok(specs) => {
197            for mut spec in specs {
198                // 如果是 unsafe trait,所有 impl 都标记为 unsafe
199                if is_unsafe_trait {
200                    spec.is_unsafe = true;
201                }
202                output.extend(generate_impl(&spec, name_ts.clone()));
203            }
204        }
205        ParseResult::Err(e) => output.extend(e),
206    }
207
208    (quote! { #trait_item #output }).into()
209}
210
211// ===========================================================================
212// batch_trait!(Trait: impl-specs; ...)
213// ===========================================================================
214
215/// 对已声明的 trait 批量生成 impl 块的函数式宏。
216///
217/// # 语法
218///
219/// ```text
220/// batch_trait!(Trait路径: impl-specs; Trait路径: impl-specs)
221/// ```
222///
223/// `:` 前是 trait 路径(支持 `sub::MyTrait`),`:` 后是 `#[batch_impl]` 同款参数,
224/// `;` 分隔不同的 trait。
225///
226/// ```
227/// # use batch_impl::batch_trait;
228/// trait A {}
229/// trait B<T> {}
230/// mod foo { pub trait C {} }
231///
232/// batch_trait!(
233///     A: usize, isize;
234///     B: <T> B<T> Vec<T>;
235///     foo::C: u32
236/// );
237/// // → impl A for usize {}  +  impl A for isize {}
238/// // → impl<T> B<T> for Vec<T> {}
239/// // → impl foo::C for u32 {}
240/// ```
241///
242/// **注意**:trait 有泛型时,必须在 `:` 后的 impl-spec 中显式写出 `Trait名<泛型>`。
243#[proc_macro]
244pub fn batch_trait(input: TokenStream) -> TokenStream {
245    let input_ts: TokenStream2 = input.into();
246    let span = core::utils::tokens_span(&input_ts);
247
248    let segments = match core::utils::split_by(input_ts, ';', span) {
249        Ok(s) => s,
250        Err(e) => return e.into(),
251    };
252
253    let mut output = TokenStream2::new();
254    for seg in segments {
255        let seg_vec: Vec<proc_macro2::TokenTree> = seg.into_iter().collect();
256        if seg_vec.is_empty() {
257            continue;
258        }
259
260        let colon = match core::utils::find_top_level_colon(&seg_vec) {
261            Some(p) => p,
262            None => {
263                output.extend(err(
264                    span,
265                    "batch_trait! 缺少 `:`(格式:Trait名: impl规格,如 A: usize, isize)",
266                ));
267                continue;
268            }
269        };
270
271        // 检测 unsafe 关键字
272        let (is_unsafe, trait_start) = if colon > 0
273            && matches!(&seg_vec[0], proc_macro2::TokenTree::Ident(id) if id == "unsafe")
274        {
275            (true, 1)
276        } else {
277            (false, 0)
278        };
279
280        let trait_path: TokenStream2 = seg_vec[trait_start..colon].iter().cloned().collect();
281        let specs_ts: TokenStream2 = seg_vec[colon + 1..].iter().cloned().collect();
282
283        let trait_ident = match syn::parse2::<syn::Path>(trait_path.clone()) {
284            Ok(p) => p.segments.last().unwrap().ident.clone(),
285            Err(e) => {
286                output.extend(err(
287                    span,
288                    &format!(
289                        "batch_trait! 中 `:` 左侧不是合法的路径: {}。示例: A, foo::Bar",
290                        e
291                    ),
292                ));
293                continue;
294            }
295        };
296
297        match parse_top_level(specs_ts, &trait_ident, trait_ident.span()) {
298            ParseResult::Ok(specs) => {
299                for mut spec in specs {
300                    if is_unsafe {
301                        spec.is_unsafe = true;
302                    }
303                    output.extend(generate_impl(&spec, trait_path.clone()));
304                }
305            }
306            ParseResult::Err(e) => output.extend(e),
307        }
308    }
309    output.into()
310}