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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 09:09:55