sz_rust_core/runtime/
websocket.rs1use std::sync::Arc;
25
26use tokio_util::sync::CancellationToken;
27
28use crate::orm::{DefaultWebSocketHandler, WebSocketHandler, WsError, WsServer};
29
30#[derive(Debug, Clone)]
32pub struct WebSocketRuntimeConfig {
33 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 pub fn new(listen_addr: impl Into<String>) -> Self {
48 Self {
49 listen_addr: listen_addr.into(),
50 }
51 }
52}
53
54pub struct WebSocketRuntime {
79 config: WebSocketRuntimeConfig,
80 server: Arc<WsServer>,
81}
82
83impl WebSocketRuntime {
84 pub fn new(config: WebSocketRuntimeConfig) -> Self {
86 let server = Arc::new(WsServer::new(&config.listen_addr));
87 Self { config, server }
88 }
89
90 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 let server_clone = server.clone();
101 let mut start_task = tokio::spawn(async move { server_clone.start(handler).await });
102
103 tokio::select! {
105 _ = token.cancelled() => {
106 let _ = server.stop().await;
107 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 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 pub async fn stop(&self) -> Result<(), WsError> {
151 self.server.stop().await
152 }
153
154 pub async fn connection_count(&self) -> usize {
156 self.server.connection_count().await
157 }
158
159 pub async fn broadcast_to_all(&self, data: Vec<u8>) -> Result<usize, WsError> {
161 self.server.broadcast_to_all(data).await
162 }
163
164 pub async fn is_running(&self) -> bool {
166 self.server.is_running().await
167 }
168
169 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 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 tokio::time::sleep(Duration::from_millis(50)).await;
209
210 token.cancel();
212
213 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 let result = runtime.broadcast_to_all(b"hello".to_vec()).await;
237 match result {
238 Ok(n) => assert_eq!(n, 0, "无连接时广播应送达 0 个接收者"),
239 Err(e) => panic!("无连接广播不应失败,实际: {:?}", e),
240 }
241 }
242
243 #[tokio::test]
244 async fn test_websocket_manual_stop() {
245 let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
246 let token = CancellationToken::new();
247 let handle = runtime.start(token.clone());
248
249 tokio::time::sleep(Duration::from_millis(50)).await;
250
251 let stop_result = runtime.stop().await;
253 assert!(
254 stop_result.is_ok(),
255 "stop() 应成功,实际: {:?}",
256 stop_result
257 );
258
259 let exit = tokio::time::timeout(Duration::from_secs(2), handle).await;
261 assert!(exit.is_ok(), "stop 后 websocket 任务应在 2s 内退出");
262 }
263
264 #[test]
265 fn test_config_accessor() {
266 let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("0.0.0.0:9999"));
267 assert_eq!(runtime.config().listen_addr, "0.0.0.0:9999");
268 }
269
270 #[tokio::test]
271 async fn test_multiple_websocket_runtimes() {
272 let rt1 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
274 let rt2 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
275
276 assert_eq!(rt1.connection_count().await, 0);
277 assert_eq!(rt2.connection_count().await, 0);
278 }
279}