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 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 let _ = runtime.stop().await;
251
252 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 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}