如何编写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
相关产品推荐
相关产品推荐

