1use crate::error::WsError;
2use async_trait::async_trait;
3use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5use std::sync::Arc;
6use tokio::sync::RwLock;
7
8#[derive(Debug, Clone)]
9pub struct WebSocketMessage {
10 pub msg_type: MessageType,
11 pub payload: Vec<u8>,
12 pub sender_id: Option<i64>,
13 pub room_id: Option<String>,
14 pub timestamp: i64,
15}
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
18pub enum MessageType {
19 #[default]
20 Text,
21 Binary,
22 Ping,
23 Pong,
24 Join,
25 Leave,
26 Subscribe,
27 Unsubscribe,
28 Notification,
29 System,
30}
31
32#[derive(Debug, Clone)]
33pub struct WebSocketConnection {
34 pub id: String,
35 pub user_id: Option<i64>,
36 pub remote_addr: Option<String>,
37 pub is_authenticated: bool,
38 pub subscriptions: Vec<String>,
39}
40
41impl WebSocketConnection {
42 pub fn new(id: impl Into<String>) -> Self {
43 Self {
44 id: id.into(),
45 user_id: None,
46 remote_addr: None,
47 is_authenticated: false,
48 subscriptions: Vec::new(),
49 }
50 }
51
52 pub fn with_user(mut self, user_id: i64) -> Self {
53 self.user_id = Some(user_id);
54 self.is_authenticated = true;
55 self
56 }
57
58 pub fn with_address(mut self, addr: impl Into<String>) -> Self {
59 self.remote_addr = Some(addr.into());
60 self
61 }
62
63 pub fn subscribe(&mut self, room: impl Into<String>) {
64 let room = room.into();
65 if !self.subscriptions.contains(&room) {
66 self.subscriptions.push(room);
67 }
68 }
69
70 pub fn unsubscribe(&mut self, room: &str) {
71 self.subscriptions.retain(|r| r != room);
72 }
73}
74
75#[async_trait]
76pub trait WebSocketHandler: Send + Sync {
77 async fn on_message(
78 &self,
79 conn: &WebSocketConnection,
80 msg: WebSocketMessage,
81 ) -> Result<Option<WebSocketMessage>, WsError>;
82
83 async fn on_connect(&self, conn: &WebSocketConnection) -> Result<(), WsError>;
84
85 async fn on_disconnect(&self, conn: &WebSocketConnection);
86
87 fn authenticate(&self, token: &str) -> Result<UserId, WsError>;
88}
89
90pub type UserId = i64;
91
92pub struct WsContext {
93 pub connection_id: String,
94 pub user_id: Option<i64>,
95 pub metadata: std::collections::HashMap<String, String>,
96}
97
98impl WsContext {
99 pub fn new(connection_id: impl Into<String>) -> Self {
100 Self {
101 connection_id: connection_id.into(),
102 user_id: None,
103 metadata: std::collections::HashMap::new(),
104 }
105 }
106
107 pub fn with_user(mut self, user_id: i64) -> Self {
108 self.user_id = Some(user_id);
109 self
110 }
111
112 pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
113 self.metadata.insert(key.into(), value.into());
114 self
115 }
116}
117
118pub struct WsMessageBuilder {
119 msg_type: MessageType,
120 payload: Vec<u8>,
121 sender_id: Option<i64>,
122 room_id: Option<String>,
123}
124
125impl WsMessageBuilder {
126 pub fn new() -> Self {
127 Self {
128 msg_type: MessageType::Text,
129 payload: Vec::new(),
130 sender_id: None,
131 room_id: None,
132 }
133 }
134
135 pub fn text(mut self, text: impl Into<String>) -> Self {
136 self.msg_type = MessageType::Text;
137 self.payload = text.into().into_bytes();
138 self
139 }
140
141 pub fn binary(mut self, data: Vec<u8>) -> Self {
142 self.msg_type = MessageType::Binary;
143 self.payload = data;
144 self
145 }
146
147 pub fn json<T: serde::Serialize>(mut self, data: &T) -> Result<Self, WsError> {
148 self.msg_type = MessageType::Text;
149 self.payload = serde_json::to_vec(data)?;
150 Ok(self)
151 }
152
153 pub fn with_sender(mut self, user_id: i64) -> Self {
154 self.sender_id = Some(user_id);
155 self
156 }
157
158 pub fn with_room(mut self, room: impl Into<String>) -> Self {
159 self.room_id = Some(room.into());
160 self
161 }
162
163 pub fn notification(mut self) -> Self {
164 self.msg_type = MessageType::Notification;
165 self
166 }
167
168 pub fn system(mut self) -> Self {
169 self.msg_type = MessageType::System;
170 self
171 }
172
173 pub fn build(self) -> WebSocketMessage {
174 WebSocketMessage {
175 msg_type: self.msg_type,
176 payload: self.payload,
177 sender_id: self.sender_id,
178 room_id: self.room_id,
179 timestamp: current_timestamp(),
180 }
181 }
182}
183
184impl Default for WsMessageBuilder {
185 fn default() -> Self {
186 Self::new()
187 }
188}
189
190fn current_timestamp() -> i64 {
191 use std::time::{SystemTime, UNIX_EPOCH};
192 SystemTime::now()
193 .duration_since(UNIX_EPOCH)
194 .unwrap_or_default()
195 .as_millis() as i64
196}
197
198pub struct DefaultWebSocketHandler {
210 connections: Arc<RwLock<HashMap<String, WebSocketConnection>>>,
211 message_log: Arc<RwLock<std::collections::VecDeque<WebSocketMessage>>>,
212}
213
214const MAX_MESSAGE_LOG: usize = 10_000;
216
217impl DefaultWebSocketHandler {
218 pub fn new() -> Self {
219 Self {
220 connections: Arc::new(RwLock::new(HashMap::new())),
221 message_log: Arc::new(RwLock::new(std::collections::VecDeque::new())),
222 }
223 }
224
225 pub async fn connection_count(&self) -> usize {
226 self.connections.read().await.len()
227 }
228
229 pub async fn is_connected(&self, connection_id: &str) -> bool {
230 self.connections.read().await.contains_key(connection_id)
231 }
232
233 pub async fn message_count(&self) -> usize {
234 self.message_log.read().await.len()
235 }
236
237 pub async fn messages(&self) -> Vec<WebSocketMessage> {
238 self.message_log.read().await.iter().cloned().collect()
239 }
240
241 pub async fn get_connection(&self, connection_id: &str) -> Option<WebSocketConnection> {
242 self.connections.read().await.get(connection_id).cloned()
243 }
244}
245
246impl Default for DefaultWebSocketHandler {
247 fn default() -> Self {
248 Self::new()
249 }
250}
251
252#[async_trait]
253impl WebSocketHandler for DefaultWebSocketHandler {
254 async fn on_message(
255 &self,
256 conn: &WebSocketConnection,
257 msg: WebSocketMessage,
258 ) -> Result<Option<WebSocketMessage>, WsError> {
259 {
261 let mut log = self.message_log.write().await;
262 if log.len() >= MAX_MESSAGE_LOG {
263 log.pop_front();
264 }
265 log.push_back(msg.clone());
266 }
267
268 match msg.msg_type {
269 MessageType::Text => {
270 let response = WebSocketMessage {
271 msg_type: MessageType::Text,
272 payload: msg.payload.clone(),
273 sender_id: conn.user_id,
274 room_id: None,
275 timestamp: current_timestamp(),
276 };
277 Ok(Some(response))
278 }
279 MessageType::Ping => {
280 let response = WebSocketMessage {
281 msg_type: MessageType::Pong,
282 payload: msg.payload.clone(),
283 sender_id: None,
284 room_id: None,
285 timestamp: current_timestamp(),
286 };
287 Ok(Some(response))
288 }
289 MessageType::Subscribe => {
290 let room = String::from_utf8_lossy(&msg.payload).to_string();
291 let ack = format!("subscribed:{}", room);
292 let response = WebSocketMessage {
293 msg_type: MessageType::System,
294 payload: ack.into_bytes(),
295 sender_id: None,
296 room_id: Some(room),
297 timestamp: current_timestamp(),
298 };
299 Ok(Some(response))
300 }
301 MessageType::Unsubscribe => {
302 let room = String::from_utf8_lossy(&msg.payload).to_string();
303 let ack = format!("unsubscribed:{}", room);
304 let response = WebSocketMessage {
305 msg_type: MessageType::System,
306 payload: ack.into_bytes(),
307 sender_id: None,
308 room_id: Some(room),
309 timestamp: current_timestamp(),
310 };
311 Ok(Some(response))
312 }
313 MessageType::Join => {
314 let room = String::from_utf8_lossy(&msg.payload).to_string();
315 let ack = format!("joined:{}", room);
316 let response = WebSocketMessage {
317 msg_type: MessageType::System,
318 payload: ack.into_bytes(),
319 sender_id: conn.user_id,
320 room_id: Some(room),
321 timestamp: current_timestamp(),
322 };
323 Ok(Some(response))
324 }
325 MessageType::Leave => {
326 let room = String::from_utf8_lossy(&msg.payload).to_string();
327 let ack = format!("left:{}", room);
328 let response = WebSocketMessage {
329 msg_type: MessageType::System,
330 payload: ack.into_bytes(),
331 sender_id: conn.user_id,
332 room_id: Some(room),
333 timestamp: current_timestamp(),
334 };
335 Ok(Some(response))
336 }
337 _ => Ok(None),
338 }
339 }
340
341 async fn on_connect(&self, conn: &WebSocketConnection) -> Result<(), WsError> {
342 self.connections
343 .write()
344 .await
345 .insert(conn.id.clone(), conn.clone());
346 Ok(())
347 }
348
349 async fn on_disconnect(&self, conn: &WebSocketConnection) {
350 self.connections.write().await.remove(&conn.id);
351 }
352
353 fn authenticate(&self, token: &str) -> Result<UserId, WsError> {
354 if let Some(id_str) = token.strip_prefix("user_id:") {
355 id_str.parse::<i64>().map_err(|_| {
356 WsError::Authentication(format!("invalid user_id in token: {}", token))
357 })
358 } else {
359 token
360 .parse::<i64>()
361 .map_err(|_| WsError::Authentication(format!("invalid token: {}", token)))
362 }
363 }
364}
365
366#[cfg(test)]
367mod tests {
368 use super::*;
369
370 fn make_conn(id: &str, user_id: Option<i64>) -> WebSocketConnection {
371 let mut conn = WebSocketConnection::new(id);
372 if let Some(uid) = user_id {
373 conn = conn.with_user(uid);
374 }
375 conn
376 }
377
378 fn make_text(payload: &[u8]) -> WebSocketMessage {
379 WebSocketMessage {
380 msg_type: MessageType::Text,
381 payload: payload.to_vec(),
382 sender_id: None,
383 room_id: None,
384 timestamp: 1000,
385 }
386 }
387
388 #[tokio::test]
389 async fn test_on_connect_tracks_connection() {
390 let handler = DefaultWebSocketHandler::new();
391 assert_eq!(handler.connection_count().await, 0);
392
393 let conn = make_conn("c1", Some(123));
394 handler.on_connect(&conn).await.unwrap();
395
396 assert_eq!(handler.connection_count().await, 1);
397 assert!(handler.is_connected("c1").await);
398 }
399
400 #[tokio::test]
401 async fn test_on_disconnect_removes_connection() {
402 let handler = DefaultWebSocketHandler::new();
403 let conn = make_conn("c1", Some(123));
404 handler.on_connect(&conn).await.unwrap();
405 assert!(handler.is_connected("c1").await);
406
407 handler.on_disconnect(&conn).await;
408 assert!(!handler.is_connected("c1").await);
409 assert_eq!(handler.connection_count().await, 0);
410 }
411
412 #[tokio::test]
413 async fn test_on_message_text_echoes_back() {
414 let handler = DefaultWebSocketHandler::new();
415 let conn = make_conn("c1", Some(42));
416 let msg = make_text(b"hello");
417
418 let response = handler.on_message(&conn, msg).await.unwrap();
419 assert!(response.is_some());
420
421 let resp = response.unwrap();
422 assert_eq!(resp.msg_type, MessageType::Text);
423 assert_eq!(resp.payload, b"hello");
424 assert_eq!(resp.sender_id, Some(42));
425 }
426
427 #[tokio::test]
428 async fn test_on_message_ping_responds_pong() {
429 let handler = DefaultWebSocketHandler::new();
430 let conn = make_conn("c1", None);
431 let msg = WebSocketMessage {
432 msg_type: MessageType::Ping,
433 payload: b"ping".to_vec(),
434 sender_id: None,
435 room_id: None,
436 timestamp: 1,
437 };
438
439 let response = handler.on_message(&conn, msg).await.unwrap();
440 let resp = response.unwrap();
441 assert_eq!(resp.msg_type, MessageType::Pong);
442 assert_eq!(resp.payload, b"ping");
443 }
444
445 #[tokio::test]
446 async fn test_on_message_subscribe_returns_system_ack() {
447 let handler = DefaultWebSocketHandler::new();
448 let conn = make_conn("c1", None);
449 let msg = WebSocketMessage {
450 msg_type: MessageType::Subscribe,
451 payload: b"room1".to_vec(),
452 sender_id: None,
453 room_id: None,
454 timestamp: 1,
455 };
456
457 let response = handler.on_message(&conn, msg).await.unwrap();
458 let resp = response.unwrap();
459 assert_eq!(resp.msg_type, MessageType::System);
460 assert_eq!(resp.payload, b"subscribed:room1");
461 assert_eq!(resp.room_id, Some("room1".to_string()));
462 }
463
464 #[tokio::test]
465 async fn test_on_message_unsubscribe_returns_system_ack() {
466 let handler = DefaultWebSocketHandler::new();
467 let conn = make_conn("c1", None);
468 let msg = WebSocketMessage {
469 msg_type: MessageType::Unsubscribe,
470 payload: b"room1".to_vec(),
471 sender_id: None,
472 room_id: None,
473 timestamp: 1,
474 };
475
476 let response = handler.on_message(&conn, msg).await.unwrap();
477 let resp = response.unwrap();
478 assert_eq!(resp.msg_type, MessageType::System);
479 assert_eq!(resp.payload, b"unsubscribed:room1");
480 }
481
482 #[tokio::test]
483 async fn test_on_message_join_returns_system_ack() {
484 let handler = DefaultWebSocketHandler::new();
485 let conn = make_conn("c1", Some(7));
486 let msg = WebSocketMessage {
487 msg_type: MessageType::Join,
488 payload: b"lobby".to_vec(),
489 sender_id: None,
490 room_id: None,
491 timestamp: 1,
492 };
493
494 let response = handler.on_message(&conn, msg).await.unwrap();
495 let resp = response.unwrap();
496 assert_eq!(resp.msg_type, MessageType::System);
497 assert_eq!(resp.payload, b"joined:lobby");
498 assert_eq!(resp.sender_id, Some(7));
499 assert_eq!(resp.room_id, Some("lobby".to_string()));
500 }
501
502 #[tokio::test]
503 async fn test_on_message_leave_returns_system_ack() {
504 let handler = DefaultWebSocketHandler::new();
505 let conn = make_conn("c1", Some(7));
506 let msg = WebSocketMessage {
507 msg_type: MessageType::Leave,
508 payload: b"lobby".to_vec(),
509 sender_id: None,
510 room_id: None,
511 timestamp: 1,
512 };
513
514 let response = handler.on_message(&conn, msg).await.unwrap();
515 let resp = response.unwrap();
516 assert_eq!(resp.payload, b"left:lobby");
517 }
518
519 #[tokio::test]
520 async fn test_on_message_binary_returns_none() {
521 let handler = DefaultWebSocketHandler::new();
522 let conn = make_conn("c1", None);
523 let msg = WebSocketMessage {
524 msg_type: MessageType::Binary,
525 payload: vec![1, 2, 3],
526 sender_id: None,
527 room_id: None,
528 timestamp: 1,
529 };
530
531 let response = handler.on_message(&conn, msg).await.unwrap();
532 assert!(response.is_none());
533 }
534
535 #[tokio::test]
536 async fn test_message_log_records_all_messages() {
537 let handler = DefaultWebSocketHandler::new();
538 let conn = make_conn("c1", None);
539
540 handler.on_message(&conn, make_text(b"m1")).await.unwrap();
541 handler
542 .on_message(
543 &conn,
544 WebSocketMessage {
545 msg_type: MessageType::Binary,
546 payload: vec![1],
547 sender_id: None,
548 room_id: None,
549 timestamp: 2,
550 },
551 )
552 .await
553 .unwrap();
554 handler.on_message(&conn, make_text(b"m3")).await.unwrap();
555
556 assert_eq!(handler.message_count().await, 3);
557 let msgs = handler.messages().await;
558 assert_eq!(msgs[0].payload, b"m1");
559 assert_eq!(msgs[1].msg_type, MessageType::Binary);
560 assert_eq!(msgs[2].payload, b"m3");
561 }
562
563 #[tokio::test]
564 async fn test_authenticate_valid_numeric_token() {
565 let handler = DefaultWebSocketHandler::new();
566 let user_id = handler.authenticate("12345").unwrap();
567 assert_eq!(user_id, 12345);
568 }
569
570 #[tokio::test]
571 async fn test_authenticate_valid_prefixed_token() {
572 let handler = DefaultWebSocketHandler::new();
573 let user_id = handler.authenticate("user_id:67890").unwrap();
574 assert_eq!(user_id, 67890);
575 }
576
577 #[test]
578 fn test_authenticate_invalid_token_returns_error() {
579 let handler = DefaultWebSocketHandler::new();
580 let result = handler.authenticate("not-a-number");
581 assert!(result.is_err());
582 assert!(matches!(result.unwrap_err(), WsError::Authentication(_)));
583 }
584
585 #[test]
586 fn test_authenticate_invalid_prefixed_token_returns_error() {
587 let handler = DefaultWebSocketHandler::new();
588 let result = handler.authenticate("user_id:abc");
589 assert!(result.is_err());
590 assert!(matches!(result.unwrap_err(), WsError::Authentication(_)));
591 }
592
593 #[tokio::test]
594 async fn test_get_connection_returns_stored_connection() {
595 let handler = DefaultWebSocketHandler::new();
596 let conn = make_conn("c1", Some(99));
597 handler.on_connect(&conn).await.unwrap();
598
599 let retrieved = handler.get_connection("c1").await.unwrap();
600 assert_eq!(retrieved.id, "c1");
601 assert_eq!(retrieved.user_id, Some(99));
602 assert!(retrieved.is_authenticated);
603 }
604
605 #[tokio::test]
606 async fn test_get_connection_not_found() {
607 let handler = DefaultWebSocketHandler::new();
608 assert!(handler.get_connection("missing").await.is_none());
609 }
610
611 #[tokio::test]
612 async fn test_multiple_connections_tracked_independently() {
613 let handler = DefaultWebSocketHandler::new();
614 let conn1 = make_conn("c1", Some(1));
615 let conn2 = make_conn("c2", Some(2));
616
617 handler.on_connect(&conn1).await.unwrap();
618 handler.on_connect(&conn2).await.unwrap();
619
620 assert_eq!(handler.connection_count().await, 2);
621
622 handler.on_disconnect(&conn1).await;
623 assert_eq!(handler.connection_count().await, 1);
624 assert!(!handler.is_connected("c1").await);
625 assert!(handler.is_connected("c2").await);
626 }
627
628 #[tokio::test]
629 async fn test_default_impl_creates_empty_handler() {
630 let handler = DefaultWebSocketHandler::default();
631 assert_eq!(handler.connection_count().await, 0);
632 assert_eq!(handler.message_count().await, 0);
633 }
634}