Skip to main content

reinhardt_websockets/
channels.rs

1//! Channel layers for distributed WebSocket systems
2//!
3//! This module provides channel layer abstractions for distributed WebSocket communication,
4//! inspired by Django Channels. Channel layers enable multiple application instances
5//! to communicate with each other and share WebSocket connections.
6
7use crate::connection::Message;
8use async_trait::async_trait;
9use std::collections::HashMap;
10use std::sync::Arc;
11use tokio::sync::{RwLock, mpsc};
12
13/// Channel layer result type
14pub type ChannelResult<T> = Result<T, ChannelError>;
15
16/// Channel layer errors
17#[derive(Debug, thiserror::Error)]
18pub enum ChannelError {
19	/// Failed to send a message through the channel.
20	#[error("Send error: {0}")]
21	SendError(String),
22	/// Failed to receive a message from the channel.
23	#[error("Receive error: {0}")]
24	ReceiveError(String),
25	/// The specified channel was not found.
26	#[error("Channel not found: {0}")]
27	ChannelNotFound(String),
28	/// The specified group was not found.
29	#[error("Group not found: {0}")]
30	GroupNotFound(String),
31	/// Failed to serialize or deserialize a message.
32	#[error("Serialization error: {0}")]
33	SerializationError(String),
34	/// Authentication is required to connect to the backing store.
35	#[error("Authentication required for Redis connection")]
36	AuthenticationRequired,
37}
38
39/// Channel message for distributed communication
40///
41/// # Examples
42///
43/// ```
44/// use reinhardt_websockets::channels::ChannelMessage;
45/// use reinhardt_websockets::Message;
46///
47/// let msg = ChannelMessage::new(
48///     "user_1".to_string(),
49///     Message::text("Hello".to_string()),
50/// );
51///
52/// assert_eq!(msg.sender(), "user_1");
53/// ```
54#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
55pub struct ChannelMessage {
56	sender: String,
57	payload: Message,
58	metadata: HashMap<String, String>,
59}
60
61impl ChannelMessage {
62	/// Create a new channel message
63	pub fn new(sender: String, payload: Message) -> Self {
64		Self {
65			sender,
66			payload,
67			metadata: HashMap::new(),
68		}
69	}
70
71	/// Add metadata to the message
72	pub fn with_metadata(mut self, key: String, value: String) -> Self {
73		self.metadata.insert(key, value);
74		self
75	}
76
77	/// Get the sender
78	pub fn sender(&self) -> &str {
79		&self.sender
80	}
81
82	/// Get the payload
83	pub fn payload(&self) -> &Message {
84		&self.payload
85	}
86
87	/// Get metadata
88	pub fn metadata(&self, key: &str) -> Option<&String> {
89		self.metadata.get(key)
90	}
91}
92
93/// Channel layer trait for distributed messaging
94#[async_trait]
95pub trait ChannelLayer: Send + Sync {
96	/// Send a message to a specific channel
97	async fn send(&self, channel: &str, message: ChannelMessage) -> ChannelResult<()>;
98
99	/// Receive a message from a specific channel
100	async fn receive(&self, channel: &str) -> ChannelResult<Option<ChannelMessage>>;
101
102	/// Add a channel to a group
103	async fn group_add(&self, group: &str, channel: &str) -> ChannelResult<()>;
104
105	/// Remove a channel from a group
106	async fn group_discard(&self, group: &str, channel: &str) -> ChannelResult<()>;
107
108	/// Send a message to all channels in a group
109	async fn group_send(&self, group: &str, message: ChannelMessage) -> ChannelResult<()>;
110}
111
112/// In-memory channel layer implementation
113///
114/// # Examples
115///
116/// ```
117/// use reinhardt_websockets::channels::{InMemoryChannelLayer, ChannelLayer, ChannelMessage};
118/// use reinhardt_websockets::Message;
119///
120/// # tokio_test::block_on(async {
121/// let layer = InMemoryChannelLayer::new();
122///
123/// let msg = ChannelMessage::new(
124///     "user_1".to_string(),
125///     Message::text("Hello".to_string()),
126/// );
127///
128/// layer.send("channel_1", msg.clone()).await.unwrap();
129/// let received = layer.receive("channel_1").await.unwrap();
130///
131/// assert!(received.is_some());
132/// # });
133/// ```
134pub struct InMemoryChannelLayer {
135	channels: Arc<RwLock<HashMap<String, mpsc::UnboundedSender<ChannelMessage>>>>,
136	receivers: Arc<RwLock<HashMap<String, mpsc::UnboundedReceiver<ChannelMessage>>>>,
137	groups: Arc<RwLock<HashMap<String, Vec<String>>>>,
138}
139
140impl InMemoryChannelLayer {
141	/// Create a new in-memory channel layer
142	pub fn new() -> Self {
143		Self {
144			channels: Arc::new(RwLock::new(HashMap::new())),
145			receivers: Arc::new(RwLock::new(HashMap::new())),
146			groups: Arc::new(RwLock::new(HashMap::new())),
147		}
148	}
149
150	/// Create or get a channel
151	async fn get_or_create_channel(&self, channel: &str) -> mpsc::UnboundedSender<ChannelMessage> {
152		let mut channels = self.channels.write().await;
153
154		if let Some(tx) = channels.get(channel) {
155			return tx.clone();
156		}
157
158		let (tx, rx) = mpsc::unbounded_channel();
159		channels.insert(channel.to_string(), tx.clone());
160
161		let mut receivers = self.receivers.write().await;
162		receivers.insert(channel.to_string(), rx);
163
164		tx
165	}
166
167	/// Get channel count
168	pub async fn channel_count(&self) -> usize {
169		let channels = self.channels.read().await;
170		channels.len()
171	}
172
173	/// Get group count
174	pub async fn group_count(&self) -> usize {
175		let groups = self.groups.read().await;
176		groups.len()
177	}
178
179	/// Clear all channels and groups
180	pub async fn clear(&self) {
181		let mut channels = self.channels.write().await;
182		let mut receivers = self.receivers.write().await;
183		let mut groups = self.groups.write().await;
184
185		channels.clear();
186		receivers.clear();
187		groups.clear();
188	}
189}
190
191impl Default for InMemoryChannelLayer {
192	fn default() -> Self {
193		Self::new()
194	}
195}
196
197#[async_trait]
198impl ChannelLayer for InMemoryChannelLayer {
199	async fn send(&self, channel: &str, message: ChannelMessage) -> ChannelResult<()> {
200		let tx = self.get_or_create_channel(channel).await;
201
202		tx.send(message)
203			.map_err(|e| ChannelError::SendError(e.to_string()))
204	}
205
206	async fn receive(&self, channel: &str) -> ChannelResult<Option<ChannelMessage>> {
207		let mut receivers = self.receivers.write().await;
208
209		if let Some(rx) = receivers.get_mut(channel) {
210			Ok(rx.try_recv().ok())
211		} else {
212			Ok(None)
213		}
214	}
215
216	async fn group_add(&self, group: &str, channel: &str) -> ChannelResult<()> {
217		let mut groups = self.groups.write().await;
218
219		let channels = groups.entry(group.to_string()).or_insert_with(Vec::new);
220
221		if !channels.contains(&channel.to_string()) {
222			channels.push(channel.to_string());
223		}
224
225		Ok(())
226	}
227
228	async fn group_discard(&self, group: &str, channel: &str) -> ChannelResult<()> {
229		let mut groups = self.groups.write().await;
230
231		if let Some(channels) = groups.get_mut(group) {
232			channels.retain(|c| c != channel);
233
234			if channels.is_empty() {
235				groups.remove(group);
236			}
237		}
238
239		Ok(())
240	}
241
242	async fn group_send(&self, group: &str, message: ChannelMessage) -> ChannelResult<()> {
243		// Collect channel IDs while holding the groups lock, then release it
244		// before calling self.send() to avoid ABBA deadlock with clear()
245		let channel_ids = {
246			let groups = self.groups.read().await;
247			groups
248				.get(group)
249				.ok_or_else(|| ChannelError::GroupNotFound(group.to_string()))?
250				.clone()
251		};
252
253		for channel in &channel_ids {
254			self.send(channel, message.clone()).await?;
255		}
256
257		Ok(())
258	}
259}
260
261/// Channel layer wrapper with additional features
262///
263/// # Examples
264///
265/// ```
266/// use reinhardt_websockets::channels::{ChannelLayerWrapper, InMemoryChannelLayer};
267///
268/// let layer = InMemoryChannelLayer::new();
269/// let wrapper = ChannelLayerWrapper::new(Box::new(layer));
270/// ```
271pub struct ChannelLayerWrapper {
272	layer: Box<dyn ChannelLayer>,
273}
274
275impl ChannelLayerWrapper {
276	/// Create a new channel layer wrapper
277	pub fn new(layer: Box<dyn ChannelLayer>) -> Self {
278		Self { layer }
279	}
280
281	/// Send a message to a channel
282	pub async fn send(&self, channel: &str, message: ChannelMessage) -> ChannelResult<()> {
283		self.layer.send(channel, message).await
284	}
285
286	/// Receive a message from a channel
287	pub async fn receive(&self, channel: &str) -> ChannelResult<Option<ChannelMessage>> {
288		self.layer.receive(channel).await
289	}
290
291	/// Add a channel to a group
292	pub async fn group_add(&self, group: &str, channel: &str) -> ChannelResult<()> {
293		self.layer.group_add(group, channel).await
294	}
295
296	/// Remove a channel from a group
297	pub async fn group_discard(&self, group: &str, channel: &str) -> ChannelResult<()> {
298		self.layer.group_discard(group, channel).await
299	}
300
301	/// Send a message to a group
302	pub async fn group_send(&self, group: &str, message: ChannelMessage) -> ChannelResult<()> {
303		self.layer.group_send(group, message).await
304	}
305}
306
307#[cfg(test)]
308mod tests {
309	use super::*;
310
311	#[test]
312	fn test_channel_message_creation() {
313		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Hello".to_string()));
314		assert_eq!(msg.sender(), "user_1");
315	}
316
317	#[test]
318	fn test_channel_message_metadata() {
319		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Hello".to_string()))
320			.with_metadata("priority".to_string(), "high".to_string());
321
322		assert_eq!(msg.metadata("priority").unwrap(), "high");
323	}
324
325	#[tokio::test]
326	async fn test_in_memory_channel_layer_send_receive() {
327		let layer = InMemoryChannelLayer::new();
328		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Hello".to_string()));
329
330		layer.send("channel_1", msg.clone()).await.unwrap();
331
332		let received = layer.receive("channel_1").await.unwrap();
333		assert!(received.is_some());
334		assert_eq!(received.unwrap().sender(), "user_1");
335	}
336
337	#[tokio::test]
338	async fn test_in_memory_channel_layer_group_add() {
339		let layer = InMemoryChannelLayer::new();
340
341		layer.group_add("group_1", "channel_1").await.unwrap();
342		layer.group_add("group_1", "channel_2").await.unwrap();
343
344		assert_eq!(layer.group_count().await, 1);
345	}
346
347	#[tokio::test]
348	async fn test_in_memory_channel_layer_group_discard() {
349		let layer = InMemoryChannelLayer::new();
350
351		layer.group_add("group_1", "channel_1").await.unwrap();
352		layer.group_add("group_1", "channel_2").await.unwrap();
353
354		layer.group_discard("group_1", "channel_1").await.unwrap();
355
356		assert_eq!(layer.group_count().await, 1);
357	}
358
359	#[tokio::test]
360	async fn test_in_memory_channel_layer_group_send() {
361		let layer = InMemoryChannelLayer::new();
362
363		layer.group_add("group_1", "channel_1").await.unwrap();
364		layer.group_add("group_1", "channel_2").await.unwrap();
365
366		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Broadcast".to_string()));
367
368		layer.group_send("group_1", msg).await.unwrap();
369
370		let received1 = layer.receive("channel_1").await.unwrap();
371		let received2 = layer.receive("channel_2").await.unwrap();
372
373		assert!(received1.is_some());
374		assert!(received2.is_some());
375	}
376
377	#[tokio::test]
378	async fn test_in_memory_channel_layer_clear() {
379		let layer = InMemoryChannelLayer::new();
380		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Test".to_string()));
381
382		layer.send("channel_1", msg).await.unwrap();
383		layer.group_add("group_1", "channel_1").await.unwrap();
384
385		assert_eq!(layer.channel_count().await, 1);
386		assert_eq!(layer.group_count().await, 1);
387
388		layer.clear().await;
389
390		assert_eq!(layer.channel_count().await, 0);
391		assert_eq!(layer.group_count().await, 0);
392	}
393
394	#[tokio::test]
395	async fn test_channel_layer_wrapper() {
396		let layer = InMemoryChannelLayer::new();
397		let wrapper = ChannelLayerWrapper::new(Box::new(layer));
398
399		let msg = ChannelMessage::new("user_1".to_string(), Message::text("Hello".to_string()));
400
401		wrapper.send("channel_1", msg.clone()).await.unwrap();
402
403		let received = wrapper.receive("channel_1").await.unwrap();
404		assert!(received.is_some());
405	}
406}