5.3 Axum 中间件实战:认证、日志、限流,构建企业级中间件栈
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) 中,中间件可以通过多种方式实现,从简单到复杂,主要有以下几种:
- 使用
tower-http的现成Layer:最简单、最快捷的方式,适用于标准化的需求(如日志、压缩、CORS)。 - 使用
from_fn创建简单的中间件:对于只需要在请求到达 handler 前进行简单处理的场景。 - 创建自定义
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_routehandler。 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
})
}
}
关键点分析:
- 状态共享: 我们把
DashMap和配置放在RateLimitState中,并用Arc包裹,使其可以在多个RateLimitService实例(对应多个连接)之间共享。 - 提取 IP: 我们通过请求的
extensions来获取ConnectInfo,从而得到客户端 IP。这需要我们在Router上添加一个tower_http::ServiceBuilderExt::into_inner()。 - 滑动窗口算法: 我们用一个
VecDeque作为滑动窗口来记录每个 IP 在时间窗口内的请求时间戳。 - 提前返回: 当请求被限流时,我们不再调用
self.inner.call(req),而是自己构建一个429响应,并将其包装在一个立即完成的Future中返回。 - 异步
FutureBoxing: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 生态的无缝支持使其在中间件方面拥有无与伦比的能力。
tower-http是你的好朋友:在自己编写中间件之前,先去tower-http的文档里找一找,绝大多数常见的需求(日志、压缩、CORS、超时等)都有现成的、高质量的实现。middleware::from_fn用于简单场景:当你的中间件是无状态的,并且只需要在请求到达 handler 前做一些检查或修改时,from_fn是最快捷、最易读的方式。它还支持提取器,非常方便。- 自定义
Layer和Service用于复杂逻辑:当你需要一个有状态的中间件(如限流器),或者需要对响应进行修改时,编写自定义的Layer提供了最大的灵活性。这是最高级的模式,需要你对tower的Servicetrait 和异步Future有更深的理解。 Layer的组合与顺序:axum允许你像堆叠乐高一样组合中间件。请求从外到内,响应从内到外,理解这个顺序对于正确配置中间件栈至关重要。
通过实战这三种不同层次的中间件,你现在已经具备了为你的 axum 应用构建企业级功能的能力。无论是简单的日志记录,还是复杂的自定义认证和限流,你都有合适的工具来优雅地解决问题。
思考题
- 在我们的
RateLimitService::call实现中,我们返回了Pin<Box<dyn Future + ...>>。为什么我们不能直接return self.inner.call(req)或者return async { Ok(response) }?这和 Rust 的类型系统有什么关系? axum::middleware::from_fn创建的中间件,其函数签名async fn<B>(req: Request<B>, next: Next<B>)中的next: Next<B>是什么?它和tower::Service的inner服务有什么关系?- 如果你想让一个中间件只对一部分路由生效,除了使用
.route_layer(),还有哪些方法可以实现?(提示:Router::nest和Router::merge也可以和.layer()结合使用)。 tower::Service::poll_ready的作用是实现背压。请设想一个场景,在我们的RateLimitLayer中,我们可以利用poll_ready做些什么来更主动地管理负载?- 一个中间件既可以修改请求,也可以修改响应。请分别简述这两种操作在自定义
Layer的call方法中应该如何实现。
实践练习
- 实现一个“维护模式”中间件:创建一个
MaintenanceLayer。当应用处于维护模式时(可以通过一个全局的Arc<AtomicBool>状态来控制),所有非/health的请求都应该被中断,并返回一个503 Service Unavailable响应。 - 为
auth_middleware添加角色检查:扩展本章的from_fn认证中间件。假设你的 token 中包含了用户的角色信息。- 在请求的
extensions中插入一个CurrentUser { id: u64, role: String }结构体。 - 创建一个新的 handler
admin_route,它需要一个新的提取器CurrentUserExtractor来从extensions中获取用户信息,并检查role是否为"admin"。如果不是,提取器应该返回错误。
- 在请求的
- 探索
tower-http的RequestBodyLimitLayer: 阅读文档,并使用RequestBodyLimitLayer来限制你的create_user接口的请求体大小不能超过 2KB。
更多推荐
所有评论(0)