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

如何在Axum自定义Extractor中获取AppState实现JWT认证?

问题

使用Axum自定义提取器实现JWT认证时,能在提取器中打印State内容,但无法正确访问其中的密钥,当前JWT验证使用硬编码密钥,需修改代码从State中获取密钥完成验证。

解决方案

需要调整自定义提取器的泛型约束,使其能正确访问AppState中的密钥,同时统一Claims定义避免重复:

修改步骤

  1. 统一Claims定义:将Claims移到authentication_token.rs中,在main.rs中引用该定义,避免重复。
  2. 调整提取器的泛型约束:让ExtractAuthorization的FromRequestParts实现针对AppState,直接从传入的state中获取密钥。
  3. 替换硬编码密钥:将提取器中的硬编码密钥改为从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 15:45:20