use std::sync::Arc;
use tokio_util::sync::CancellationToken;
use crate::orm::{DefaultWebSocketHandler, WebSocketHandler, WsError, WsServer};
#[derive(Debug, Clone)]
pub struct WebSocketRuntimeConfig {
pub listen_addr: String,
}
impl Default for WebSocketRuntimeConfig {
fn default() -> Self {
Self {
listen_addr: "0.0.0.0:2346".to_string(),
}
}
}
impl WebSocketRuntimeConfig {
pub fn new(listen_addr: impl Into<String>) -> Self {
Self {
listen_addr: listen_addr.into(),
}
}
}
pub struct WebSocketRuntime {
config: WebSocketRuntimeConfig,
server: Arc<WsServer>,
}
impl WebSocketRuntime {
pub fn new(config: WebSocketRuntimeConfig) -> Self {
let server = Arc::new(WsServer::new(&config.listen_addr));
Self { config, server }
}
pub fn start(&self, token: CancellationToken) -> tokio::task::JoinHandle<Result<(), WsError>> {
let server = self.server.clone();
let handler: Arc<dyn WebSocketHandler> = Arc::new(DefaultWebSocketHandler::new());
tokio::spawn(async move {
let server_clone = server.clone();
let mut start_task = tokio::spawn(async move { server_clone.start(handler).await });
tokio::select! {
_ = token.cancelled() => {
let _ = server.stop().await;
let _ = (&mut start_task).await;
Ok(())
}
result = &mut start_task => {
match result {
Ok(inner) => inner,
Err(e) => Err(WsError::Connection(format!("start task panicked: {}", e))),
}
}
}
})
}
pub fn start_with_handler(
&self,
handler: Arc<dyn WebSocketHandler>,
token: CancellationToken,
) -> tokio::task::JoinHandle<Result<(), WsError>> {
let server = self.server.clone();
tokio::spawn(async move {
let server_clone = server.clone();
let mut start_task = tokio::spawn(async move { server_clone.start(handler).await });
tokio::select! {
_ = token.cancelled() => {
let _ = server.stop().await;
let _ = (&mut start_task).await;
Ok(())
}
result = &mut start_task => {
match result {
Ok(inner) => inner,
Err(e) => Err(WsError::Connection(format!("start task panicked: {}", e))),
}
}
}
})
}
pub async fn stop(&self) -> Result<(), WsError> {
self.server.stop().await
}
pub async fn connection_count(&self) -> usize {
self.server.connection_count().await
}
pub async fn broadcast_to_all(&self, data: Vec<u8>) -> Result<usize, WsError> {
self.server.broadcast_to_all(data).await
}
pub async fn is_running(&self) -> bool {
self.server.is_running().await
}
pub fn config(&self) -> &WebSocketRuntimeConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn test_websocket_runtime_config_default() {
let config = WebSocketRuntimeConfig::default();
assert_eq!(config.listen_addr, "0.0.0.0:2346");
}
#[test]
fn test_websocket_runtime_config_custom() {
let config = WebSocketRuntimeConfig::new("127.0.0.1:8080");
assert_eq!(config.listen_addr, "127.0.0.1:8080");
}
#[tokio::test]
async fn test_websocket_runtime_creation() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
assert_eq!(runtime.config().listen_addr, "127.0.0.1:0");
assert!(!runtime.is_running().await);
assert_eq!(runtime.connection_count().await, 0);
}
#[tokio::test]
async fn test_websocket_start_and_cancel() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
let token = CancellationToken::new();
let handle = runtime.start(token.clone());
tokio::time::sleep(Duration::from_millis(50)).await;
token.cancel();
let result = tokio::time::timeout(Duration::from_secs(2), handle).await;
assert!(result.is_ok(), "websocket task should stop on cancel");
}
#[tokio::test]
async fn test_websocket_start_with_handler_and_cancel() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
let handler: Arc<dyn WebSocketHandler> = Arc::new(DefaultWebSocketHandler::new());
let token = CancellationToken::new();
let handle = runtime.start_with_handler(handler, token.clone());
tokio::time::sleep(Duration::from_millis(50)).await;
token.cancel();
let result = tokio::time::timeout(Duration::from_secs(2), handle).await;
assert!(result.is_ok(), "websocket task should stop on cancel");
}
#[tokio::test]
async fn test_websocket_broadcast_no_connections() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
let result = runtime.broadcast_to_all(b"hello".to_vec()).await;
let _ = result;
}
#[tokio::test]
async fn test_websocket_manual_stop() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
let token = CancellationToken::new();
let handle = runtime.start(token.clone());
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = runtime.stop().await;
let _ = tokio::time::timeout(Duration::from_secs(2), handle).await;
}
#[test]
fn test_config_accessor() {
let runtime = WebSocketRuntime::new(WebSocketRuntimeConfig::new("0.0.0.0:9999"));
assert_eq!(runtime.config().listen_addr, "0.0.0.0:9999");
}
#[tokio::test]
async fn test_multiple_websocket_runtimes() {
let rt1 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
let rt2 = WebSocketRuntime::new(WebSocketRuntimeConfig::new("127.0.0.1:0"));
assert_eq!(rt1.connection_count().await, 0);
assert_eq!(rt2.connection_count().await, 0);
}
}