Skip to main content

batch_impl/
lib.rs

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