//! WebSocket service implementation for real-time model serving use crate::error::{ApiError, ApiResult}; use axum::extract::ws::{Message, WebSocket as AxumWebSocket}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::net::SocketAddr; use std::sync::Arc; use tokio::sync::{Mutex, RwLock}; use uuid::Uuid; /// WebSocket server configuration #[derive(Debug, Clone)] pub struct WebSocketConfig { pub port: u16, pub max_connections: usize, pub heartbeat_interval: u64, // seconds pub message_buffer_size: usize, } impl Default for WebSocketConfig { fn default() -> Self { Self { port: 8080, max_connections: 1000, heartbeat_interval: 30, message_buffer_size: 1024, } } } impl WebSocketConfig { pub fn set_max_connections(&mut self, max: usize) { self.max_connections = max; } pub fn set_heartbeat_interval(&mut self, seconds: u64) { self.heartbeat_interval = seconds; } } /// WebSocket message types #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] #[serde(tag = "type", content = "data")] pub enum WebSocketMessage { Ping(Vec), Pong(Vec), Heartbeat(u64), Broadcast { content: String }, Text(String), Binary(Vec), Close, } /// WebSocket request types #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum WebSocketRequest { Inference { model_name: String, input_data: Vec, stream: bool, }, ModelInfo { model_name: String, }, Subscribe { topic: String, }, Unsubscribe { topic: String, }, } /// WebSocket response types #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum WebSocketResponse { InferenceResult { model_name: String, outputs: Vec, latency_ms: f64, }, StreamToken { model_name: String, token: String, is_final: bool, }, ModelInfo { name: String, model_type: String, is_loaded: bool, }, Error { message: String, code: u16, }, Acknowledgment { request_id: String, }, } /// WebSocket connection wrapper #[derive(Debug)] pub struct WebSocketConnection { id: Uuid, socket: Arc>, is_active: Arc>, } impl WebSocketConnection { fn new(socket: AxumWebSocket) -> Self { Self { id: Uuid::new_v4(), socket: Arc::new(Mutex::new(socket)), is_active: Arc::new(RwLock::new(true)), } } #[must_use] pub fn connection_id(&self) -> Uuid { self.id } #[must_use] pub fn is_active(&self) -> bool { futures::executor::block_on(async { *self.is_active.read().await }) } pub async fn send(&self, msg: WebSocketMessage) -> ApiResult<()> { let message = match msg { WebSocketMessage::Ping(data) => Message::Ping(data), WebSocketMessage::Pong(data) => Message::Pong(data), WebSocketMessage::Text(text) => Message::Text(text), WebSocketMessage::Binary(data) => Message::Binary(data), WebSocketMessage::Close => Message::Close(None), _ => Message::Text(serde_json::to_string(&msg)?), }; let mut socket = self.socket.lock().await; socket .send(message) .await .map_err(|e| ApiError::WebSocket(e.to_string()))?; Ok(()) } pub async fn recv(&self) -> ApiResult { let mut socket = self.socket.lock().await; match socket.recv().await { Some(Ok(Message::Ping(data))) => Ok(WebSocketMessage::Pong(data)), Some(Ok(Message::Text(text))) => Ok(WebSocketMessage::Text(text)), Some(Ok(Message::Binary(data))) => Ok(WebSocketMessage::Binary(data)), Some(Ok(Message::Close(_))) => Ok(WebSocketMessage::Close), _ => Err(ApiError::WebSocket("Failed to receive message".to_string())), } } pub async fn send_json(&self, data: &T) -> ApiResult<()> { let json = serde_json::to_string(data)?; self.send(WebSocketMessage::Text(json)).await } pub async fn recv_json Deserialize<'de>>(&self) -> ApiResult { let msg = self.recv().await?; match msg { WebSocketMessage::Text(text) => { serde_json::from_str(&text).map_err(|e| ApiError::WebSocket(e.to_string())) } _ => Err(ApiError::WebSocket("Expected text message".to_string())), } } pub async fn try_recv(&self) -> ApiResult { self.recv().await } pub async fn close(&self) -> ApiResult<()> { let mut is_active = self.is_active.write().await; *is_active = false; self.send(WebSocketMessage::Close).await } } /// Connection manager for WebSocket connections #[derive(Debug)] pub struct ConnectionManager { connections: Arc>>>, max_connections: usize, } impl Default for ConnectionManager { fn default() -> Self { Self::new() } } impl ConnectionManager { #[must_use] pub fn new() -> Self { Self { connections: Arc::new(RwLock::new(HashMap::new())), max_connections: 1000, } } pub async fn add_connection( &self, ws: AxumWebSocket, ) -> Result, crate::ApiError> { let connections = self.connections.read().await; if connections.len() >= self.max_connections { return Err(crate::ApiError::WebSocket( "Connection limit reached".to_string(), )); } drop(connections); let conn = Arc::new(WebSocketConnection::new(ws)); let mut connections = self.connections.write().await; connections.insert(conn.connection_id(), conn.clone()); Ok(conn) } #[must_use] pub fn connection_count(&self) -> usize { futures::executor::block_on(async { self.connections.read().await.len() }) } pub async fn get_connection(&self, conn_id: Uuid) -> Option> { let connections = self.connections.read().await; connections.get(&conn_id).cloned() } pub async fn broadcast(&self, msg: WebSocketMessage) -> ApiResult<()> { let connections = self.connections.read().await; for conn in connections.values() { if conn.is_active() { conn.send(msg.clone()).await?; } } Ok(()) } pub async fn remove_connection(&self, conn_id: Uuid) { let mut connections = self.connections.write().await; connections.remove(&conn_id); } } /// WebSocket server for real-time model serving #[derive(Debug)] pub struct WebSocketServer { config: WebSocketConfig, models: Arc>>, connection_manager: Arc, is_ready: Arc>, } impl WebSocketServer { #[must_use] pub fn new(config: WebSocketConfig) -> Self { Self { config, models: Arc::new(RwLock::new(HashMap::new())), connection_manager: Arc::new(ConnectionManager::new()), is_ready: Arc::new(RwLock::new(true)), } } #[must_use] pub fn port(&self) -> u16 { self.config.port } #[must_use] pub fn max_connections(&self) -> usize { self.config.max_connections } #[must_use] pub fn is_ready(&self) -> bool { futures::executor::block_on(async { *self.is_ready.read().await }) } pub fn register_model(&mut self, name: &str, model_type: &str) { futures::executor::block_on(async { let mut models = self.models.write().await; models.insert(name.to_string(), model_type.to_string()); }); } pub async fn start(&self) -> ApiResult { let addr = SocketAddr::from(([127, 0, 0, 1], self.config.port)); Ok(addr) } pub async fn handle_connection(&self, ws: AxumWebSocket) { let conn = Arc::new(WebSocketConnection::new(ws)); let mut connections = self.connection_manager.connections.write().await; connections.insert(conn.connection_id(), conn.clone()); // Handle messages loop { match conn.recv().await { Ok(WebSocketMessage::Close) => break, Ok(msg) => { // Process message self.process_message(&conn, msg).await; } Err(_) => break, } } // Remove connection self.connection_manager .remove_connection(conn.connection_id()) .await; } async fn process_message(&self, conn: &WebSocketConnection, msg: WebSocketMessage) { match msg { WebSocketMessage::Ping(data) => { let _ = conn.send(WebSocketMessage::Pong(data)).await; } WebSocketMessage::Text(text) => { if let Ok(request) = serde_json::from_str::(&text) { self.handle_request(conn, request).await; } } _ => {} } } async fn handle_request(&self, conn: &WebSocketConnection, request: WebSocketRequest) { match request { WebSocketRequest::Inference { model_name, input_data, stream, } => { if stream { self.handle_streaming_inference(conn, &model_name, input_data) .await; } else { self.handle_inference(conn, &model_name, input_data).await; } } WebSocketRequest::ModelInfo { model_name } => { self.handle_model_info(conn, &model_name).await; } _ => {} } } async fn handle_inference( &self, conn: &WebSocketConnection, model_name: &str, _input_data: Vec, ) { let models = self.models.read().await; if !models.contains_key(model_name) { let error = WebSocketResponse::Error { message: format!("Model {model_name} not found"), code: 404, }; let _ = conn.send_json(&error).await; return; } // Simulate inference let response = WebSocketResponse::InferenceResult { model_name: model_name.to_string(), outputs: vec![0.1, 0.2, 0.3], latency_ms: 10.0, }; let _ = conn.send_json(&response).await; } async fn handle_streaming_inference( &self, conn: &WebSocketConnection, model_name: &str, _input_data: Vec, ) { let models = self.models.read().await; if !models.contains_key(model_name) { let error = WebSocketResponse::Error { message: format!("Model {model_name} not found"), code: 404, }; let _ = conn.send_json(&error).await; return; } // Simulate streaming tokens for i in 0..5 { let response = WebSocketResponse::StreamToken { model_name: model_name.to_string(), token: format!("token_{i}"), is_final: i == 4, }; if conn.send_json(&response).await.is_err() { break; } tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; } } async fn handle_model_info(&self, conn: &WebSocketConnection, model_name: &str) { let models = self.models.read().await; if let Some(model_type) = models.get(model_name) { let response = WebSocketResponse::ModelInfo { name: model_name.to_string(), model_type: model_type.clone(), is_loaded: true, }; let _ = conn.send_json(&response).await; } else { let error = WebSocketResponse::Error { message: format!("Model {model_name} not found"), code: 404, }; let _ = conn.send_json(&error).await; } } } #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn test_websocket_config() { let config = WebSocketConfig::default(); assert_eq!(config.port, 8080); assert_eq!(config.max_connections, 1000); } #[test] fn test_websocket_message_serialization() { let msg = WebSocketMessage::Text("hello".to_string()); let serialized = serde_json::to_string(&msg).unwrap(); assert!(serialized.contains("hello")); } #[test] fn test_websocket_request_serialization() { let req = WebSocketRequest::Inference { model_name: "test".to_string(), input_data: vec![1.0, 2.0], stream: true, }; let serialized = serde_json::to_string(&req).unwrap(); assert!(serialized.contains("test")); } #[test] fn test_websocket_response_serialization() { let resp = WebSocketResponse::StreamToken { model_name: "test".to_string(), token: "hello".to_string(), is_final: false, }; let serialized = serde_json::to_string(&resp).unwrap(); assert!(serialized.contains("hello")); } }