Skip to main content

reinhardt_core/signals/
websocket_integration.rs

1//! WebSocket Integration - Real-time signal propagation to clients
2//!
3//! This module provides WebSocket integration for signals, allowing real-time
4//! signal propagation to connected WebSocket clients.
5//!
6//! # Examples
7//!
8//! ```rust,no_run
9//! use reinhardt_core::signals::websocket_integration::WebSocketSignalBridge;
10//! use reinhardt_core::signals::post_save;
11//!
12//! # #[tokio::main]
13//! # async fn main() -> Result<(), reinhardt_core::signals::SignalError> {
14//! # #[derive(Clone, serde::Serialize)]
15//! # struct User;
16//! # let user = User;
17//! // Create a WebSocket bridge
18//! let bridge = WebSocketSignalBridge::new();
19//!
20//! // Connect signals to WebSocket broadcast
21//! bridge.connect_signal(post_save::<User>(), "user.saved").await;
22//!
23//! // When a signal is emitted, it will be broadcast to WebSocket clients
24//! post_save::<User>().send(user).await?;
25//! # Ok(())
26//! # }
27//! ```
28
29use 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/// WebSocket message format for signal events
39///
40/// # Examples
41///
42/// ```
43/// use reinhardt_core::signals::websocket_integration::WebSocketMessage;
44/// use serde_json::json;
45///
46/// let msg = WebSocketMessage::new("user.created", json!({"id": 123}));
47/// assert_eq!(msg.event_type, "user.created");
48/// ```
49#[derive(Debug, Clone, Serialize, Deserialize)]
50pub struct WebSocketMessage<T> {
51	/// Event type identifier
52	pub event_type: String,
53	/// Event payload
54	pub payload: T,
55	/// Message timestamp (Unix timestamp in milliseconds)
56	pub timestamp: u64,
57}
58
59impl<T> WebSocketMessage<T> {
60	/// Create a new WebSocket message
61	///
62	/// # Examples
63	///
64	/// ```
65	/// use reinhardt_core::signals::websocket_integration::WebSocketMessage;
66	///
67	/// let msg = WebSocketMessage::new("notification", "Hello");
68	/// assert_eq!(msg.event_type, "notification");
69	/// assert_eq!(msg.payload, "Hello");
70	/// ```
71	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
87/// WebSocket client connection trait
88///
89/// Implement this trait to integrate with your WebSocket server
90pub trait WebSocketClient: Send + Sync {
91	/// Send a message to this client
92	fn send_message(&self, message: String) -> Result<(), SignalError>;
93
94	/// Get the client ID
95	fn client_id(&self) -> &str;
96
97	/// Check if the client is still connected
98	fn is_connected(&self) -> bool;
99}
100
101/// In-memory WebSocket client for testing
102///
103/// # Examples
104///
105/// ```
106/// use reinhardt_core::signals::websocket_integration::{MockWebSocketClient, WebSocketClient};
107///
108/// let client = MockWebSocketClient::new("client-1");
109/// assert_eq!(client.client_id(), "client-1");
110/// assert!(client.is_connected());
111/// ```
112pub struct MockWebSocketClient {
113	id: String,
114	messages: Arc<RwLock<Vec<String>>>,
115	connected: Arc<RwLock<bool>>,
116}
117
118impl MockWebSocketClient {
119	/// Create a new mock WebSocket client
120	///
121	/// # Examples
122	///
123	/// ```
124	/// use reinhardt_core::signals::websocket_integration::MockWebSocketClient;
125	///
126	/// let client = MockWebSocketClient::new("test-client");
127	/// ```
128	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	/// Get all messages received by this client
137	///
138	/// # Examples
139	///
140	/// ```
141	/// use reinhardt_core::signals::websocket_integration::{MockWebSocketClient, WebSocketClient};
142	///
143	/// let client = MockWebSocketClient::new("test");
144	/// client.send_message("Hello".to_string()).unwrap();
145	///
146	/// let messages = client.messages();
147	/// assert_eq!(messages.len(), 1);
148	/// assert_eq!(messages[0], "Hello");
149	/// ```
150	pub fn messages(&self) -> Vec<String> {
151		self.messages.read().clone()
152	}
153
154	/// Disconnect this client
155	///
156	/// # Examples
157	///
158	/// ```
159	/// use reinhardt_core::signals::websocket_integration::{MockWebSocketClient, WebSocketClient};
160	///
161	/// let client = MockWebSocketClient::new("test");
162	/// assert!(client.is_connected());
163	///
164	/// client.disconnect();
165	/// assert!(!client.is_connected());
166	/// ```
167	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
190/// WebSocket signal bridge
191///
192/// Bridges signals to WebSocket clients, broadcasting signal events
193/// to connected clients in real-time
194///
195/// # Examples
196///
197/// ```
198/// use reinhardt_core::signals::websocket_integration::WebSocketSignalBridge;
199///
200/// let bridge = WebSocketSignalBridge::new();
201/// ```
202pub struct WebSocketSignalBridge {
203	clients: Arc<RwLock<HashMap<String, Arc<dyn WebSocketClient>>>>,
204}
205
206impl WebSocketSignalBridge {
207	/// Create a new WebSocket signal bridge
208	///
209	/// # Examples
210	///
211	/// ```
212	/// use reinhardt_core::signals::websocket_integration::WebSocketSignalBridge;
213	///
214	/// let bridge = WebSocketSignalBridge::new();
215	/// ```
216	pub fn new() -> Self {
217		Self {
218			clients: Arc::new(RwLock::new(HashMap::new())),
219		}
220	}
221
222	/// Add a WebSocket client to the bridge
223	///
224	/// # Examples
225	///
226	/// ```
227	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, MockWebSocketClient};
228	/// use std::sync::Arc;
229	///
230	/// let bridge = WebSocketSignalBridge::new();
231	/// let client = Arc::new(MockWebSocketClient::new("client-1"));
232	/// bridge.add_client(client);
233	///
234	/// assert_eq!(bridge.client_count(), 1);
235	/// ```
236	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	/// Remove a WebSocket client from the bridge
243	///
244	/// # Examples
245	///
246	/// ```
247	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, MockWebSocketClient};
248	/// use std::sync::Arc;
249	///
250	/// let bridge = WebSocketSignalBridge::new();
251	/// let client = Arc::new(MockWebSocketClient::new("client-1"));
252	/// bridge.add_client(client);
253	///
254	/// bridge.remove_client("client-1");
255	/// assert_eq!(bridge.client_count(), 0);
256	/// ```
257	pub fn remove_client(&self, client_id: &str) {
258		self.clients.write().remove(client_id);
259	}
260
261	/// Get the number of connected clients
262	///
263	/// # Examples
264	///
265	/// ```
266	/// use reinhardt_core::signals::websocket_integration::WebSocketSignalBridge;
267	///
268	/// let bridge = WebSocketSignalBridge::new();
269	/// assert_eq!(bridge.client_count(), 0);
270	/// ```
271	pub fn client_count(&self) -> usize {
272		self.clients.read().len()
273	}
274
275	/// Broadcast a message to all connected clients
276	///
277	/// # Examples
278	///
279	/// ```
280	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, MockWebSocketClient};
281	/// use std::sync::Arc;
282	///
283	/// let bridge = WebSocketSignalBridge::new();
284	/// let client = Arc::new(MockWebSocketClient::new("client-1"));
285	/// bridge.add_client(client.clone());
286	///
287	/// bridge.broadcast("Hello all".to_string()).unwrap();
288	///
289	/// let messages = client.messages();
290	/// assert_eq!(messages.len(), 1);
291	/// ```
292	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	/// Connect a signal to WebSocket broadcast
315	///
316	/// When the signal is emitted, it will be serialized and broadcast to all clients
317	///
318	/// # Examples
319	///
320	/// ```rust,no_run
321	/// use reinhardt_core::signals::websocket_integration::WebSocketSignalBridge;
322	/// use reinhardt_core::signals::post_save;
323	/// # use serde::{Serialize, Deserialize};
324	///
325	/// # #[tokio::main]
326	/// # async fn main() {
327	/// # #[derive(Debug, Clone, Serialize, Deserialize)]
328	/// # struct User { id: Option<i64> }
329	/// let bridge = WebSocketSignalBridge::new();
330	/// bridge.connect_signal(post_save::<User>(), "user.saved").await;
331	/// # }
332	/// ```
333	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	/// Clean up disconnected clients
364	///
365	/// # Examples
366	///
367	/// ```
368	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, MockWebSocketClient};
369	/// use std::sync::Arc;
370	///
371	/// let bridge = WebSocketSignalBridge::new();
372	/// let client = Arc::new(MockWebSocketClient::new("client-1"));
373	/// bridge.add_client(client.clone());
374	///
375	/// client.disconnect();
376	/// bridge.cleanup_disconnected();
377	///
378	/// assert_eq!(bridge.client_count(), 0);
379	/// ```
380	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
408/// Typed WebSocket signal broadcaster
409///
410/// A type-safe wrapper for broadcasting specific signal types to WebSocket clients
411///
412/// # Examples
413///
414/// ```
415/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, TypedWebSocketBroadcaster};
416///
417/// let bridge = WebSocketSignalBridge::new();
418/// let broadcaster = TypedWebSocketBroadcaster::<String>::new(bridge, "string_event");
419/// ```
420pub 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	/// Create a new typed broadcaster
434	///
435	/// # Examples
436	///
437	/// ```
438	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, TypedWebSocketBroadcaster};
439	///
440	/// let bridge = WebSocketSignalBridge::new();
441	/// let broadcaster = TypedWebSocketBroadcaster::<String>::new(bridge, "test");
442	/// ```
443	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	/// Broadcast a typed message
452	///
453	/// # Examples
454	///
455	/// ```
456	/// use reinhardt_core::signals::websocket_integration::{WebSocketSignalBridge, TypedWebSocketBroadcaster, MockWebSocketClient};
457	/// use std::sync::Arc;
458	///
459	/// let bridge = WebSocketSignalBridge::new();
460	/// let client = Arc::new(MockWebSocketClient::new("client-1"));
461	/// bridge.add_client(client.clone());
462	///
463	/// let broadcaster = TypedWebSocketBroadcaster::new(bridge, "message");
464	/// broadcaster.broadcast("Hello".to_string()).unwrap();
465	///
466	/// let messages = client.messages();
467	/// assert_eq!(messages.len(), 1);
468	/// ```
469	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}