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}