Axum框架中WebSocket连接的Origin来源过滤配置咨询
你遇到的这个问题我之前也帮人排查过,确实Axum自带的CorsLayer是专门针对普通HTTP/HTTPS请求的CORS规则校验,WebSocket(WS/WSS)的连接握手虽然用的是GET请求,但它的Origin校验并不属于CORS规范的覆盖范围——CORS管的是跨域HTTP请求的权限,而WebSocket的Origin是浏览器自带的安全校验机制,所以你之前配置的CorsLayer对WebSocket连接完全起不到过滤作用,所有请求都能进来,这是正常的~
你现在在handler里用TypedHeader提取Origin然后校验的思路是没问题的,而且已经能实现需求,不过如果之后要加更多WebSocket路由,每个handler都写一遍校验确实麻烦,下面给你两种更灵活的方案:
方案一:完善handler内的校验(适合当前只有一个WebSocket路由的场景)
这种方式简单直接,适合你现在只有一个WebSocket路由的情况,把校验逻辑补全就行:
use axum::{ extract::{Path, TypedHeader}, response::{IntoResponse, StatusCode, WebSocketUpgrade}, }; use axum_extra::extract::TypedHeader; use headers::Origin; async fn client_handler( ws: WebSocketUpgrade, Path(session_id): Path<String>, TypedHeader(origin): TypedHeader<Origin>, ) -> Result<impl IntoResponse, StatusCode> { // 精准校验来源主机名 if origin.hostname() != "www.example.com" { return Err(StatusCode::FORBIDDEN); } // 校验通过,完成WebSocket升级,后续处理连接逻辑 Ok(ws.on_upgrade(|socket| async move { // 这里写你的WebSocket消息处理逻辑,比如 handle_socket(socket, session_id).await; })) }
缺点就是无法复用校验逻辑,之后新增WebSocket路由时要重复写这段校验代码。
方案二:封装全局Origin校验中间件(适合多WebSocket路由的场景)
如果之后可能会加更多WebSocket路由,把校验逻辑封装成自定义中间件挂载到路由层是更优雅的选择,所有经过的请求都会先过校验,不用每个handler重复写:
首先编写中间件代码:
use axum::{ extract::TypedHeader, http::StatusCode, middleware::Next, response::IntoResponse, RequestPartsExt, }; use headers::Origin; // 定义WebSocket Origin校验中间件 async fn ws_origin_middleware<B>( mut req: axum::Request<B>, next: Next<B>, ) -> Result<impl IntoResponse, StatusCode> { // 从请求中提取Origin头,如果没有则返回400(非浏览器发起的请求可能没有这个头) let TypedHeader(origin) = req.extract_parts::<TypedHeader<Origin>>() .await .map_err(|_| StatusCode::BAD_REQUEST)?; // 校验来源主机名是否符合要求 if origin.hostname() != "www.example.com" { return Err(StatusCode::FORBIDDEN); } // 校验通过,继续执行后续的handler逻辑 Ok(next.run(req).await.into_response()) }
然后在路由配置中挂载这个中间件:
use axum::{Router, routing::get}; use tower_http::cors::CorsLayer; use axum::http::{Method, header::CONTENT_TYPE}; use axum::http::HeaderValue; let app = Router::new() // 给WebSocket路由挂载自定义的Origin校验中间件 .route("/client/:session_id", get(client_handler)) .route_layer(axum::middleware::from_fn(ws_origin_middleware)) // 普通HTTP请求的CORS配置保持你原来的设置即可 .layer( CorsLayer::new() .allow_origin("https://www.example.com".parse::<HeaderValue>().unwrap()) .allow_methods([Method::GET, Method::POST]) .allow_headers([CONTENT_TYPE]), );
这样不管之后加多少WebSocket路由,只要把它们挂载到同一个路由组下,就能自动复用Origin校验逻辑,不用重复写代码。如果你的场景允许非浏览器发起的WebSocket请求(这类请求可能没有Origin头),可以把提取Origin的逻辑改成可选的,比如用Option<TypedHeader<Origin>>,然后根据业务需求判断是否放行。
内容来源于stack exchange

