本文又名《重生之我要成为rust魔法师》《五题三库粉碎魔法梦》
学习一下rust宏,参考资料:
- David Tolnay 此领域绕不开的名字
- TRPL
- The Little Book of Rust Macros
- Rust for Rustaceans 第七章 Macro
basic
宏是什么?
- 是用来生成代码的代码,所以宏编程又叫元编程(meta programming)
- 是一个工具,能让编译器替我写代码
宏有哪些种类?
- 声明宏(declarative macro) 看起来类似于一个
match表达式,info!println!就属于此类 - 过程宏(procedural macro) 看起来像一个标签,比如常用的
#[derive(debug)]#[tokio::main]
宏与函数有什么区别?
- 宏支持可变数量的参数
- 宏在编译器解释代码前展开,所以可以做到函数做不到的事,比如为一个类型实现一个特征
- 宏更复杂
宏只能用在允许的地方,包括表达式、impl块等等大部分地方,但是不包括标识符和match arm,所以以下场景是不可以的:
fn main(){
arg! = 1; // expected identifier
let choice = 1;
match choice {
choice!(0), // expected match arm
_ => ,
}
}
声明宏
声明宏特指那些用macro_rules!语法定义的宏,可以方便的定义函数类宏
值得说的一点是,之所以它叫这个一点也不直观的名字,来自于它的写法。相比于告诉宏的输入如何被翻译成输出,声明宏的模式是对于像A的输入,输出应该是B,程序员声明了若干这样的规则,编译器只需要解析并重写即可
macro_rules!的语法大致是:
macro_rules! $name {
$rule0 ;
$rule1 ;
// …
$ruleN ;
}
对于每个$rule,其形如 ($pattern) => {$expansion}
一个最简单的例子像这样:
macro_rules! four {
() => {1 + 3};
}
// 以下三种都能匹配成功
four!()
four![]
four!{}
但是这个例子无法处理输入,看起来只是长得很奇怪的函数
captures
在pattern的位置可以声明这个分支想要的参数是什么样的,语法形如:
macro_rules! one_expression {
($e:expr) => {...};
}
e是在expansion里用的标识符(变量名),expr是它的类型,具体而言,这里可以有:
- expr 表达式,比如
1+1,hello_world - ident 标识符
x - ty 类型
i32 - stmt 语句
x+1 - item 一个项,比如一个函数、结构体、模块(moduel)等等
- block 代码块,
{……} - path 带双冒号的作用域路径,比如
foo,::std::mem::replace,transmute::<_, int> - tt 一个token tree,简单得说,万能通配符
- pat 模式,可以放在match后面的东西,比如
Some(x),(a,b) - meta 元属性,包含在#[…]或#![…]里的东西,比如
derive(Debug),inline - …
其中以上内容可以组合,用起来像这样:
macro_rules! multiply_add {
($a:expr, $b:expr, $c:expr) => {$a * ($b + $c)};
}
Repetitions
刚才的例子里虽然支持了参数,但是数量是严格写死的,还是只能算一种奇怪的函数
为了支持可变数量参数,需要Repetitions语法,它形如$ ( ... ) sep rep
其中省略号里的是刚才的captures,seq是任意分隔符,比如, ;,req用于控制重复多少次,具体而言:
*0或任意次+1或任意次?0或1次
一个例子如下:
macro_rules! vec_strs {
(
// Start a repetition:
$(
// Each repeat must contain an expression...
$element:expr
)
// ...separated by commas...
,
// ...zero or more times.
*
) => {
// Enclose the expansion in a block so that we can use
// multiple statements.
{
let mut v = Vec::new();
// Start a repetition:
$(
// Each repeat will contain the following statement, with
// $element replaced with the corresponding expression.
v.push(format!("{}", $element));
)*
v
}
};
}
使用
简单的说,当发现自己在写重复的、结构高度相似的代码时,就可以用宏。
比如测试:
macro_rules! test_battery {
( $( $t:ty as $name:ident ),* ) => {
$(
mod $name {
use super::*;
#[test]
fn frobnified() {
test_inner::<$t>(1, true)
}
#[test]
fn unfrobnified() {
test_inner::<$t>(1, false)
}
}
)*
}
}
fn test_inner<T>(val: i32, flag: bool) {
// ... 测试逻辑
}
test_battery! {
u8 as u8_tests,
i128 as i128_tests
}
它等价于:
fn test_inner<T>(val: i32, flag: bool) {
// ... 测试逻辑
}
mod u8_tests {
use super::*;
#[test] fn frobnified() { test_inner::<u8>(1, true) }
#[test] fn unfrobnified() { test_inner::<u8>(1, false) }
}
mod u16_tests {
use super::*;
#[test] fn frobnified() { test_inner::<u16>(1, true) }
#[test] fn unfrobnified() { test_inner::<u16>(1, false) }
}
mod i128_tests {
use super::*;
#[test] fn frobnified() { test_inner::<i128>(1, true) }
#[test] fn unfrobnified() { test_inner::<i128>(1, false) }
}
自定义一个特征时,对于已有的类型显然都实现以下是再好不过的了:
macro_rules! clone_from_copy {
($($t:ty) , *) => {
$(impl Clone for $t {
fn clone(&self) -> Self {* self}
}) *
}
}
fn main() {
clone_from_copy![bool, f32, f64, u8, i8 /*...*/];
}
过程宏
过程宏又分为以下几类:
- 自定义派生宏(custom derive macros),比如
#[derive(debug)] - 属性宏(attriute macros),比如
#[test] - 类函数宏(function-like macros),比如
sql!,看起来和声明宏没什么两样
类函数宏
这是最简单的过程宏形式,就像声明宏一样,它只是做简单的代码替换
但是macro_rules!会严格限制作用域,而过程宏没有,可能会污染上下文,这就引出了宏的卫生性问题,需要在代码中显示声明内部变量名是宏内私有(Span::mixed_site)还是对外生效(Span::call_site)
大致在两种情况下,用类函数宏是合理的:
- 声明宏变得越来越臃肿、难以维护 时
- 需要一个编译器需要执行,但是
const fn又做不到的函数 时。比如phf库,在编译时把提供的一堆字符串算成一个完美的哈希表。
属性宏
属性宏也替换其作用域的item,它的输入除了宏的部分(属性名及其参数),还包括附加到的整个item
他可以很容易地把一个函数变成另一个模样,就像#[tokio::main] #[test]做的那样
属性宏是权力最大的宏,它的使用场景也比较多:
- 生成测试
- 框架胶水,典型的比如#[tokio::main],其实是重写了整个main函数
- 透明中间件
- 类型转换:改定义,比如增加字段。
派生宏
派生宏的目标和前两种不太一样,它不替换只附加。
它的限制最多:只能追加,不能带参数,用辅助属性来传递额外信息
它应当且只应当用在一个地方——在可能的情况下,自动实现一个特征,同时需要满足两个条件—— 1. 使用频率极高,否则对不起写宏花的时间 2. 逻辑必须符合直觉
典型的正面案例就是#[derive(Serialize,Deserialize)] #[derive(Debug)] #[derive(Clone)]
开销
过程宏会增长编译时间,具体体现在两个方面:
- 引入一些很重的依赖,写过程宏需要的
syncrate在所有feature都开启的情况下需要数十秒的时间编译,所以应当关闭不需要的feature,同时在debug模式下编译 - 容易在不知不觉间生成大量的代码
mechanism
以上的内容可以对宏有个大致的认识,但是还不够,对于宏如何工作,我们还是一无所知,这也导致编写过程宏变得困难
声明宏
编译器处理源代码的第一步,是把代码字符处理成一个个token。这一步只有数字、标点、字符串、标识符的区分,不在乎变量名和关键字的差别,空格和注释也会在这一步被舍弃掉。
当源代码被处理成token流之后,编译器开始为他们赋予含义,比如被()包裹的区域组成了一个group,!标识了一个宏等等,这一步称为parsing,最终会生成一个abstract syntax tree(AST)来描述代码的结构。比如:
let x = || 4
它包括 let(keyword) x(identifier) =(punctuation) 两个|(punctuation)和 4(literal)。编译器在这一步知道了 let 是关键字,它正在声明一个模式为 x 的变量,而等号右边是一个完整的“无参闭包”。这棵 AST 树比干巴巴的 Token 包含了丰富得多的语法意义。
宏决定了对于一个特定的token序列要转换成什么。编译器在解析时遇到宏调用,它会初步的评估传给宏的代码,这一步可以有语法错误,但必须符合Token Tree的规则,比如括号必须匹配。但是在parsing过程中,编译器只是解析token,而并不完成宏替换,而是记住输入序列并延迟解析。在后续的解析过程中,编译器会得到一个语法树,并替换回调用时的语法树中。
声明宏的输出永远是合法的rust代码。其返回的内容必须是一个表达式(1+1)、一个语句(let x = 1)、一个item、一个类型或一个match 模式。因而声明宏是卫生的,无效的代码会死在编译期。
过程宏
过程宏的核心是TokenStream类型,它由一个个TokenTree组成。一个TokenTree可以是一个单个的token,比如一个标识符、标点或字面量,也可以是一个被各种括号包裹的另一个TokenStream。David Tolnay的syn库提供了把原始的token流变成结构化的 Rust 抽象语法树(AST)的能力。
过程宏的目标不仅仅是解析TokenStream,还要能生成代码,这有两种主要方法:手动构建TokenStream,逐个TokenTree地拓展,或者使用TokenStream的FromStr方法,或者混合这两种……或者用毫无疑问的quote!工具!
每一个TokenTree都有一个span。代码报错时可以追溯到源代码的具体位置正是得益于它。每个token的span标记了这个token的源头,比如这样一个为给定类型实现Debug特征的声明宏:
macro_rules! name_as_debug {
($t:ty) => {
impl ::core::fmt::Debug for $t {
fn fmt(&self, f: &mut ::core::fmt::Formatter<'_>) -> ::core::fmt::Result
{ ::core::write!(f, ::core::stringify!($t)) }
} }; }
假设使用了name_as_debug(u31),报错显然是发生在宏内部的,但是用户期望报错能定位到实际调用处。此处的$t的span是实际上映射到它的代码的,这个信息一直可以关联到最终的编译错误,于是报错信息里因此包含了用户代码的错误信息。
span可以结合compile_error!宏,让自定义错误信息也能追溯到用户代码。
rust宏的卫生性也得益于span。当我们构建一个Ident(标识符) token 时,也同时为它提供了一个span,其决定了这个标识符的作用域。如果用Span::call_site(),意思是把这个变量当做在调用处声明的,等于是直接暴露在外部作用域,完全没有卫生性。如果用Span::mixed_site(),意思是把变量当做在宏内部声明的,局部变量就可以对外隐藏,是卫生的,但是对于types moduels等等其它东西依然是外部可访问的,因而是mixed
procedural macros workshop
终于来到了写过程宏这一步。
rust宏领域里那个绕不开的名字,syn和quote和proc-macro2三个重量级库的作者,David Tolnay有一个练习题项目procedural macros workshop,用来讲解如何写过程宏再合适不过。
一点理论知识
首先,由于过程宏是一段rust代码,需要编译才能运行,但是宏展开时rust代码还没有编译,所以过程宏必须存在于一个独立的crate中预先编译好: cargo new my_macro --lib
对于写过程宏必要的三个库,首先来逐个讲解一下:
syn库负责解析原始的Token流,它完成的最重要的工作就是把一切token都映射到了一个具体的、结构清晰的类型上,比如对于item,他是这样一个枚举类型:
pub enum Item {
Const(ItemConst),
Enum(ItemEnum),
ExternCrate(ItemExternCrate),
Fn(ItemFn),
//...
}
quote库允许用写稍有不同的rust代码的方法生成token流
proc_macro2拓展了系统库,使之可以用在非过程宏上下文中,同时提供了Span用于定位上下文
但是显然掌握这些知识连入门都不够格,但是不管了,直接做题吧
builder
第一关,也是作者推荐的入口点,写一个builder派生宏,从一个框架到解析结构体的内部字段,到支持Result和Option,到支持追加,到错误处理,最后把基础类型换成绝对路径以避免用户的重定义覆盖。最终lib.rs代码:
use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::{parse_macro_input, Data, DeriveInput, Fields};
// 引入辅助属性
#[proc_macro_derive(Builder, attributes(builder))]
pub fn derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as DeriveInput);
let ident = &ast.ident; // 传入的结构体的名字
// 提取结构体字段
let fields = match &ast.data {
Data::Struct(data_struct) => match &data_struct.fields {
Fields::Named(fields_named) => &fields_named.named,
_ => panic!("panic!"),
},
_ => panic!("panic!"),
};
// 属性检查
for f in fields {
if let std::result::Result::Err(e) = get_each_attr(f) {
// 把 syn::Error 转换成 compile_error!("...") 标记流返回
return proc_macro::TokenStream::from(e.to_compile_error());
}
}
// builder结构体名
let builder_ident = format_ident!("{}Builder", ident);
// 变量名,变量初始化和变量类型,用vec使得可以反复迭代
let f_names: Vec<_> = fields.iter().map(|f| &f.ident).collect();
// builder结构体的字段,如果原类型不是Option,就包一层Option
let builder_fields = fields.iter().map(|f| {
let name = &f.ident;
let ty = &f.ty;
if get_inner_type(ty, "Option").is_some() {
// 用原类型
quote! { #name: #ty }
} else {
// 包一层 Option
quote! { #name: std::option::Option<#ty> }
}
});
// 处理setter,如果输入是Option 需要提取出来,如果类型是一个Vec,考虑追加模式
let setters = fields.iter().map(|f| {
let name = f.ident.as_ref().unwrap();
let ty = &f.ty;
if let std::option::Option::Some(inner_ty) = get_inner_type(ty, "Option") {
quote! {
pub fn #name(&mut self, value: #inner_ty) -> &mut Self {
self.#name = std::option::Option::Some(value);
self
}
}
} else if let std::option::Option::Some(each_name) = get_each_attr(f).unwrap() {
// 如果提供的是一个参数
let inner_ty = get_inner_type(ty, "Vec").unwrap();
let mut method = quote! {
pub fn #each_name(&mut self, value : #inner_ty) -> &mut Self{
let vec = self.#name.get_or_insert_with(std::vec::Vec::new);
vec.push(value);
self
}
};
// 提供覆盖
if each_name != *name {
method.extend(quote! {
pub fn #name(&mut self, value : #ty) -> &mut Self {
self.#name = std::option::Option::Some(value);
self
}
});
}
method
} else {
quote! {
pub fn #name(&mut self, value: #ty) -> &mut Self {
self.#name = std::option::Option::Some(value);
self
}
}
}
});
// 构建函数
let build_fields = fields.iter().map(|f| {
let name = &f.ident;
let ty = &f.ty;
if get_inner_type(ty, "Option").is_some() {
// 可为空 直接 clone
quote! { #name: self.#name.clone() }
} else if get_each_attr(f).unwrap().is_some() {
quote! {#name: self.#name.clone().unwrap_or_else(std::vec::Vec::new)}
} else {
// 必填 如果是 None 就报错
quote! {
#name: self.#name.clone().ok_or(format!("{} is missing", stringify!(#name)))?
}
}
});
// 组合,分为:builder结构体的声明与初始化,每个字段的setter函数
let expanded = quote! {
pub struct #builder_ident {
#(
#builder_fields,
)*
}
impl #ident {
pub fn builder() -> #builder_ident {
#builder_ident {
#(
#f_names: std::option::Option::None,
)*
}
}
}
impl #builder_ident {
#( #setters )*
pub fn build(&mut self) -> std::result::Result<#ident, std::boxed::Box<dyn std::error::Error>>{
std::result::Result::Ok(
#ident{
#( #build_fields, )*
}
)
}
}
};
// 输出宏生成的代码
eprintln!(
"============ TOKENS ============\n{}\n================================",
expanded
);
TokenStream::from(expanded)
}
// 提取器
fn get_inner_type<'a>(ty: &'a syn::Type, ident_name: &str) -> std::option::Option<&'a syn::Type> {
// 判断是不是一个路径类型
if let syn::Type::Path(syn::TypePath { path, .. }) = ty {
// 提取末尾
if let std::option::Option::Some(segment) = path.segments.last() {
if segment.ident == ident_name {
// 这里使用传入的 ident_name
if let syn::PathArguments::AngleBracketed(syn::AngleBracketedGenericArguments {
args,
..
}) = &segment.arguments
{
if let std::option::Option::Some(syn::GenericArgument::Type(inner_ty)) =
args.first()
{
return std::option::Option::Some(inner_ty);
}
}
}
}
}
std::option::Option::None
}
// 提取属性标识符,检查报错
fn get_each_attr(
f: &syn::Field,
) -> std::result::Result<std::option::Option<syn::Ident>, syn::Error> {
for attr in &f.attrs {
if attr.path().is_ident("builder") {
let mut each_ident = std::option::Option::None;
// 解析属性内部键值对
let res = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("each") {
let value = meta.value()?; // 拿到 "=" 后面的内容
let s: syn::LitStr = value.parse()?; // 解析成字符串常量
each_ident = std::option::Option::Some(syn::Ident::new(&s.value(), s.span()));
std::result::Result::Ok(())
} else {
std::result::Result::Err(meta.error("expected `builder(each = \"...\")`"))
}
});
if let std::result::Result::Err(_) = res {
return std::result::Result::Err(syn::Error::new_spanned(
&attr.meta,
"expected `builder(each = \"...\")`",
));
}
return std::result::Result::Ok(each_ident);
}
}
std::result::Result::Ok(std::option::Option::None)
}
debug
依旧是派生宏,一点点支持自定义格式、泛型参数、幽灵类型(phantom)、关联类型和逃生舱:
use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Data, DeriveInput, Fields, Type, TypePath};
#[proc_macro_derive(CustomDebug, attributes(debug))]
pub fn derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as DeriveInput);
let ident = ast.ident;
let ident_string = ident.to_string();
// 提取字段
let fields = match &ast.data {
Data::Struct(data_struct) => match &data_struct.fields {
Fields::Named(fields_named) => &fields_named.named,
_ => panic!("panic!"),
},
_ => panic!("panic!"),
};
// 提取用户传入的辅助属性
let custom_bound = get_struct_custom_bound(&ast.attrs);
let mut generics = ast.generics.clone();
if let std::option::Option::Some(bound_string) = custom_bound {
// 对于有"逃生舱"的情况,直接用作where子句
let predicate: syn::WherePredicate = syn::parse_str(&bound_string).unwrap();
generics.make_where_clause().predicates.push(predicate);
} else {
// 没有的情况下,自己匹配
// 获取所有泛型参数的名字
let mut type_params = std::collections::HashSet::new();
for param in &mut generics.params {
if let syn::GenericParam::Type(type_param) = param {
type_params.insert(type_param.ident.to_string());
}
}
// 找出所有在 PhantomData 中出现的泛型参数
let mut phantom_generics = std::collections::HashSet::new();
let mut associated_bounds = Vec::new();
let mut associated_type_hosts = std::collections::HashSet::new();
for field in fields {
// 记录幽灵类型的泛型参数
if let Some(generic_name) = get_phantom_data_generic_name(&field.ty).unwrap() {
phantom_generics.insert(generic_name);
}
// 记录关联类型
let assoc_types = get_associated_types(&field.ty, &type_params);
for ty_path in assoc_types {
// 记录宿主名字 (比如 "T")
let host_ident = ty_path.path.segments[0].ident.to_string();
associated_type_hosts.insert(host_ident);
// 记录完整的关联类型路径 (比如 T::Value)
associated_bounds.push(ty_path);
}
}
// 在正常的为所有泛型加 Debug 约束的基础上,排除幽灵类型,排除关联类型
for param in &mut generics.params {
if let syn::GenericParam::Type(ref mut type_param) = *param {
let param_name = type_param.ident.to_string();
// 如果既不是幽灵泛型,也不是关联类型的宿主,才加约束
if !phantom_generics.contains(¶m_name)
&& !associated_type_hosts.contains(¶m_name)
{
type_param.bounds.push(syn::parse_quote!(std::fmt::Debug));
}
}
}
if !associated_bounds.is_empty() {
let where_clause = generics.make_where_clause();
for bound_ty in associated_bounds {
where_clause.predicates.push(syn::parse_quote! {
#bound_ty: std::fmt::Debug
});
}
}
}
// 方便的分隔函数
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
// 字段名和值
let field_string = fields.iter().map(|f| f.ident.as_ref().unwrap().to_string());
let field_value = fields.iter().map(|f| {
let ident = f.ident.as_ref().unwrap();
if let Some(format_str) = get_debug_attr(f).unwrap() {
quote! {
&format_args!(#format_str, &self.#ident)
}
} else {
quote! {
&self.#ident
}
}
});
let expanded = quote! {
impl #impl_generics std::fmt::Debug for #ident #ty_generics #where_clause {
fn fmt(&self, fmt: &mut std::fmt::Formatter) -> std::fmt::Result {
// 用标准库函数方便地构造
fmt.debug_struct(#ident_string)
#(
.field(#field_string, #field_value)
)*
.finish()
}
}
};
eprintln!(" TOKENS:\n {}", expanded);
// Hand the output tokens back to the compiler
TokenStream::from(expanded)
}
// 提取辅助属性,用于支持自定义格式
fn get_debug_attr(
f: &syn::Field,
) -> std::result::Result<std::option::Option<std::string::String>, syn::Error> {
for attr in &f.attrs {
if attr.path().is_ident("debug") {
if let syn::Meta::NameValue(meta_name_value) = &attr.meta {
// 匹配 表达式->字面量->字符
if let syn::Expr::Lit(expr_lit) = &meta_name_value.value {
// 确保这个字面量是一个字符串
if let syn::Lit::Str(lit_str) = &expr_lit.lit {
return std::result::Result::Ok(std::option::Option::Some(lit_str.value()));
}
}
}
}
}
std::result::Result::Ok(std::option::Option::None)
}
// 考虑phantom类型
// marker: PhantomData<T>
fn get_phantom_data_generic_name(
ty: &syn::Type,
) -> std::result::Result<std::option::Option<std::string::String>, syn::Error> {
if let syn::Type::Path(syn::TypePath { path, .. }) = ty {
if let Some(segment) = path.segments.last() {
if segment.ident == "PhantomData" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(syn::Type::Path(inner_path))) =
args.args.first()
{
if let Some(inner_segment) = inner_path.path.segments.first() {
return std::result::Result::Ok(std::option::Option::Some(
inner_segment.ident.to_string(),
));
}
}
}
}
}
}
std::result::Result::Ok(std::option::Option::None)
}
// 识别关联函数
fn get_associated_types(
ty: &Type,
type_params: &std::collections::HashSet<String>,
) -> Vec<TypePath> {
let mut associated_types = Vec::new();
if let Type::Path(type_path) = ty {
// 检查当前路径是不是 T::Value 的形式
if type_path.path.segments.len() >= 2 {
let first_segment = &type_path.path.segments[0].ident;
if type_params.contains(&first_segment.to_string()) {
associated_types.push(type_path.clone());
}
}
// 无论是不是 T::Value,都要检查有没有尖括号参数 (比如 Vec<...>),有的话就递归进去
for segment in &type_path.path.segments {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
for arg in &args.args {
if let syn::GenericArgument::Type(inner_ty) = arg {
// 递归调用,把深层找到的关联类型合并进来
associated_types.extend(get_associated_types(inner_ty, type_params));
}
}
}
}
}
associated_types
}
// 提取结构体上的 #[debug(bound = "...")] 属性
fn get_struct_custom_bound(attrs: &[syn::Attribute]) -> Option<String> {
for attr in attrs {
if attr.path().is_ident("debug") {
// 解析 #[debug(bound = "...")]
let mut res = None;
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("bound") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
res = Some(s.value());
}
Ok(())
});
if res.is_some() {
return res;
}
}
}
None
}
seq
类函数宏,需要新建一个结构体承载每一个token,然后对于body部分,要能支持任意的代码(比如宏,比如自定义规则的"gadget"),要既能解析..也能解析..=
use proc_macro::TokenStream;
use proc_macro2::{Delimiter, Group, Literal, TokenStream as TokenStream2, TokenTree};
use quote::quote;
use syn::parse::{Parse, ParseStream, Result};
use syn::{parse_macro_input, DeriveInput, Ident, LitInt, Token};
//syn::Ident, Token![in], syn::LitInt,
struct Seq {
iden: Ident,
in_token: Token![in],
start: LitInt, // 对应 0
inclusive: bool,
end: LitInt, // 对应 8
body: TokenStream2, // 对应 { ... } 里面的所有代码
}
impl Parse for Seq {
fn parse(input: ParseStream) -> Result<Self> {
let iden: Ident = input.parse()?;
let in_token = input.parse::<Token![in]>()?;
let start: LitInt = input.parse()?;
let inclusive = if input.peek(Token![..=]) {
input.parse::<Token![..=]>()?; // 确认是 ..=,吃掉
true
} else if input.peek(Token![..]) {
input.parse::<Token![..]>()?; // 确认是 ..,吃掉
false
} else {
//啥都不是,报错吧
return Err(input.error("expected `..` or `..=`"));
};
let end: LitInt = input.parse()?;
// 解析大括号 { ... }
// content 是一个临时的游标,指向大括号内部的代码流
let content;
syn::braced!(content in input);
let body: TokenStream2 = content.parse()?;
Ok(Seq {
iden,
in_token,
start,
inclusive,
end,
body,
})
}
}
// 单次展开
fn expand_body_single(body: &TokenStream2, iden: &Ident, index: usize) -> TokenStream2 {
let mut expanded = TokenStream2::new();
let tokens: Vec<TokenTree> = body.clone().into_iter().collect();
let mut slice = tokens.as_slice();
while !slice.is_empty() {
// 处理f~N模式
if let [TokenTree::Ident(prefix), TokenTree::Punct(punct), TokenTree::Ident(ident), rest @ ..] =
slice
{
if punct.as_char() == '~' && ident == iden {
// 1. 拼接新名字
let new_name = format!("{}{}", prefix.to_string(), index);
// 传入span以提供调用出的位置信息
let new_ident = Ident::new(&new_name, prefix.span());
expanded.extend(quote!(#new_ident));
slice = rest;
continue;
}
}
let tt = &slice[0];
match tt {
// 如果是 Group(()、[]、{})
TokenTree::Group(g) => {
// 递归调用自己,去括号里面继续寻找并替换
let inner_expanded = expand_body_single(&g.stream(), iden, index);
// 把替换后的内容重新包回原来的括号类型中
let mut new_group = Group::new(g.delimiter(), inner_expanded);
// 传入span
new_group.set_span(g.span());
expanded.extend(quote!(#new_group));
}
// 如果是标识符 (Ident),且名字刚好就是我们要找的 (比如 `N`)
TokenTree::Ident(ref i) if i == iden => {
// 生成一个新的数字字面量,比如 0, 1, 2
let mut lit = Literal::usize_unsuffixed(index);
lit.set_span(i.span()); // 继承 N 的位置信息
expanded.extend(quote!(#lit));
}
// 情况 3 Punction Literal,其它标识符,不变
_ => {
expanded.extend(quote!(#tt));
}
}
slice = &slice[1..];
}
expanded
}
// 匹配外部模式
fn find_and_expand_pound(
body: TokenStream2,
target_iden: &Ident,
start: usize,
end: usize,
) -> (TokenStream2, bool) {
let mut expanded = TokenStream2::new();
let mut found_pattern = false; // 记录是否找到了 #(...) *
let tokens: Vec<TokenTree> = body.into_iter().collect();
let mut slice = tokens.as_slice();
while !slice.is_empty() {
// 匹配 # (…) *
if let [TokenTree::Punct(hash), TokenTree::Group(group), TokenTree::Punct(star), rest @ ..] =
slice
{
if hash.as_char() == '#'
&& group.delimiter() == Delimiter::Parenthesis
&& star.as_char() == '*'
{
found_pattern = true;
// 找到了之后内部展开 start..end 次
for i in start..end {
let inner_expanded = expand_body_single(&group.stream(), target_iden, i);
expanded.extend(inner_expanded);
}
// 前进
slice = rest;
continue;
}
}
// 如果不是 #(...)*,那就原样输出(如果是普通的 Group,需要递归进去找)
let tt = &slice[0];
match tt {
TokenTree::Group(g) => {
// 递归往更深层的括号里去找
let (inner_expanded, inner_found) =
find_and_expand_pound(g.stream(), target_iden, start, end);
if inner_found {
found_pattern = true; // 只要里面找到了,外面也要标记为 true
}
let mut new_group = Group::new(g.delimiter(), inner_expanded);
new_group.set_span(g.span());
expanded.extend(quote!(#new_group));
}
_ => {
expanded.extend(quote!(#tt)); // 不是 Group,直接搬运
}
}
slice = &slice[1..];
}
(expanded, found_pattern)
}
#[proc_macro]
pub fn seq(input: TokenStream) -> TokenStream {
// Parse the input tokens into a syntax tree
let seq = parse_macro_input!(input as Seq);
// base10_parse() 把数字字符串解析成对应的数值类型
let start = seq.start.base10_parse::<usize>().unwrap();
let mut end = seq.end.base10_parse::<usize>().unwrap();
// ..=模式,多匹配一个
if seq.inclusive {
end += 1;
}
let (expanded, found) = find_and_expand_pound(seq.body.clone(), &seq.iden, start, end);
if found {
// 模式 A:如果找到了,说明外部代码不要重复,直接返回处理后的结果
proc_macro::TokenStream::from(expanded)
} else {
// 模式 B:如果没找到,全局循环展开
let mut stream = TokenStream2::new();
(start..end).for_each(|i| {
let exp = expand_body_single(&seq.body, &seq.iden, i);
stream.extend(exp);
});
proc_macro::TokenStream::from(stream)
}
}
sorted
属性宏
use proc_macro::TokenStream;
use quote::quote;
use syn::visit_mut::VisitMut;
use syn::{parse_macro_input, visit_mut, Error, ExprMatch, Item, ItemFn};
#[proc_macro_attribute]
pub fn sorted(args: TokenStream, input: TokenStream) -> TokenStream {
let _ = args;
// 保留一份以便报错用
let original_input = proc_macro2::TokenStream::from(input.clone());
let item = parse_macro_input!(input as Item);
match check_and_expand(item) {
Ok(expanded) => TokenStream::from(expanded),
Err(e) => {
let mut output = e.to_compile_error();
output.extend(original_input);
proc_macro::TokenStream::from(output)
}
}
}
// 实现排序识别逻辑 -> 紧邻的两个比较
fn check_and_expand(item: Item) -> syn::Result<proc_macro2::TokenStream> {
match item {
Item::Enum(item_enum) => {
let mut variants_iter = item_enum.variants.iter();
// 以第一个元素作为 try_fold
if let Some(first) = variants_iter.next() {
variants_iter.try_fold(first, |prev, current| {
let current_name = current.ident.to_string();
let prev_name = prev.ident.to_string();
if current_name < prev_name {
let should_be_before = item_enum
.variants
.iter()
.find(|v| v.ident.to_string() > current_name)
.unwrap();
let msg = format!(
"{} should sort before {}",
current.ident, should_be_before.ident
);
Err(Error::new(current.ident.span(), msg))
} else {
Ok(current)
}
})?;
}
Ok(quote!(#item_enum))
}
_ => Err(Error::new(
proc_macro2::Span::call_site(),
"expected enum or match expression",
)),
}
}
#[proc_macro_attribute]
pub fn check(args: TokenStream, input: TokenStream) -> TokenStream {
let _args = args;
let mut item_fn = parse_macro_input!(input as ItemFn);
let mut visitor = MatchVisitor { errors: Vec::new() };
// 让 Visitor 深入函数体,寻找并修改所有的 match 表达式
visitor.visit_item_fn_mut(&mut item_fn);
// 把收集到的错误转化为 TokenStream,并把修改后的函数拼在后面
let mut output = proc_macro2::TokenStream::new();
for err in visitor.errors {
output.extend(err.to_compile_error());
}
// 把去掉了 #[sorted] 属性的合法函数代码吐给编译器
output.extend(quote!(#item_fn));
TokenStream::from(output)
}
struct MatchVisitor {
errors: Vec<syn::Error>,
}
impl VisitMut for MatchVisitor {
fn visit_expr_match_mut(&mut self, node: &mut ExprMatch) {
let sorted_index = node
.attrs
.iter_mut()
.position(|attr| attr.path().is_ident("sorted"));
if let std::option::Option::Some(idx) = sorted_index {
// 删除#[sorted]标签
node.attrs.remove(idx);
let mut prev_name = String::new();
for arm in &node.arms {
// 对于通配符syn::Wild要特判
if let syn::Pat::Wild(_) = arm.pat {
continue;
}
let (current_name, span_node) = match get_pat_info(&arm.pat) {
Some(info) => info,
None => {
self.errors.push(syn::Error::new_spanned(
&arm.pat,
"unsupported by #[sorted]",
));
return;
}
};
if current_name >= prev_name {
prev_name = current_name;
} else {
// 乱序逻辑
let should_be_before = node
.arms
.iter()
.find(|a| get_pat_info(&a.pat).is_some_and(|(n, _)| n > current_name))
.unwrap();
let (should_be_name, _) = get_pat_info(&should_be_before.pat).unwrap();
let msg = format!("{} should sort before {}", current_name, should_be_name);
self.errors.push(syn::Error::new_spanned(span_node, msg));
return;
}
}
}
// Delegate to the default impl to visit nested expressions.
visit_mut::visit_expr_match_mut(self, node);
}
}
fn path_to_string(path: &syn::Path) -> String {
path.segments
.iter()
.map(|segment| segment.ident.to_string())
.collect::<Vec<_>>()
.join("::")
}
// 提取模式的字符串名称
fn get_pat_info(pat: &syn::Pat) -> Option<(String, &dyn quote::ToTokens)> {
match pat {
syn::Pat::Ident(pat_ident) => {
let name = pat_ident.ident.to_string();
let first_char = name.chars().next().unwrap();
// 对于_other,不处理
if first_char.is_lowercase() || first_char == '_' {
None
} else {
Some((name, &pat_ident.ident))
}
}
syn::Pat::Path(pat_path) => Some((path_to_string(&pat_path.path), &pat_path.path)),
syn::Pat::Struct(pat_struct) => Some((path_to_string(&pat_struct.path), &pat_struct.path)),
syn::Pat::TupleStruct(pat_tuple) => {
Some((path_to_string(&pat_tuple.path), &pat_tuple.path))
}
_ => None,
}
}
bitfield
最终关。分为两个文件.
// impl/src/lib.rs
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::{format_ident, quote};
use syn::{parse_macro_input, Fields, Item, Item::Struct};
#[proc_macro_attribute]
pub fn bitfield(args: TokenStream, input: TokenStream) -> TokenStream {
let item = parse_macro_input!(input as Item);
let item_struct = match &item {
Struct(item_struct) => item_struct,
_ => panic!("怎么不是结构体"),
};
let ident = &item_struct.ident;
let fields = match &item_struct.fields {
Fields::Named(fields_named) => &fields_named.named,
_ => panic!("这是什么字段"),
};
// 提取类型
let field_tys = fields.iter().map(|field| &field.ty);
// 可见性保持一致
let vis = &item_struct.vis;
// 计算大小
let total_bits = quote! {
0usize #( + <#field_tys as ::bitfield::Specifier>::BITS )*
};
let size = quote! { (#total_bits) / 8usize };
let (methods, bits_assertions) = process_fields(fields);
let output = quote! {
#[repr(C)]
#vis struct #ident {
data : [u8; #size],
}
const _: () = {
struct ZeroMod8;
struct OneMod8;
struct TwoMod8;
struct ThreeMod8;
struct FourMod8;
struct FiveMod8;
struct SixMod8;
struct SevenMod8;
trait TotalSizeIsMultipleOfEightBits {}
struct _AssertTotalSizeIsMultipleOfEightBits
where
<::bitfield::checks::Check<{ (#total_bits) % 8usize }>
as ::bitfield::checks::MapMod8>::Output:
::bitfield::checks::TotalSizeIsMultipleOfEightBits;
#( #bits_assertions )*
};
impl #ident {
pub fn new() -> Self {
Self {
data: [0u8; #size],
}
}
#( #methods )*
}
};
/*eprintln!(
"============ TOKENS ============\n{}\n================================",
output
);*/
output.into()
}
use syn::ItemEnum;
#[proc_macro_derive(BitfieldSpecifier)]
pub fn derive_bitfield_specifier(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as ItemEnum);
let ident = input.ident;
let variants = input.variants;
// 检查是否是 2 的幂次方
let count = variants.len();
if !count.is_power_of_two() {
let err = syn::Error::new(
Span::call_site(),
"BitfieldSpecifier expected a number of variants which is a power of 2",
);
return err.to_compile_error().into();
}
let bits = count.trailing_zeros() as usize;
let variant_idents: Vec<_> = variants.iter().map(|v| &v.ident).collect();
// 检查
let checks = variants.iter().map(|variant| {
let v_ident = &variant.ident;
// 获取span
quote::quote_spanned! { v_ident.span() =>
const _: () = {
struct True;
struct False;
trait DiscriminantInRange {}
type Arg = <::bitfield::checks::CheckInRange<{
(#ident::#v_ident as u64) < (1u64 << #bits)
}> as ::bitfield::checks::MapBool>::Output;
struct _AssertDiscriminantInRange
where
Arg: ::bitfield::checks::DiscriminantInRange;
};
}
});
let expanded = quote! {
impl ::bitfield::Specifier for #ident {
const BITS: usize = #bits;
type Accessor = Self; // 枚举的 Accessor 就是它自己
fn from_u64(val: u64) -> Self::Accessor {
#(
if val == (Self::#variant_idents as u64) {
return Self::#variant_idents;
}
)*
unreachable!("无效的 bit 组合")
}
fn into_u64(val: Self::Accessor) -> u64 {
val as u64
}
}
#( #checks )*
};
expanded.into()
}
#[proc_macro]
pub fn derive_b_types(input: TokenStream) -> TokenStream {
let mut output = quote! {};
(1..=64).for_each(|index| {
let index = index as usize;
let ident = format_ident!("B{}", index);
let accessor_ty = if index <= 8 {
quote! { u8 }
} else if index <= 16 {
quote! { u16 }
} else if index <= 32 {
quote! { u32 }
} else {
quote! { u64 }
};
let extand = quote! {
pub enum #ident {}
impl crate::Specifier for #ident {
const BITS : usize = #index;
type Accessor = #accessor_ty;
fn from_u64(val: u64) -> Self::Accessor { val as #accessor_ty }
fn into_u64(val: Self::Accessor) -> u64 { val as u64 }
}
};
output.extend(extand);
});
/*eprintln!(
"============ TOKENS ============\n{}\n================================",
output
);*/
proc_macro::TokenStream::from(output)
}
// 遍历item的各字段,处理辅助属性以及生成getter和setter
fn process_fields(
fields: &syn::punctuated::Punctuated<syn::Field, syn::token::Comma>,
) -> (Vec<proc_macro2::TokenStream>, Vec<proc_macro2::TokenStream>) {
let mut methods = Vec::new();
let mut bits_assertions = Vec::new();
let mut current_offset = quote! { 0usize }; // 记录偏移
fields.iter().for_each(|field| {
let field_name = field.ident.as_ref().unwrap();
let field_ty = &field.ty;
let getter_name = format_ident!("get_{}", field_name);
let setter_name = format_ident!("set_{}", field_name);
// 寻找并解析出 #[bits = N] 属性
let mut declared_bits = None;
for attr in &field.attrs {
if attr.path().is_ident("bits") {
if let syn::Meta::NameValue(meta) = &attr.meta {
if let syn::Expr::Lit(syn::ExprLit {
lit: syn::Lit::Int(lit_int),
..
}) = &meta.value
{
if let Ok(n) = lit_int.base10_parse::<usize>() {
// 保留把 lit_int.span()
declared_bits = Some((n, lit_int.span()));
}
}
}
}
}
// 如果写了 #[bits = N],生成编译期相等断言
if let Some((n, span)) = declared_bits {
bits_assertions.push(quote::quote_spanned! { span =>
const _: [(); #n] = [(); <#field_ty as ::bitfield::Specifier>::BITS];
});
}
// 生成该字段的 getter 和 setter
methods.push(quote! {
pub fn #getter_name(&self) -> <#field_ty as ::bitfield::Specifier>::Accessor {
type Accessor = <#field_ty as ::bitfield::Specifier>::Accessor;
let offset = #current_offset;
let bits = <#field_ty as ::bitfield::Specifier>::BITS;
let mut val = 0u64;
for i in 0..bits {
let bit_idx = offset + i;
let byte_idx = bit_idx / 8;
let bit_in_byte = bit_idx % 8;
let bit = (self.data[byte_idx] >> bit_in_byte) & 1;
val |= (bit as u64) << i;
}
<#field_ty as ::bitfield::Specifier>::from_u64(val)
}
pub fn #setter_name(&mut self, val : <#field_ty as ::bitfield::Specifier>::Accessor) {
let offset = #current_offset;
let bits = <#field_ty as ::bitfield::Specifier>::BITS;
let val = <#field_ty as ::bitfield::Specifier>::into_u64(val);
// 逐位写入
for i in 0..bits {
let bit_idx = offset + i;
let byte_idx = bit_idx / 8;
let bit_in_byte = bit_idx % 8;
let bit = (val >> i) & 1;
if bit == 1 {
self.data[byte_idx] |= 1 << bit_in_byte;
} else {
self.data[byte_idx] &= !(1 << bit_in_byte);
}
}
}
});
// 4. 更新下一个字段的位移
current_offset = quote! {
#current_offset + <#field_ty as ::bitfield::Specifier>::BITS
};
});
(methods, bits_assertions)
}
// bitfield/src/lib.rs
// Crates that have the "proc-macro" crate type are only allowed to export
// procedural macros. So we cannot have one crate that defines procedural macros
// alongside other types of public APIs like traits and structs.
//
// For this project we are going to need a #[bitfield] macro but also a trait
// and some structs. We solve this by defining the trait and structs in this
// crate, defining the attribute macro in a separate bitfield-impl crate, and
// then re-exporting the macro from this crate so that users only have one crate
// that they need to import.
//
// From the perspective of a user of this crate, they get all the necessary APIs
// (macro, trait, struct) through the one bitfield crate.
pub use bitfield_impl::*;
pub trait Specifier {
const BITS: usize;
// 利用一个关联类型处理类型窄化
type Accessor;
// 处理内部所用的u64的类型转换
fn from_u64(val: u64) -> Self::Accessor;
fn into_u64(val: Self::Accessor) -> u64;
}
impl Specifier for bool {
const BITS: usize = 1;
type Accessor = bool;
fn from_u64(val: u64) -> Self::Accessor {
val != 0 // 只要不是 0 就是 true
}
fn into_u64(val: Self::Accessor) -> u64 {
val as u64 // true 变 1,false 变 0
}
}
pub mod checks {
// 用一个神奇的特征提供报错信息
pub trait TotalSizeIsMultipleOfEightBits {}
pub trait DiscriminantInRange {
type Check;
}
pub struct True;
pub struct False;
pub struct ZeroMod8;
pub struct OneMod8;
pub struct TwoMod8;
pub struct ThreeMod8;
pub struct FourMod8;
pub struct FiveMod8;
pub struct SixMod8;
pub struct SevenMod8;
// 只为取模后为0的实现这个特征
impl TotalSizeIsMultipleOfEightBits for ZeroMod8 {}
// 处理枚举类型但是值超出了默认
impl DiscriminantInRange for True {
type Check = ();
}
pub struct CheckInRange<const B: bool>;
pub struct Check<const R: usize>;
pub trait MapMod8 {
type Output;
}
pub trait MapBool {
type Output;
}
impl MapBool for CheckInRange<true> {
type Output = True;
}
impl MapBool for CheckInRange<false> {
type Output = False;
}
impl MapMod8 for Check<0> {
type Output = ZeroMod8;
}
impl MapMod8 for Check<1> {
type Output = OneMod8;
}
impl MapMod8 for Check<2> {
type Output = TwoMod8;
}
impl MapMod8 for Check<3> {
type Output = ThreeMod8;
}
impl MapMod8 for Check<4> {
type Output = FourMod8;
}
impl MapMod8 for Check<5> {
type Output = FiveMod8;
}
impl MapMod8 for Check<6> {
type Output = SixMod8;
}
impl MapMod8 for Check<7> {
type Output = SevenMod8;
}
}
bitfield_impl::derive_b_types!();