如何在Axum的from_fn中向认证中间件传递可选参数?
在Axum中控制认证中间件注入内容的实现方案
要实现控制认证中间件向请求扩展(extensions)中注入user_id或完整User模型,核心是让中间件能接收路由传递的full_user参数,以下是两种可行方案:
方案一:使用高阶函数(推荐)
通过定义接收full_user参数的高阶函数,返回符合Axum中间件签名的闭包,实现路由级别的参数控制。
改造中间件代码
use axum::{ extract::State, middleware::Next, response::Response, http::{Request, StatusCode}, }; use mongodb::{Client, Collection, bson::doc, bson::oid::ObjectId}; use std::env; // 假设已定义User模型和verify_token验证函数 // struct User { /* 你的User结构体定义 */ } // fn verify_token(token: &str, secret: &str) -> Result<YourClaimsType, YourErrorType> { /* 验证逻辑 */ } pub fn auth(full_user: bool) -> impl Fn(State<Client>, Request<impl axum::body::HttpBody>, Next<impl axum::body::HttpBody>) -> impl std::future::Future<Output = Result<Response, StatusCode>> + Clone { move |State(client): State<Client>, mut req: Request<_>, next: Next<_>| async move { // 提取并验证Authorization Header let auth_token = req.headers() .get(http::header::AUTHORIZATION) .and_then(|h| h.to_str().ok()) .ok_or(StatusCode::UNAUTHORIZED)?; // 读取JWT密钥并验证token let jwt_secret = env::var("JWT_SECRET") .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let token_claims = verify_token(auth_token, &jwt_secret) .map_err(|_| StatusCode::UNAUTHORIZED)?; // 解析user_id并查询用户 let user_id = ObjectId::parse_str(&token_claims.sub) .map_err(|_| StatusCode::UNAUTHORIZED)?; let user_collection: Collection<User> = client.database("Merume").collection("users"); let user = user_collection.find_one(doc! {"_id": user_id}, None).await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? .ok_or(StatusCode::UNAUTHORIZED)?; // 根据full_user参数选择注入内容 if full_user { req.extensions_mut().insert(user); } else { req.extensions_mut().insert(user_id); } Ok(next.run(req).await) } }
路由层使用示例
// 需要注入完整User的路由 .route("/api/full-user", post(handlers::full_user_handler)) .route_layer(middleware::from_fn_with_state( client.clone(), auth_middleware::auth(true), )) // 只需要注入user_id的路由 .route("/api/only-user-id", post(handlers::user_id_handler)) .route_layer(middleware::from_fn_with_state( client.clone(), auth_middleware::auth(false), ))
处理函数示例
// 接收完整User的handler pub async fn full_user_handler( Extension(user): Extension<User>, // 其他请求参数 ) -> impl IntoResponse { // 业务逻辑,直接使用完整User对象 format!("Hello, {}", user.username) } // 接收user_id的handler pub async fn user_id_handler( Extension(user_id): Extension<ObjectId>, // 其他请求参数 ) -> impl IntoResponse { // 业务逻辑,使用user_id format!("User ID: {}", user_id) }
方案二:使用自定义State结构体
如果偏好通过State传递参数,可以定义包含client和full_user的结构体,将其作为中间件的状态。
定义自定义State
pub struct AuthState { pub client: Client, pub full_user: bool, }
修改中间件代码
pub async fn auth<B>( State(state): State<AuthState>, mut req: Request<B>, next: Next<B>, ) -> Result<Response, StatusCode> { // 验证逻辑和方案一一致,替换client为state.client let auth_token = req.headers() .get(http::header::AUTHORIZATION) .and_then(|h| h.to_str().ok()) .ok_or(StatusCode::UNAUTHORIZED)?; let jwt_secret = env::var("JWT_SECRET") .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let token_claims = verify_token(auth_token, &jwt_secret) .map_err(|_| StatusCode::UNAUTHORIZED)?; let user_id = ObjectId::parse_str(&token_claims.sub) .map_err(|_| StatusCode::UNAUTHORIZED)?; let user_collection: Collection<User> = state.client.database("Merume").collection("users"); let user = user_collection.find_one(doc! {"_id": user_id}, None).await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? .ok_or(StatusCode::UNAUTHORIZED)?; // 根据state.full_user选择注入内容 if state.full_user { req.extensions_mut().insert(user); } else { req.extensions_mut().insert(user_id); } Ok(next.run(req).await) }
路由层使用示例
.route("/api/full-user", post(handlers::full_user_handler)) .route_layer(middleware::from_fn_with_state( AuthState { client: client.clone(), full_user: true, }, auth_middleware::auth, )) .route("/api/only-user-id", post(handlers::user_id_handler)) .route_layer(middleware::from_fn_with_state( AuthState { client: client.clone(), full_user: false, }, auth_middleware::auth, ))
为什么你之前的写法失败?
你尝试直接在闭包中给auth函数额外传参,但Axum要求中间件函数的签名必须匹配State<T>, Request<B>, Next<B>,原auth函数没有full_user参数,导致签名不匹配。通过高阶函数或自定义State的方式,能让参数符合中间件的要求,同时实现路由级别的行为控制。
内容的提问来源于stack exchange,提问作者Azashiro Merume
相关产品推荐
相关产品推荐

