使用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"
错误原因
- 返回类型不匹配:
call方法要求返回BoxFuture<'static, Result<Response, S::Error>>,但错误分支直接返回了Err((StatusCode, &str)),这是一个Result类型而非BoxFuture。 - 错误类型不兼容: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(利用IntoResponsetrait,(StatusCode, &str)可以自动转换成响应),然后通过Ok(response)返回,这样既符合返回类型要求,又能正确返回错误状态码和信息。 poll_ready方法中需要将内部服务的错误转换成axum::Error,使用map_err(Into::into)完成转换。
内容的提问来源于stack exchange,提问作者baby195lxl
相关产品推荐
相关产品推荐

