5.3 Axum 中间件实战:认证、日志、限流,构建企业级中间件栈

引言:横切关注点与中间件的力量

在构建 Web 应用时,除了核心的业务逻辑(例如,创建用户、查询商品),我们还需要处理大量与业务逻辑无关,但对整个应用至关重要的“横切关注点”(Cross-Cutting Concerns)。这些关注点包括:

  • 日志记录 (Logging):记录每个请求的详细信息。
  • 认证 (Authentication):验证用户身份,保护路由。
  • 授权 (Authorization):检查用户是否有权限访问特定资源。
  • 超时控制 (Timeout):防止单个请求耗时过长,拖垮整个服务。
  • 请求限流 (Rate Limiting):防止恶意攻击或服务滥用。
  • CORS 处理 (Cross-Origin Resource Sharing):允许或拒绝跨域请求。
  • 压缩 (Compression):为响应体启用 Gzip/Brotli 压缩。

如果将这些逻辑与业务逻辑混杂在每个 handler 中,代码将变得臃肿、重复且难以维护。中间件 (Middleware) 就是为了解决这个问题而生的。它是一种在请求被 handler 处理之前和响应被发送给客户端之后,对请求和响应进行处理的机制。

正如我们在 5.1 节中学到的,axum 完全拥抱 tower 的 Service/Layer 模型。在 axum 中,中间件就是 tower::Layer。这使得我们可以利用 tower 和 tower-http 生态中海量的、经过生产环境考验的中间件,也可以轻松地编写自己的中间件。

本章,我们将通过实战来构建一个企业级的中间件栈,涵盖日志、认证、限流等核心功能。

中间件的类型

在 axum (以及 tower) 中,中间件可以通过多种方式实现,从简单到复杂,主要有以下几种:

  1. 使用 tower-http 的现成 Layer:最简单、最快捷的方式,适用于标准化的需求(如日志、压缩、CORS)。
  2. 使用 from_fn 创建简单的中间件:对于只需要在请求到达 handler 前进行简单处理的场景。
  3. 创建自定义 Layer:最强大、最灵活的方式,可以完全控制请求到响应的整个生命周期,适用于复杂的逻辑(如自定义认证)。

实战 1:使用 tower-http 快速搭建基础中间件栈

我们先从最简单的开始,利用 tower-http 快速为我们的应用添加日志和 CORS 支持。

环境准备:

[dependencies]
axum = "0.6"
tokio = { version = "1", features = ["full"] }
tower = "0.4"
tower-http = { version = "0.4", features = ["trace", "cors"] }
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }

构建中间件栈:

use axum::{routing::get, Router};
use tower_http::{cors::{CorsLayer, Any}, trace::TraceLayer};
use std::net::SocketAddr;

#[tokio::main]
async fn main() {
    // 初始化 tracing
    tracing_subscriber::fmt()
        .with_env_filter(tracing_subscriber::EnvFilter::from_default_env())
        .init();

    // 创建 CORS layer
    let cors_layer = CorsLayer::new()
        .allow_origin(Any) // 允许所有来源
        .allow_methods(Any) // 允许所有 HTTP 方法
        .allow_headers(Any); // 允许所有 HTTP 头

    // 创建应用
    let app = Router::new()
        .route("/", get(handler))
        // 使用 .layer() 应用中间件
        // 执行顺序:从下到上(请求),从上到下(响应)
        .layer(cors_layer)       // 2. CORS layer
        .layer(TraceLayer::new_for_http()); // 1. Trace layer

    let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
    tracing::info!("listening on {}", addr);
    axum::Server::bind(&addr)
        .serve(app.into_make_service())
        .await
        .unwrap();
}

async fn handler() -> &'static str {
    "Hello, Middleware!"
}

分析:

  • TraceLayer: 这是 tower-http 提供的一个功能非常强大的日志中间件。它会在收到请求时打印一条日志,在发送响应时再打印一条日志,并包含 status, latency(延迟)等信息。
  • CorsLayer: 轻松配置 CORS 策略。.allow_origin(Any) 为了方便演示,在生产环境中你应该指定允许的来源。
  • 执行顺序: 当一个请求到来时,它会先经过 TraceLayer,再经过 CorsLayer,最后到达我们的 handler。当 handler 生成响应后,响应会先经过 CorsLayer(在这里它可能会添加 Access-Control-Allow-Origin 等头部),然后再经过 TraceLayer(在这里它会记录响应状态和延迟)。

运行 RUST_LOG=tower_http=debug cargo run 并用 curl -v http://127.0.0.1:3000 请求,你将看到详细的日志输出和 CORS 相关的响应头。

实战 2:使用 from_fn 创建简单的认证中间件

有时候,我们只需要一个简单的函数作为中间件,例如检查一个 Authorization 头。axum 提供了 axum::middleware::from_fn 来将一个异步函数转换成一个中间件 Layer。

这个函数需要有特定的签名:

async fn my_middleware<B>(
    request: Request<B>,
    next: Next<B>
) -> Result<Response, AppError>
  • request: 传入的请求。
  • next: 代表“下一个”服务(可以是另一个中间件,也可以是最终的 handler)。你必须调用 next.run(request).await 来继续处理流程。
  • 返回值必须是一个 Response。

编写认证中间件:

use axum::{
    async_trait,
    extract::{FromRequestParts, TypedHeader},
    headers::{authorization::Bearer, Authorization},
    http::{Request, StatusCode},
    middleware::{self, Next},
    response::{IntoResponse, Response},
    routing::get,
    Router,
};

// 认证中间件函数
async fn auth_middleware<B>(
    // 我们可以像在 handler 中一样使用提取器!
    TypedHeader(auth_header): TypedHeader<Authorization<Bearer>>,
    // 注意:这里的 request 必须放在 Next 之前
    request: Request<B>,
    next: Next<B>,
) -> Result<Response, StatusCode> {
    let token = auth_header.token();
    
    // 模拟检查 token
    if token == "secret-token-123" {
        // Token 有效,继续处理请求
        println!("认证成功!");
        let response = next.run(request).await;
        Ok(response)
    } else {
        // Token 无效,直接返回 401 Unauthorized
        println!("认证失败!");
        Err(StatusCode::UNAUTHORIZED)
    }
}

// 定义一个 handler
async fn protected_route() -> &'static str {
    "Welcome to the protected area!"
}

#[tokio::main]
async fn main() {
    let app = Router::new()
        .route("/protected", get(protected_route))
        // 使用 from_fn 将函数转为中间件,并用 route_layer 应用到单个路由
        .route_layer(middleware::from_fn(auth_middleware));

    // ... 启动服务器 ...
}

分析:

  • 在中间件中使用提取器: from_fn 创建的中间件函数签名非常灵活,它可以像 handler 一样使用提取器。这里我们直接提取了 Authorization<Bearer> 头部。如果请求中没有这个头部或格式不正确,axum 会自动返回 400 Bad Request,我们的中间件函数甚至不会被调用。
  • next.run(request).await: 这是中间件的核心。调用它会将请求传递给内层的服务(可能是另一个中间件或最终的 handler)。如果你不调用它,请求处理链就会在此中断。
  • 提前返回: 在认证失败的分支中,我们直接 return Err(StatusCode::UNAUTHORIZED),它会被转换成一个 401 响应,请求不会到达 protected_route handler。
  • route_layer: 这个中间件只对 /protected 路由生效。

用 curl 测试:

# 失败:没有 token
curl -v http://127.0.0.1:3000/protected
# > HTTP/1.1 401 Unauthorized

# 成功:提供了正确的 token
curl -H "Authorization: Bearer secret-token-123" http://127.0.0.1:3000/protected
# Welcome to the protected area!

实战 3:创建自定义 Layer 实现请求限流

from_fn 很方便,但它的能力有限。例如,如果你的中间件需要维护自己的状态(比如限流器的计数),或者需要修改响应,那么创建完整的自定义 Layer 是更健壮的选择。

让我们来实现一个简单的、基于 IP 的内存请求限流器。

环境准备:
我们需要一个并发的哈希表来存储每个 IP 的访问记录。

[dependencies]
# ...
dashmap = "5.4"
# ...

1. 定义 Layer 和 Service

use tower::{Layer, Service};
use std::sync::Arc;
use std::collections::VecDeque;
use std::time::{Instant, Duration};
use dashmap::DashMap;
use axum::http::{Request, StatusCode};
use axum::response::{IntoResponse, Response};
use std::net::SocketAddr;
use axum::extract::ConnectInfo;


// 我们的 Layer
#[derive(Clone)]
pub struct RateLimitLayer {
    state: Arc<RateLimitState>,
}

struct RateLimitState {
    // 使用 DashMap 来实现线程安全的 IP 访问记录
    records: DashMap<String, VecDeque<Instant>>,
    // 允许在 `period` 时间内发生 `max_requests` 次请求
    max_requests: u64,
    period: Duration,
}

impl RateLimitLayer {
    pub fn new(max_requests: u64, period: Duration) -> Self {
        Self {
            state: Arc::new(RateLimitState {
                records: DashMap::new(),
                max_requests,
                period,
            }),
        }
    }
}

// 实现 Layer trait
impl<S> Layer<S> for RateLimitLayer {
    type Service = RateLimitService<S>;

    fn layer(&self, inner: S) -> Self::Service {
        RateLimitService {
            inner,
            state: self.state.clone(),
        }
    }
}

// 我们包装后的 Service
#[derive(Clone)]
pub struct RateLimitService<S> {
    inner: S,
    state: Arc<RateLimitState>,
}

2. 为 RateLimitService 实现 Service Trait

use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};


impl<S, ReqBody> Service<Request<ReqBody>> for RateLimitService<S>
where
    // 内部服务 S 必须是处理 Request<ReqBody> 的 Service
    S: Service<Request<ReqBody>, Response = Response> + Send + 'static,
    S::Future: Send + 'static,
{
    type Response = S::Response;
    type Error = S::Error;
    // 我们需要返回一个自定义的 Future,因为我们有可能会提前返回 429 响应
    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;

    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        // 直接委托给内部服务
        self.inner.poll_ready(cx)
    }

    fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
        // 从请求中提取 IP 地址
        // ConnectInfo 必须作为 app 的 layer 添加才能被提取
        let ip = req.extensions()
            .get::<ConnectInfo<SocketAddr>>()
            .map(|ci| ci.0.ip().to_string())
            .unwrap_or_else(|| "unknown".to_string());
            
        let state = self.state.clone();
        let now = Instant::now();
        
        let mut records = state.records.entry(ip).or_default();
        
        // 移除时间窗口之外的旧记录
        while let Some(front) = records.front() {
            if now.duration_since(*front) > state.period {
                records.pop_front();
            } else {
                break;
            }
        }
        
        // 检查请求次数是否超限
        if (records.len() as u64) >= state.max_requests {
            println!("IP {} 被限流", records.key());
            // 如果超限,立即返回一个 429 Too Many Requests 响应
            let response = (
                StatusCode::TOO_MANY_REQUESTS,
                "Too many requests",
            ).into_response();
            
            // Box::pin 将其包装成一个立即完成的 Future
            return Box::pin(async { Ok(response) });
        }
        
        // 未超限,记录本次请求时间
        records.push_back(now);
        
        // 必须释放 DashMap 的 `RefMut` (通过 drop),否则会死锁!
        drop(records);

        // 调用内部服务
        let future = self.inner.call(req);
        
        // 将内部服务的 Future 返回
        Box::pin(async move {
            future.await
        })
    }
}

关键点分析:

  1. 状态共享: 我们把 DashMap 和配置放在 RateLimitState 中,并用 Arc 包裹,使其可以在多个 RateLimitService 实例(对应多个连接)之间共享。
  2. 提取 IP: 我们通过请求的 extensions 来获取 ConnectInfo,从而得到客户端 IP。这需要我们在 Router 上添加一个 tower_http::ServiceBuilderExt::into_inner()。
  3. 滑动窗口算法: 我们用一个 VecDeque 作为滑动窗口来记录每个 IP 在时间窗口内的请求时间戳。
  4. 提前返回: 当请求被限流时,我们不再调用 self.inner.call(req),而是自己构建一个 429 响应,并将其包装在一个立即完成的 Future 中返回。
  5. 异步 Future Boxing: call 方法的返回值类型是 Pin<Box<dyn Future + ...>>。这是一种类型擦除技术,因为我们的 call 方法有两个可能的返回路径:一个是立即返回的我们自己创建的 Future(限流时),另一个是内部服务返回的 Future。它们的类型不同,所以我们需要将它们都“擦除”为通用的 dyn Future 类型。Box::pin 是完成这个工作的标准方式。

3. 在应用中使用

// main.rs
// ...
// axum::extract::ConnectInfo 需要这个
use axum::extract::connect_info::IntoMakeServiceWithConnectInfo;
use std::net::SocketAddr;

#[tokio::main]
async fn main() {
    let rate_limit_layer = RateLimitLayer::new(5, Duration::from_secs(10)); // 10秒内最多5次请求

    let app = Router::new()
        .route("/", get(|| async { "Hello!" }))
        .layer(rate_limit_layer);
        
    let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
    let make_service = app.into_make_service_with_connect_info::<SocketAddr>();

    axum::Server::bind(&addr)
        .serve(make_service)
        .await
        .unwrap();
}

现在,如果你在 10 秒内连续访问 http://127.0.0.1:3000 超过 5 次,你就会收到 429 Too Many Requests 响应。

总结

中间件是构建健壮、可维护的 Web 应用的基石,而 axum 对 tower 生态的无缝支持使其在中间件方面拥有无与伦比的能力。

  1. tower-http 是你的好朋友:在自己编写中间件之前,先去 tower-http 的文档里找一找,绝大多数常见的需求(日志、压缩、CORS、超时等)都有现成的、高质量的实现。
  2. middleware::from_fn 用于简单场景:当你的中间件是无状态的,并且只需要在请求到达 handler 前做一些检查或修改时,from_fn 是最快捷、最易读的方式。它还支持提取器,非常方便。
  3. 自定义 Layer 和 Service 用于复杂逻辑:当你需要一个有状态的中间件(如限流器),或者需要对响应进行修改时,编写自定义的 Layer 提供了最大的灵活性。这是最高级的模式,需要你对 tower 的 Service trait 和异步 Future 有更深的理解。
  4. Layer 的组合与顺序:axum 允许你像堆叠乐高一样组合中间件。请求从外到内,响应从内到外,理解这个顺序对于正确配置中间件栈至关重要。

通过实战这三种不同层次的中间件,你现在已经具备了为你的 axum 应用构建企业级功能的能力。无论是简单的日志记录,还是复杂的自定义认证和限流,你都有合适的工具来优雅地解决问题。

思考题

  1. 在我们的 RateLimitService::call 实现中,我们返回了 Pin<Box<dyn Future + ...>>。为什么我们不能直接 return self.inner.call(req) 或者 return async { Ok(response) }?这和 Rust 的类型系统有什么关系?
  2. axum::middleware::from_fn 创建的中间件,其函数签名 async fn<B>(req: Request<B>, next: Next<B>) 中的 next: Next<B> 是什么?它和 tower::Service 的 inner 服务有什么关系?
  3. 如果你想让一个中间件只对一部分路由生效,除了使用 .route_layer(),还有哪些方法可以实现?(提示:Router::nest 和 Router::merge 也可以和 .layer() 结合使用)。
  4. tower::Service::poll_ready 的作用是实现背压。请设想一个场景,在我们的 RateLimitLayer 中,我们可以利用 poll_ready 做些什么来更主动地管理负载?
  5. 一个中间件既可以修改请求,也可以修改响应。请分别简述这两种操作在自定义 Layer 的 call 方法中应该如何实现。

实践练习

  1. 实现一个“维护模式”中间件:创建一个 MaintenanceLayer。当应用处于维护模式时(可以通过一个全局的 Arc<AtomicBool> 状态来控制),所有非 /health 的请求都应该被中断,并返回一个 503 Service Unavailable 响应。
  2. 为 auth_middleware 添加角色检查:扩展本章的 from_fn 认证中间件。假设你的 token 中包含了用户的角色信息。
    • 在请求的 extensions 中插入一个 CurrentUser { id: u64, role: String } 结构体。
    • 创建一个新的 handler admin_route,它需要一个新的提取器 CurrentUserExtractor 来从 extensions 中获取用户信息,并检查 role 是否为 "admin"。如果不是,提取器应该返回错误。
  3. 探索 tower-http 的 RequestBodyLimitLayer: 阅读文档,并使用 RequestBodyLimitLayer 来限制你的 create_user 接口的请求体大小不能超过 2KB。
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐