如何在Axum自定义Extractor中获取AppState实现JWT认证?
问题
使用Axum自定义提取器实现JWT认证时,能在提取器中打印State内容,但无法正确访问其中的密钥,当前JWT验证使用硬编码密钥,需修改代码从State中获取密钥完成验证。
解决方案
需要调整自定义提取器的泛型约束,使其能正确访问AppState中的密钥,同时统一Claims定义避免重复:
修改步骤
- 统一Claims定义:将
Claims移到authentication_token.rs中,在main.rs中引用该定义,避免重复。 - 调整提取器的泛型约束:让
ExtractAuthorization的FromRequestParts实现针对AppState,直接从传入的state中获取密钥。 - 替换硬编码密钥:将提取器中的硬编码密钥改为从
AppState读取。
修改后的完整代码
main.rs
use axum::body::Bytes; use axum::extract::{Json, Request, State}; use axum::{ routing::{get, post}, Router, }; use jsonwebtoken::{encode, Algorithm, EncodingKey, Header, Validation}; use serde::{Deserialize, Serialize}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tower_http::trace::TraceLayer; use tracing::Span; use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt}; mod authentication_token; use authentication_token::{ExtractAuthorization, Claims}; #[derive(Clone, Debug)] pub struct AppState { pub secret: String, } #[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> { use jsonwebtoken::decode; let token = req .headers() .get("Authorization") .unwrap() .to_str() .unwrap(); let payload = decode::<Claims>( token, &jsonwebtoken::DecodingKey::from_secret(state.secret.as_bytes()), &Validation::new(Algorithm::HS256), ) .unwrap(); Json(payload.claims) } async fn protected(_auth_token: ExtractAuthorization, req: Request) -> String { println!("{:?}", req); "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("/register", post(register)) .route("/login", post(login)) .route("/protected", get(protected)) .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(); }
authentication_token.rs
use axum::{ async_trait, extract::FromRequestParts, http::{header::AUTHORIZATION, request::Parts, StatusCode}, }; use jsonwebtoken::{ decode, errors::Error as JwtError, Algorithm, DecodingKey, TokenData, Validation, }; use serde::{Deserialize, Serialize}; pub struct ExtractAuthorization { pub id: usize, } #[derive(Serialize, Deserialize, Debug, PartialEq)] pub struct Claims { pub id: usize, pub exp: usize, } // 导入全局的AppState use crate::AppState; #[async_trait] impl FromRequestParts<AppState> for ExtractAuthorization { type Rejection = (StatusCode, &'static str); async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result<Self, Self::Rejection> { println!("{:?}", state); // 从State中获取密钥,替换硬编码值 let secret = &state.secret; let auth_header = parts.headers.get(AUTHORIZATION).ok_or(( StatusCode::BAD_REQUEST, "Authorization header missing", ))?; let auth_token = auth_header.to_str().map_err(|_| ( StatusCode::BAD_REQUEST, "Invalid characters in Authorization header", ))?; if auth_token.is_empty() { return Err(( StatusCode::BAD_REQUEST, "Authentication token cannot be empty", )); } let token_result: Result<TokenData<Claims>, JwtError> = decode::<Claims>( auth_token, &DecodingKey::from_secret(secret.as_bytes()), &Validation::new(Algorithm::HS256), ); match token_result { Ok(token) => Ok(ExtractAuthorization { id: token.claims.id, }), Err(_) => Err((StatusCode::UNAUTHORIZED, "Invalid or expired token")), } } }
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"
关键修改说明
- 在
authentication_token.rs中导入crate::AppState,并将FromRequestParts的泛型参数指定为AppState,这样就能直接访问state.secret。 - 移除了提取器中的硬编码密钥,改为从传入的
state中读取。 - 统一了
Claims的定义,避免重复代码,同时让类型保持一致。 - 优化了错误信息和对应状态码,使其更符合HTTP语义。
内容的提问来源于stack exchange,提问作者baby195lxl
相关产品推荐
相关产品推荐

