1use crate::connection::Message;
8use async_trait::async_trait;
9use std::collections::HashMap;
10use std::sync::Arc;
11use tokio::sync::{RwLock, mpsc};
12
13pub type ChannelResult<T> = Result<T, ChannelError>;
15
16#[derive(Debug, thiserror::Error)]
18pub enum ChannelError {
19 #[error("Send error: {0}")]
21 SendError(String),
22 #[error("Receive error: {0}")]
24 ReceiveError(String),
25 #[error("Channel not found: {0}")]
27 ChannelNotFound(String),
28 #[error("Group not found: {0}")]
30 GroupNotFound(String),
31 #[error("Serialization error: {0}")]
33 SerializationError(String),
34 #[error("Authentication required for Redis connection")]
36 AuthenticationRequired,
37}
38
39#[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 pub fn new(sender: String, payload: Message) -> Self {
64 Self {
65 sender,
66 payload,
67 metadata: HashMap::new(),
68 }
69 }
70
71 pub fn with_metadata(mut self, key: String, value: String) -> Self {
73 self.metadata.insert(key, value);
74 self
75 }
76
77 pub fn sender(&self) -> &str {
79 &self.sender
80 }
81
82 pub fn payload(&self) -> &Message {
84 &self.payload
85 }
86
87 pub fn metadata(&self, key: &str) -> Option<&String> {
89 self.metadata.get(key)
90 }
91}
92
93#[async_trait]
95pub trait ChannelLayer: Send + Sync {
96 async fn send(&self, channel: &str, message: ChannelMessage) -> ChannelResult<()>;
98
99 async fn receive(&self, channel: &str) -> ChannelResult<Option<ChannelMessage>>;
101
102 async fn group_add(&self, group: &str, channel: &str) -> ChannelResult<()>;
104
105 async fn group_discard(&self, group: &str, channel: &str) -> ChannelResult<()>;
107
108 async fn group_send(&self, group: &str, message: ChannelMessage) -> ChannelResult<()>;
110}
111
112pub 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 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 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 pub async fn channel_count(&self) -> usize {
169 let channels = self.channels.read().await;
170 channels.len()
171 }
172
173 pub async fn group_count(&self) -> usize {
175 let groups = self.groups.read().await;
176 groups.len()
177 }
178
179 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 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
261pub struct ChannelLayerWrapper {
272 layer: Box<dyn ChannelLayer>,
273}
274
275impl ChannelLayerWrapper {
276 pub fn new(layer: Box<dyn ChannelLayer>) -> Self {
278 Self { layer }
279 }
280
281 pub async fn send(&self, channel: &str, message: ChannelMessage) -> ChannelResult<()> {
283 self.layer.send(channel, message).await
284 }
285
286 pub async fn receive(&self, channel: &str) -> ChannelResult<Option<ChannelMessage>> {
288 self.layer.receive(channel).await
289 }
290
291 pub async fn group_add(&self, group: &str, channel: &str) -> ChannelResult<()> {
293 self.layer.group_add(group, channel).await
294 }
295
296 pub async fn group_discard(&self, group: &str, channel: &str) -> ChannelResult<()> {
298 self.layer.group_discard(group, channel).await
299 }
300
301 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}