Skip to main content

sz_rust_core/runtime/
websocket.rs

1//! sz-orm-websocket 服务端接入
2//!
3//! ## PHP 对齐
4//!
5//! 对齐 PHP `think-worker` 的 WebSocket 服务端模型:
6//!
7//! ```php
8//! $worker = new Workerman\Worker("websocket://0.0.0.0:2346");
9//! $worker->onConnect = function($connection) { /* ... */ };
10//! $worker->onMessage = function($connection, $data) { /* ... */ };
11//! $worker->onClose = function($connection) { /* ... */ };
12//! $worker->count = 4;  // worker 进程数
13//! Worker::runAll();
14//! ```
15//!
16//! Rust 端复用 `sz_orm_websocket::WsServer`(基于 tokio-tungstenite)。
17//!
18//! ## 设计
19//!
20//! - `WebSocketRuntime`:封装 WsServer,提供 start/stop lifecycle
21//! - 默认使用 `DefaultWebSocketHandler`(echo 模式)
22//! - 监听 `CancellationToken` 优雅停止
23
24use std::sync::Arc;
25
26use tokio_util::sync::CancellationToken;
27
28use sz_orm_websocket::{DefaultWebSocketHandler, WebSocketHandler, WsError, WsServer};
29
30/// WebSocket 运行时配置
31#[derive(Debug, Clone)]
32pub struct WebSocketRuntimeConfig {
33    /// 监听地址(如 "0.0.0.0:2346")
34    pub listen_addr: String,
35}
36
37impl Default for WebSocketRuntimeConfig {
38    fn default() -> Self {
39        Self {
40            listen_addr: "0.0.0.0:2346".to_string(),
41        }
42    }
43}
44
45impl WebSocketRuntimeConfig {
46    /// 创建新配置
47    pub fn new(listen_addr: impl Into<String>) -> Self {
48        Self {
49            listen_addr: listen_addr.into(),
50        }
51    }
52}
53
54/// WebSocket 运行时
55///
56/// 封装 `sz_orm_websocket::WsServer`,提供服务端 lifecycle 管理。
57///
58/// ## 设计
59///
60/// - 默认使用 `DefaultWebSocketHandler`(echo 模式),可替换为自定义 handler
61/// - `start` 方法 spawn 后台任务,返回 `JoinHandle<()>`
62/// - `stop` 方法调用 `WsServer::stop()`,触发 oneshot shutdown
63/// - 监听 `CancellationToken`:收到信号后自动调用 `WsServer::stop()`
64///
65/// ## 用法
66///
67/// ```rust,ignore
68/// use sz_rust_core::runtime::websocket::{WebSocketRuntime, WebSocketRuntimeConfig};
69/// use tokio_util::sync::CancellationToken;
70///
71/// let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("0.0.0.0:2346"));
72/// let token = CancellationToken::new();
73/// let handle = runtime.start(token.clone());
74/// // ... 业务运行 ...
75/// token.cancel();
76/// let _ = handle.await;
77/// ```
78pub struct WebSocketRuntime {
79    config: WebSocketRuntimeConfig,
80    server: Arc<WsServer>,
81}
82
83impl WebSocketRuntime {
84    /// 创建 WebSocket 运行时,使用默认 handler
85    pub fn new(config: WebSocketRuntimeConfig) -> Self {
86        let server = Arc::new(WsServer::new(&config.listen_addr));
87        Self { config, server }
88    }
89
90    /// 启动 WebSocket 服务(返回 JoinHandle,调用方持有)
91    ///
92    /// - 使用 `DefaultWebSocketHandler` 处理连接
93    /// - 监听 `token.cancelled()`,收到信号后调用 `server.stop()`
94    pub fn start(&self, token: CancellationToken) -> tokio::task::JoinHandle<Result<(), WsError>> {
95        let server = self.server.clone();
96        let handler: Arc<dyn WebSocketHandler> = Arc::new(DefaultWebSocketHandler::new());
97
98        tokio::spawn(async move {
99            // 启动 server(在子任务中运行,避免阻塞)
100            let server_clone = server.clone();
101            let mut start_task = tokio::spawn(async move { server_clone.start(handler).await });
102
103            // 监听 cancel,select! 通过 &mut 引用避免 move
104            tokio::select! {
105                _ = token.cancelled() => {
106                    let _ = server.stop().await;
107                    // start_task 尚未被 move,可以 await
108                    let _ = (&mut start_task).await;
109                    Ok(())
110                }
111                result = &mut start_task => {
112                    match result {
113                        Ok(inner) => inner,
114                        Err(e) => Err(WsError::Connection(format!("start task panicked: {}", e))),
115                    }
116                }
117            }
118        })
119    }
120
121    /// 使用自定义 handler 启动 WebSocket 服务
122    pub fn start_with_handler(
123        &self,
124        handler: Arc<dyn WebSocketHandler>,
125        token: CancellationToken,
126    ) -> tokio::task::JoinHandle<Result<(), WsError>> {
127        let server = self.server.clone();
128
129        tokio::spawn(async move {
130            let server_clone = server.clone();
131            let mut start_task = tokio::spawn(async move { server_clone.start(handler).await });
132
133            tokio::select! {
134                _ = token.cancelled() => {
135                    let _ = server.stop().await;
136                    let _ = (&mut start_task).await;
137                    Ok(())
138                }
139                result = &mut start_task => {
140                    match result {
141                        Ok(inner) => inner,
142                        Err(e) => Err(WsError::Connection(format!("start task panicked: {}", e))),
143                    }
144                }
145            }
146        })
147    }
148
149    /// 手动停止服务
150    pub async fn stop(&self) -> Result<(), WsError> {
151        self.server.stop().await
152    }
153
154    /// 获取当前连接数
155    pub async fn connection_count(&self) -> usize {
156        self.server.connection_count().await
157    }
158
159    /// 广播消息到所有连接
160    pub async fn broadcast_to_all(&self, data: Vec<u8>) -> Result<usize, WsError> {
161        self.server.broadcast_to_all(data).await
162    }
163
164    /// 是否仍在运行
165    pub async fn is_running(&self) -> bool {
166        self.server.is_running().await
167    }
168
169    /// 获取配置
170    pub fn config(&self) -> &WebSocketRuntimeConfig {
171        &self.config
172    }
173}
174
175#[cfg(test)]
176mod tests {
177    use super::*;
178    use std::time::Duration;
179
180    #[test]
181    fn test_websocket_runtime_config_default() {
182        let config = WebSocketRuntimeConfig::default();
183        assert_eq!(config.listen_addr, "0.0.0.0:2346");
184    }
185
186    #[test]
187    fn test_websocket_runtime_config_custom() {
188        let config = WebSocketRuntimeConfig::new("127.0.0.1:8080");
189        assert_eq!(config.listen_addr, "127.0.0.1:8080");
190    }
191
192    #[tokio::test]
193    async fn test_websocket_runtime_creation() {
194        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
195        assert_eq!(runtime.config().listen_addr, "127.0.0.1:0");
196        assert!(!runtime.is_running().await);
197        assert_eq!(runtime.connection_count().await, 0);
198    }
199
200    #[tokio::test]
201    async fn test_websocket_start_and_cancel() {
202        // 使用 port 0 让 OS 分配可用端口
203        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
204        let token = CancellationToken::new();
205        let handle = runtime.start(token.clone());
206
207        // 给 server 一点时间启动
208        tokio::time::sleep(Duration::from_millis(50)).await;
209
210        // 触发关闭
211        token.cancel();
212
213        // 等待任务退出
214        let result = tokio::time::timeout(Duration::from_secs(2), handle).await;
215        assert!(result.is_ok(), "websocket task should stop on cancel");
216    }
217
218    #[tokio::test]
219    async fn test_websocket_start_with_handler_and_cancel() {
220        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
221        let handler: Arc<dyn WebSocketHandler> = Arc::new(DefaultWebSocketHandler::new());
222        let token = CancellationToken::new();
223        let handle = runtime.start_with_handler(handler, token.clone());
224
225        tokio::time::sleep(Duration::from_millis(50)).await;
226        token.cancel();
227
228        let result = tokio::time::timeout(Duration::from_secs(2), handle).await;
229        assert!(result.is_ok(), "websocket task should stop on cancel");
230    }
231
232    #[tokio::test]
233    async fn test_websocket_broadcast_no_connections() {
234        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
235        // 没有连接时广播应返回 0
236        let result = runtime.broadcast_to_all(b"hello".to_vec()).await;
237        // 可能返回 Ok(0) 或 Err,取决于 WsServer 实现
238        let _ = result;
239    }
240
241    #[tokio::test]
242    async fn test_websocket_manual_stop() {
243        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
244        let token = CancellationToken::new();
245        let handle = runtime.start(token.clone());
246
247        tokio::time::sleep(Duration::from_millis(50)).await;
248
249        // 手动停止
250        let _ = runtime.stop().await;
251
252        // 等待任务退出
253        let _ = tokio::time::timeout(Duration::from_secs(2), handle).await;
254    }
255
256    #[test]
257    fn test_config_accessor() {
258        let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("0.0.0.0:9999"));
259        assert_eq!(runtime.config().listen_addr, "0.0.0.0:9999");
260    }
261
262    #[tokio::test]
263    async fn test_multiple_websocket_runtimes() {
264        // 验证可以创建多个 runtime 实例(不启动)
265        let rt1 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
266        let rt2 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
267
268        assert_eq!(rt1.connection_count().await, 0);
269        assert_eq!(rt2.connection_count().await, 0);
270    }
271}