You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何编写Rust宏及选择类型?实现FromRequest的Scope复用宏

实现可复用的Scope验证FromRequest宏

问题描述

我有一段为ReadMe结构体实现FromRequest trait的异步代码,现在需要为多个结构体复用该实现,每个结构体对应不同的Scope名称(比如示例中的"ReadMe")。目前只能通过复制代码并修改Scope名称来实现,希望改成通过#[derive(Scope("ReadMe"))]的宏方式复用。请问如何编写这个Rust宏,应该选择哪种宏类型?

原实现代码如下:

#[rocket::async_trait]
impl<'r> FromRequest<'r> for ReadMe {
    type Error = AuthErrors;

    async fn from_request(request: &'r Request<'_>) -> request::Outcome<Self, Self::Error> {
        let db = request
            .guard::<&State<mongodb::Database>>()
            .await
            .expect("No database state!");
        let token = request.headers().get_one("Authorization");

        match token {
            Some(token) => {
                // 验证有效性
                let user = get_user_by_token(token.to_string(), db.inner().to_owned()).await;

                match user {
                    Some(user) => {
                        if user.tokenData.scopes.split(" ").any(|x| x == "ReadMe") {
                            Outcome::Success(ReadMe(Some(BasicScope {
                                user: user.user,
                                value: true,
                            })))
                        } else {
                            Outcome::Failure((Status::Forbidden, AuthErrors::ScopeMissing))
                        }
                    }
                    None => Outcome::Failure((Status::Unauthorized, AuthErrors::Invalid)), // TODO?: 是否需要返回其他值?
                }
            }
            None => Outcome::Failure((Status::Unauthorized, AuthErrors::Missing)),
        }
    }
}

解决方案

1. 统一结构体结构

首先让所有需要复用该逻辑的结构体遵循相同结构:包裹一个Option<BasicScope>,示例如下:

use rocket::{request::{FromRequest, Outcome, request}, http::Status, State};
use mongodb::Database;

// 假设你已定义以下类型
#[derive(Debug)]
enum AuthErrors {
    Missing,
    Invalid,
    ScopeMissing,
}

#[derive(Debug)]
struct BasicScope {
    user: User, // 替换为你的实际用户类型
    value: bool,
}

struct User; // 示例用户类型
async fn get_user_by_token(token: String, db: Database) -> Option<{ tokenData: { scopes: String }, user: User }> {
    // 你的用户查询逻辑
    None
}

2. 使用过程宏实现自定义Derive

这里选择过程宏(proc-macro),它能在编译时解析结构体属性并生成对应的trait实现,完全符合#[derive(...)]的使用需求。

第一步:添加依赖

在Cargo.toml中添加过程宏相关依赖:

[dependencies]
rocket = "0.5.0-rc.3" # 根据你的Rocket版本调整
mongodb = "2.0"
proc-macro2 = "1.0"
quote = "1.0"
syn = { version = "2.0", features = ["full", "derive"] }

[lib]
proc-macro = true

第二步:编写过程宏代码

use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, DeriveInput, LitStr};

#[proc_macro_derive(Scope, attributes(scope))]
pub fn derive_scope(input: TokenStream) -> TokenStream {
    let input = parse_macro_input!(input as DeriveInput);
    let struct_name = &input.ident;
    
    // 提取#[scope("xxx")]属性中的Scope名称
    let scope_name = input.attrs.iter()
        .find(|attr| attr.path().is_ident("scope"))
        .and_then(|attr| attr.parse_args::<LitStr>().ok())
        .map(|lit| lit.value())
        .expect("必须添加#[scope(\"scope_name\")]属性");

    let expanded = quote! {
        #[rocket::async_trait]
        impl<'r> FromRequest<'r> for #struct_name {
            type Error = AuthErrors;

            async fn from_request(request: &'r Request<'_>) -> request::Outcome<Self, Self::Error> {
                let db = request
                    .guard::<&State<mongodb::Database>>()
                    .await
                    .expect("No database state!");
                let token = request.headers().get_one("Authorization");

                match token {
                    Some(token) => {
                        let user = get_user_by_token(token.to_string(), db.inner().to_owned()).await;

                        match user {
                            Some(user) => {
                                if user.tokenData.scopes.split(" ").any(|x| x == #scope_name) {
                                    Outcome::Success(#struct_name(Some(BasicScope {
                                        user: user.user,
                                        value: true,
                                    })))
                                } else {
                                    Outcome::Failure((Status::Forbidden, AuthErrors::ScopeMissing))
                                }
                            }
                            None => Outcome::Failure((Status::Unauthorized, AuthErrors::Invalid)),
                        }
                    }
                    None => Outcome::Failure((Status::Unauthorized, AuthErrors::Missing)),
                }
            }
        }
    };

    expanded.into()
}

第三步:使用宏

在需要的结构体上添加#[derive(Scope)]和#[scope("xxx")]属性即可:

// 导入你的宏
use your_macro_crate::Scope;

#[derive(Debug, Scope)]
#[scope("ReadMe")]
struct ReadMe(Option<BasicScope>);

#[derive(Debug, Scope)]
#[scope("WritePost")]
struct WritePost(Option<BasicScope>);

3. 优化:提取通用验证逻辑

为了让宏生成的代码更简洁易维护,可以把验证逻辑提取成通用函数:

async fn validate_scope(request: &Request<'_>, scope: &str) -> Outcome<BasicScope, AuthErrors> {
    let db = request
        .guard::<&State<mongodb::Database>>()
        .await
        .expect("No database state!");
    let token = request.headers().get_one("Authorization");

    match token {
        Some(token) => {
            let user = get_user_by_token(token.to_string(), db.inner().to_owned()).await;

            match user {
                Some(user) => {
                    if user.tokenData.scopes.split(" ").any(|x| x == scope) {
                        Outcome::Success(BasicScope {
                            user: user.user,
                            value: true,
                        })
                    } else {
                        Outcome::Failure((Status::Forbidden, AuthErrors::ScopeMissing))
                    }
                }
                None => Outcome::Failure((Status::Unauthorized, AuthErrors::Invalid)),
            }
        }
        None => Outcome::Failure((Status::Unauthorized, AuthErrors::Missing)),
    }
}

然后修改宏生成的代码:

let expanded = quote! {
    #[rocket::async_trait]
    impl<'r> FromRequest<'r> for #struct_name {
        type Error = AuthErrors;

        async fn from_request(request: &'r Request<'_>) -> request::Outcome<Self, Self::Error> {
            match validate_scope(request, #scope_name).await {
                Outcome::Success(scope) => Outcome::Success(#struct_name(Some(scope))),
                Outcome::Failure(err) => Outcome::Failure(err),
                Outcome::Forward(forward) => Outcome::Forward(forward),
            }
        }
    }
};

内容的提问来源于stack exchange,提问作者piano1029

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 06:55:45