reinhardt_core/signals/
websocket_integration.rs1use super::error::SignalError;
30use super::signal::Signal;
31use parking_lot::RwLock;
32use serde::{Deserialize, Serialize, de::DeserializeOwned};
33use std::collections::HashMap;
34use std::fmt;
35use std::marker::PhantomData;
36use std::sync::Arc;
37
38#[derive(Debug, Clone, Serialize, Deserialize)]
50pub struct WebSocketMessage<T> {
51 pub event_type: String,
53 pub payload: T,
55 pub timestamp: u64,
57}
58
59impl<T> WebSocketMessage<T> {
60 pub fn new(event_type: impl Into<String>, payload: T) -> Self {
72 use std::time::{SystemTime, UNIX_EPOCH};
73
74 let timestamp = SystemTime::now()
75 .duration_since(UNIX_EPOCH)
76 .unwrap_or_default()
77 .as_millis() as u64;
78
79 Self {
80 event_type: event_type.into(),
81 payload,
82 timestamp,
83 }
84 }
85}
86
87pub trait WebSocketClient: Send + Sync {
91 fn send_message(&self, message: String) -> Result<(), SignalError>;
93
94 fn client_id(&self) -> &str;
96
97 fn is_connected(&self) -> bool;
99}
100
101pub struct MockWebSocketClient {
113 id: String,
114 messages: Arc<RwLock<Vec<String>>>,
115 connected: Arc<RwLock<bool>>,
116}
117
118impl MockWebSocketClient {
119 pub fn new(id: impl Into<String>) -> Self {
129 Self {
130 id: id.into(),
131 messages: Arc::new(RwLock::new(Vec::new())),
132 connected: Arc::new(RwLock::new(true)),
133 }
134 }
135
136 pub fn messages(&self) -> Vec<String> {
151 self.messages.read().clone()
152 }
153
154 pub fn disconnect(&self) {
168 *self.connected.write() = false;
169 }
170}
171
172impl WebSocketClient for MockWebSocketClient {
173 fn send_message(&self, message: String) -> Result<(), SignalError> {
174 if !self.is_connected() {
175 return Err(SignalError::new("Client is disconnected"));
176 }
177 self.messages.write().push(message);
178 Ok(())
179 }
180
181 fn client_id(&self) -> &str {
182 &self.id
183 }
184
185 fn is_connected(&self) -> bool {
186 *self.connected.read()
187 }
188}
189
190pub struct WebSocketSignalBridge {
203 clients: Arc<RwLock<HashMap<String, Arc<dyn WebSocketClient>>>>,
204}
205
206impl WebSocketSignalBridge {
207 pub fn new() -> Self {
217 Self {
218 clients: Arc::new(RwLock::new(HashMap::new())),
219 }
220 }
221
222 pub fn add_client(&self, client: Arc<dyn WebSocketClient>) {
237 self.clients
238 .write()
239 .insert(client.client_id().to_string(), client);
240 }
241
242 pub fn remove_client(&self, client_id: &str) {
258 self.clients.write().remove(client_id);
259 }
260
261 pub fn client_count(&self) -> usize {
272 self.clients.read().len()
273 }
274
275 pub fn broadcast(&self, message: String) -> Result<(), SignalError> {
293 let clients = self.clients.read();
294 let mut errors = Vec::new();
295
296 for client in clients.values() {
297 if client.is_connected()
298 && let Err(e) = client.send_message(message.clone())
299 {
300 errors.push(e);
301 }
302 }
303
304 if !errors.is_empty() {
305 return Err(SignalError::new(format!(
306 "Failed to send to {} clients",
307 errors.len()
308 )));
309 }
310
311 Ok(())
312 }
313
314 pub async fn connect_signal<T>(&self, signal: Signal<T>, event_type: impl Into<String>)
334 where
335 T: Serialize + Send + Sync + 'static,
336 {
337 let clients = Arc::clone(&self.clients);
338 let event_type = event_type.into();
339
340 signal.connect(move |instance| {
341 let clients = Arc::clone(&clients);
342 let event_type = event_type.clone();
343
344 async move {
345 let message = WebSocketMessage::new(&event_type, &*instance);
346 let json = serde_json::to_string(&message)
347 .map_err(|e| SignalError::new(format!("Serialization error: {}", e)))?;
348
349 let clients_read = clients.read();
350 for client in clients_read.values() {
351 if client.is_connected()
352 && let Err(e) = client.send_message(json.clone())
353 {
354 eprintln!("Failed to send WebSocket message: {}", e);
355 }
356 }
357
358 Ok(())
359 }
360 });
361 }
362
363 pub fn cleanup_disconnected(&self) {
381 let mut clients = self.clients.write();
382 clients.retain(|_, client| client.is_connected());
383 }
384}
385
386impl Default for WebSocketSignalBridge {
387 fn default() -> Self {
388 Self::new()
389 }
390}
391
392impl Clone for WebSocketSignalBridge {
393 fn clone(&self) -> Self {
394 Self {
395 clients: Arc::clone(&self.clients),
396 }
397 }
398}
399
400impl fmt::Debug for WebSocketSignalBridge {
401 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
402 f.debug_struct("WebSocketSignalBridge")
403 .field("client_count", &self.client_count())
404 .finish()
405 }
406}
407
408pub struct TypedWebSocketBroadcaster<T>
421where
422 T: Serialize + DeserializeOwned + Send + Sync + 'static,
423{
424 bridge: WebSocketSignalBridge,
425 event_type: String,
426 _phantom: PhantomData<T>,
427}
428
429impl<T> TypedWebSocketBroadcaster<T>
430where
431 T: Serialize + DeserializeOwned + Send + Sync + 'static,
432{
433 pub fn new(bridge: WebSocketSignalBridge, event_type: impl Into<String>) -> Self {
444 Self {
445 bridge,
446 event_type: event_type.into(),
447 _phantom: PhantomData,
448 }
449 }
450
451 pub fn broadcast(&self, payload: T) -> Result<(), SignalError> {
470 let message = WebSocketMessage::new(&self.event_type, payload);
471 let json = serde_json::to_string(&message)
472 .map_err(|e| SignalError::new(format!("Serialization error: {}", e)))?;
473
474 self.bridge.broadcast(json)
475 }
476}
477
478#[cfg(test)]
479mod tests {
480 use super::*;
481 use std::sync::Arc;
482
483 #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
484 struct TestPayload {
485 message: String,
486 }
487
488 #[test]
489 fn test_websocket_message_creation() {
490 let msg = WebSocketMessage::new("test", "payload");
491 assert_eq!(msg.event_type, "test");
492 assert_eq!(msg.payload, "payload");
493 assert!(msg.timestamp > 0);
494 }
495
496 #[test]
497 fn test_mock_websocket_client() {
498 let client = MockWebSocketClient::new("test-client");
499 assert_eq!(client.client_id(), "test-client");
500 assert!(client.is_connected());
501
502 client.send_message("Hello".to_string()).unwrap();
503 let messages = client.messages();
504 assert_eq!(messages.len(), 1);
505 assert_eq!(messages[0], "Hello");
506 }
507
508 #[test]
509 fn test_mock_client_disconnect() {
510 let client = MockWebSocketClient::new("test");
511 client.disconnect();
512 assert!(!client.is_connected());
513
514 let result = client.send_message("test".to_string());
515 assert!(result.is_err());
516 }
517
518 #[test]
519 fn test_websocket_bridge_add_remove_client() {
520 let bridge = WebSocketSignalBridge::new();
521 let client = Arc::new(MockWebSocketClient::new("client-1"));
522
523 bridge.add_client(client.clone());
524 assert_eq!(bridge.client_count(), 1);
525
526 bridge.remove_client("client-1");
527 assert_eq!(bridge.client_count(), 0);
528 }
529
530 #[test]
531 fn test_websocket_bridge_broadcast() {
532 let bridge = WebSocketSignalBridge::new();
533
534 let client1 = Arc::new(MockWebSocketClient::new("client-1"));
535 let client2 = Arc::new(MockWebSocketClient::new("client-2"));
536
537 bridge.add_client(client1.clone());
538 bridge.add_client(client2.clone());
539
540 bridge.broadcast("Test message".to_string()).unwrap();
541
542 assert_eq!(client1.messages().len(), 1);
543 assert_eq!(client2.messages().len(), 1);
544 assert_eq!(client1.messages()[0], "Test message");
545 }
546
547 #[tokio::test]
548 async fn test_websocket_bridge_connect_signal() {
549 let bridge = WebSocketSignalBridge::new();
550 let client = Arc::new(MockWebSocketClient::new("client-1"));
551 bridge.add_client(client.clone());
552
553 let signal = Signal::new(crate::signals::SignalName::custom("test_signal"));
554 bridge.connect_signal(signal.clone(), "test.event").await;
555
556 signal.send("test payload".to_string()).await.unwrap();
557
558 let messages = client.messages();
559 assert_eq!(messages.len(), 1);
560
561 let parsed: WebSocketMessage<String> = serde_json::from_str(&messages[0]).unwrap();
562 assert_eq!(parsed.event_type, "test.event");
563 }
564
565 #[test]
566 fn test_websocket_bridge_cleanup_disconnected() {
567 let bridge = WebSocketSignalBridge::new();
568
569 let client1 = Arc::new(MockWebSocketClient::new("client-1"));
570 let client2 = Arc::new(MockWebSocketClient::new("client-2"));
571
572 bridge.add_client(client1.clone());
573 bridge.add_client(client2.clone());
574
575 assert_eq!(bridge.client_count(), 2);
576
577 client1.disconnect();
578 bridge.cleanup_disconnected();
579
580 assert_eq!(bridge.client_count(), 1);
581 }
582
583 #[test]
584 fn test_typed_websocket_broadcaster() {
585 let bridge = WebSocketSignalBridge::new();
586 let client = Arc::new(MockWebSocketClient::new("client-1"));
587 bridge.add_client(client.clone());
588
589 let broadcaster = TypedWebSocketBroadcaster::new(bridge, "typed.event");
590
591 let payload = TestPayload {
592 message: "Hello".to_string(),
593 };
594
595 broadcaster.broadcast(payload.clone()).unwrap();
596
597 let messages = client.messages();
598 assert_eq!(messages.len(), 1);
599
600 let parsed: WebSocketMessage<TestPayload> = serde_json::from_str(&messages[0]).unwrap();
601 assert_eq!(parsed.event_type, "typed.event");
602 assert_eq!(parsed.payload, payload);
603 }
604}