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

使用tower组件实现Axum JWT认证中间件编译报错求助

解决Axum自定义JWT中间件编译错误

问题场景

在Axum中实现自定义JWT认证中间件时,验证失败返回错误的分支无法通过编译,错误提示为:

Err(_) => Err((StatusCode::BAD_REQUEST, "bad request")),
   |                         ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ expected `Pin<Box<dyn Future<Output = ...> + Send>>`, found `Result<_, (StatusCode, &str)>`

原代码

自定义中间件代码(custom_middleware.rs)

use axum::http::StatusCode;
use axum::{extract::Request, response::Response};
use futures_util::future::BoxFuture;
use jsonwebtoken::{
    decode, errors::Error as JwtError, Algorithm, DecodingKey, TokenData, Validation,
};
use serde::{Deserialize, Serialize};
use std::task::{Context, Poll};
use tower::{Layer, Service};

#[derive(Serialize, Deserialize)]
pub struct Claims {
    pub id: usize,
    pub exp: usize,
}

#[derive(Clone)]
pub struct MyLayer;

impl<S> Layer<S> for MyLayer {
    type Service = MyMiddleware<S>;

    fn layer(&self, inner: S) -> Self::Service {
        MyMiddleware { inner }
    }
}

#[derive(Clone)]
pub struct MyMiddleware<S> {
    inner: S,
}

impl<S> Service<Request> for MyMiddleware<S>
where
    S: Service<Request, Response = Response> + Send + 'static,
    S::Future: Send + 'static,
{
    type Response = S::Response;
    type Error = S::Error;
    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;

    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        self.inner.poll_ready(cx)
    }

    fn call(&mut self, req: Request) -> Self::Future {
        match has_permission(&req) {
            Ok(_) => {
                let future = self.inner.call(req);
                Box::pin(async move {
                    let response: Response = future.await?;
                    Ok(response)
                })
            }
            Err(_) => Err((StatusCode::BAD_REQUEST, "bad request")),
        }
    }
}

fn has_permission(req: &Request) -> Result<TokenData<Claims>, (StatusCode, &'static str)> {
    let secret = "baby195lxl";
    let authorization_header_option = req.headers().get("authorization");
    if authorization_header_option.is_none() {
        return Err((StatusCode::BAD_REQUEST, "authorization header is none"));
    }
    let authentication_token: String = authorization_header_option
        .unwrap()
        .to_str()
        .unwrap_or("")
        .to_string();

    if authentication_token.is_empty() {
        return Err((StatusCode::BAD_REQUEST, "authorization header is empty"));
    }
    let token_result: Result<TokenData<Claims>, JwtError> = decode::<Claims>(
        &authentication_token,
        &DecodingKey::from_secret(secret.as_bytes()),
        &Validation::new(Algorithm::HS256),
    );
    match token_result {
        Ok(_token) => Ok(_token),
        Err(_e) => Err((StatusCode::UNAUTHORIZED, "Token Error")),
    }
}

主程序代码(main.rs)

use axum::{
    body::Bytes,
    extract::{Json, Request, State},
    routing::{get, post},
    Router,
};
use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
use serde::Deserialize;

use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tower_http::trace::TraceLayer;
use tracing::Span;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};

mod custom_middleware;
use custom_middleware::Claims;
use custom_middleware::MyLayer;

mod state;
use state::AppState;

#[derive(Deserialize, Debug, PartialEq)]
struct User {
    account: usize,
    password: String,
}

async fn register(State(state): State<AppState>, Json(user): Json<User>) -> String {
    let store_user = User {
        account: 195,
        password: "world".to_string(),
    };
    if user == store_user {
        let expiration = SystemTime::now() + Duration::from_secs(30 * 60);
        let exp_timestamp = expiration.duration_since(UNIX_EPOCH).unwrap().as_secs();
        let claims = Claims {
            id: user.account,
            exp: exp_timestamp as usize,
        };
        let token = encode(
            &Header::default(),
            &claims,
            &EncodingKey::from_secret(state.secret.as_bytes()),
        )
        .unwrap();
        token
    } else {
        "hello, world!".to_string()
    }
}

async fn login(State(state): State<AppState>, req: Request) -> Json<Claims> {
    let token = req
        .headers()
        .get("Authorization")
        .unwrap()
        .to_str()
        .unwrap();
    let payload = decode::<Claims>(
        token,
        &DecodingKey::from_secret(state.secret.as_bytes()),
        &Validation::new(Algorithm::HS256),
    )
    .unwrap();
    Json(payload.claims)
}

async fn protected(_req: Request) -> String {
    "World!".to_string()
}
#[tokio::main]
async fn main() {
    let state = AppState {
        secret: "baby195lxl".to_string(),
    };
    tracing_subscriber::registry()
        .with(tracing_subscriber::EnvFilter::new("debug"))
        .with(tracing_subscriber::fmt::layer())
        .init();

    let app = Router::new()
        .route("/protected", get(protected))
        .layer(MyLayer)
        .route("/register", post(register))
        .route("/login", post(login))
        .with_state(state)
        .layer(TraceLayer::new_for_http().on_body_chunk(
            |chunk: &Bytes, latency: Duration, _span: &Span| {
                tracing::debug!("streaming {} bytes in {:?}", chunk.len(), latency);
            },
        ));

    let listener = tokio::net::TcpListener::bind("127.0.0.1:5000")
        .await
        .unwrap();
    tracing::debug!("listening on {}", listener.local_addr().unwrap());
    axum::serve(listener, app).await.unwrap();
}

Cargo.toml配置

[dependencies]
axum = "^0.7"
tokio = { version = "^1.36", features = ["full"] }
tower-http = { version = "^0.5", features = ["trace"] }
tracing = "^0.1"
tracing-subscriber = { version = "^0.3", features = ["env-filter"] }
serde = { version = "1.0", features = ["derive"] }
jsonwebtoken = "9.2.0"
tower = "0.4.13"
futures-util = "0.3.30"

错误原因

  1. 返回类型不匹配:call方法要求返回BoxFuture<'static, Result<Response, S::Error>>,但错误分支直接返回了Err((StatusCode, &str)),这是一个Result类型而非BoxFuture。
  2. 错误类型不兼容:Axum的Service默认错误类型是axum::Error,自定义的(StatusCode, &str)元组无法直接作为错误返回。

解决方案

修改中间件的call方法,在错误分支中构建一个包含错误信息的Response,并将其包装到BoxFuture中返回;同时统一错误处理逻辑:

修正后的custom_middleware.rs关键部分

// 引入axum的Error和IntoResponse trait
use axum::{http::StatusCode, extract::Request, response::{Response, IntoResponse}, Error};
// ... 其他导入保持不变

impl<S> Service<Request> for MyMiddleware<S>
where
    S: Service<Request, Response = Response> + Send + 'static,
    S::Future: Send + 'static,
    S::Error: Into<Error>, // 增加约束,让内部服务的错误能转换成axum::Error
{
    type Response = S::Response;
    type Error = Error; // 将错误类型改为axum::Error
    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;

    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        self.inner.poll_ready(cx).map_err(Into::into)
    }

    fn call(&mut self, req: Request) -> Self::Future {
        match has_permission(&req) {
            Ok(_) => {
                let future = self.inner.call(req);
                Box::pin(async move {
                    let response = future.await.map_err(Into::into)?;
                    Ok(response)
                })
            }
            Err((status, msg)) => {
                // 构建错误响应并包装到BoxFuture中
                Box::pin(async move {
                    let response = (status, msg).into_response();
                    Ok(response)
                })
            }
        }
    }
}

// has_permission函数保持不变

补充说明

  • 把中间件的Error类型改为axum::Error,并添加S::Error: Into<Error>的约束,确保内部服务的错误能被正确转换。
  • 错误分支不再返回Err,而是直接构建符合Axum响应格式的Response(利用IntoResponse trait,(StatusCode, &str)可以自动转换成响应),然后通过Ok(response)返回,这样既符合返回类型要求,又能正确返回错误状态码和信息。
  • poll_ready方法中需要将内部服务的错误转换成axum::Error,使用map_err(Into::into)完成转换。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 04:34:51