Skip to main content

sz_rust_core/
websocket_route.rs

1//! WebSocket 原生路由 — 基于 axum WebSocketUpgrade
2//!
3//! 将 WebSocket 处理集成到主 HTTP 路由中,无需独立端口。
4//!
5//! ## PHP 对齐
6//!
7//! 对齐 PHP `think-worker` WebSocket + Workerman 的 `websocket://` 协议:
8//!
9//! ```php
10//! $worker = new Worker("websocket://0.0.0.0:2346");
11//! $worker->onMessage = function($connection, $data) {
12//!     $connection->send("echo: $data");
13//! };
14//! ```
15//!
16//! Rust 端通过 axum 的 `WebSocketUpgrade` 提取器,在主 HTTP 端口上
17//! 处理 WebSocket 升级请求(如 `GET /ws/chat`),无需独立端口。
18//!
19//! ## 设计
20//!
21//! - [`WsHandler`] trait:简化 WebSocket 事件处理(on_connect/on_message/on_close)
22//! - [`ws_handler()`]:将 `WsHandler` 转为 axum handler
23//! - [`EchoWsHandler`]:默认回显处理器
24//! - [`crate::router::RouterBuilder::ws`]:注册 WebSocket 路由
25//!
26//! ## 用法
27//!
28//! ```ignore
29//! use sz_rust_core::router::RouterBuilder;
30//! use sz_rust_core::websocket_route::{ws_handler, EchoWsHandler};
31//!
32//! let router = RouterBuilder::new()
33//!     .ws("/ws/echo", ws_handler(EchoWsHandler::new()))
34//!     .build();
35//! ```
36
37use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
38use std::sync::Arc;
39
40// ============================================================================
41// WsHandler trait — 简化 WebSocket 事件处理
42// ============================================================================
43
44/// WebSocket 事件处理器 trait
45///
46/// 对齐 PHP Workerman 的 `onConnect` / `onMessage` / `onClose` 事件模型。
47/// 用户实现此 trait,通过 [`ws_handler()`] 转为 axum handler。
48pub trait WsHandler: Send + Sync + 'static {
49    /// 连接建立时调用
50    ///
51    /// 默认空实现。可用于注册连接、发送欢迎消息等。
52    fn on_connect(&self) {}
53
54    /// 收到文本/二进制消息时调用
55    ///
56    /// 返回 `Some(Message)` 则自动回发给客户端,返回 `None` 则不回发。
57    /// 默认返回 `None`。
58    fn on_message(&self, _msg: Message) -> Option<Message> {
59        None
60    }
61
62    /// 连接关闭时调用
63    ///
64    /// 默认空实现。可用于清理资源、广播离线通知等。
65    fn on_close(&self) {}
66}
67
68/// 将 [`WsHandler`] 转为 axum `MethodRouter`,可直接注册到路由
69///
70/// 返回 `MethodRouter` 而非裸 handler,避免 `impl Fn` 在 axum `Handler` trait
71/// 推导上的已知限制(`impl Trait` 返回类型无法被泛型 bound 正确解析)。
72///
73/// ## 用法
74///
75/// ```ignore
76/// use sz_rust_core::websocket_route::{ws_handler, EchoWsHandler};
77///
78/// let router = axum::Router::new()
79///     .route("/ws/echo", ws_handler(EchoWsHandler::new()));
80/// ```
81pub fn ws_handler<H: WsHandler>(handler: H) -> axum::routing::MethodRouter<()> {
82    let handler = Arc::new(handler);
83    axum::routing::get(move |ws: WebSocketUpgrade| async move {
84        let handler = handler.clone();
85        ws.on_upgrade(move |socket| handle_ws_connection(socket, handler))
86    })
87}
88
89/// 处理 WebSocket 连接生命周期
90///
91/// 依次调用 `on_connect` → 循环 `on_message` → `on_close`。
92async fn handle_ws_connection(mut socket: WebSocket, handler: Arc<dyn WsHandler>) {
93    handler.on_connect();
94
95    // 循环接收消息,调用 on_message,有回复则发送
96    while let Some(Ok(msg)) = socket.recv().await {
97        // 处理 Close 消息:退出循环
98        if matches!(msg, Message::Close(_)) {
99            break;
100        }
101
102        if let Some(reply) = handler.on_message(msg) {
103            // 发送回复,失败则退出
104            if socket.send(reply).await.is_err() {
105                break;
106            }
107        }
108    }
109
110    handler.on_close();
111}
112
113// ============================================================================
114// EchoWsHandler — 默认回显处理器
115// ============================================================================
116
117/// 回显 WebSocket 处理器
118///
119/// 将收到的消息原样回发给客户端。适用于心跳检测、调试等场景。
120#[derive(Debug, Default, Clone)]
121pub struct EchoWsHandler;
122
123impl EchoWsHandler {
124    /// 创建回显处理器
125    pub fn new() -> Self {
126        Self
127    }
128}
129
130impl WsHandler for EchoWsHandler {
131    fn on_message(&self, msg: Message) -> Option<Message> {
132        Some(msg)
133    }
134}
135
136// ============================================================================
137// 测试
138// ============================================================================
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143    use axum::body::Body;
144    use axum::http::{Method, Request, StatusCode};
145    use tower::ServiceExt;
146
147    /// 测试 EchoWsHandler 基本行为
148    #[test]
149    fn test_echo_handler_returns_message() {
150        let handler = EchoWsHandler::new();
151        let msg = Message::text("hello");
152        let result = handler.on_message(msg);
153        assert!(result.is_some());
154    }
155
156    /// 测试 EchoWsHandler 默认构造
157    #[test]
158    fn test_echo_handler_default() {
159        let handler = EchoWsHandler;
160        let msg = Message::text("test");
161        assert!(handler.on_message(msg).is_some());
162    }
163
164    /// 测试自定义 WsHandler
165    struct NoReplyHandler;
166    impl WsHandler for NoReplyHandler {
167        fn on_message(&self, _msg: Message) -> Option<Message> {
168            None
169        }
170    }
171
172    #[test]
173    fn test_custom_handler_no_reply() {
174        let handler = NoReplyHandler;
175        let msg = Message::text("hello");
176        assert!(handler.on_message(msg).is_none());
177    }
178
179    /// 测试自定义 WsHandler 有回复
180    struct PrefixHandler;
181    impl WsHandler for PrefixHandler {
182        fn on_message(&self, _msg: Message) -> Option<Message> {
183            Some(Message::text("prefix: reply"))
184        }
185    }
186
187    #[test]
188    fn test_custom_handler_with_reply() {
189        let handler = PrefixHandler;
190        let msg = Message::text("input");
191        let reply = handler.on_message(msg).unwrap();
192        assert_eq!(reply.to_text().unwrap(), "prefix: reply");
193    }
194
195    /// 测试 on_connect 和 on_close 默认实现不 panic
196    #[test]
197    fn test_default_lifecycle_hooks_no_panic() {
198        let handler = EchoWsHandler::new();
199        handler.on_connect();
200        handler.on_close();
201    }
202
203    /// 测试 WebSocket 路由注册(HTTP 层面验证 400 Bad Request)
204    ///
205    /// 非 WebSocket 请求(无 Upgrade 头)访问 WebSocket 路由时,
206    /// axum 0.8 返回 400 Bad Request(要求 `Upgrade: websocket` 头)。
207    #[tokio::test]
208    async fn test_ws_route_registered_as_get() {
209        let router = axum::Router::new().route("/ws/echo", ws_handler(EchoWsHandler::new()));
210
211        // 普通 GET 请求(无 Upgrade 头)应返回 400
212        let request = Request::builder()
213            .method(Method::GET)
214            .uri("/ws/echo")
215            .body(Body::empty())
216            .unwrap();
217        let response = router.oneshot(request).await.unwrap();
218        // axum 0.8 对非 WebSocket 请求返回 400 Bad Request
219        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
220    }
221
222    /// 测试 WebSocket 路由 404
223    #[tokio::test]
224    async fn test_ws_route_not_found() {
225        let router = axum::Router::new().route("/ws/echo", ws_handler(EchoWsHandler::new()));
226
227        let request = Request::builder()
228            .method(Method::GET)
229            .uri("/ws/nonexistent")
230            .body(Body::empty())
231            .unwrap();
232        let response = router.oneshot(request).await.unwrap();
233        assert_eq!(response.status(), StatusCode::NOT_FOUND);
234    }
235
236    /// 测试 POST 方法访问 WebSocket 路由返回 405
237    #[tokio::test]
238    async fn test_ws_route_rejects_post() {
239        let router = axum::Router::new().route("/ws/echo", ws_handler(EchoWsHandler::new()));
240
241        let request = Request::builder()
242            .method(Method::POST)
243            .uri("/ws/echo")
244            .body(Body::empty())
245            .unwrap();
246        let response = router.oneshot(request).await.unwrap();
247        assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
248    }
249}