Macro in rust

本文又名《重生之我要成为rust魔法师》《五题三库粉碎魔法梦》

学习一下rust宏,参考资料:

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+1hello_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

大致在两种情况下,用类函数宏是合理的:

  1. 声明宏变得越来越臃肿、难以维护 时
  2. 需要一个编译器需要执行,但是const fn又做不到的函数 时。比如phf库,在编译时把提供的一堆字符串算成一个完美的哈希表。

属性宏

属性宏也替换其作用域的item,它的输入除了宏的部分(属性名及其参数),还包括附加到的整个item

他可以很容易地把一个函数变成另一个模样,就像#[tokio::main] #[test]做的那样

属性宏是权力最大的宏,它的使用场景也比较多:

  • 生成测试
  • 框架胶水,典型的比如#[tokio::main],其实是重写了整个main函数
  • 透明中间件
  • 类型转换:改定义,比如增加字段。

派生宏

派生宏的目标和前两种不太一样,它不替换只附加。

它的限制最多:只能追加,不能带参数,用辅助属性来传递额外信息

它应当且只应当用在一个地方——在可能的情况下,自动实现一个特征,同时需要满足两个条件—— 1. 使用频率极高,否则对不起写宏花的时间 2. 逻辑必须符合直觉

典型的正面案例就是#[derive(Serialize,Deserialize)] #[derive(Debug)] #[derive(Clone)]

开销

过程宏会增长编译时间,具体体现在两个方面:

  • 引入一些很重的依赖,写过程宏需要的syn crate在所有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地拓展,或者使用TokenStreamFromStr方法,或者混合这两种……或者用毫无疑问的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宏领域里那个绕不开的名字,synquoteproc-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(&param_name)
                    && !associated_type_hosts.contains(&param_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!();