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