Skip to main content

sz_orm_websocket/
handler.rs

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
198/// Default WebSocket handler that tracks connections and echoes messages.
199///
200/// - Text messages are echoed back to the sender with the sender's user_id.
201/// - Ping messages are answered with a Pong.
202/// - Subscribe/Unsubscribe/Join/Leave messages return a System acknowledgement.
203/// - All other messages are logged but produce no response.
204///
205/// # 安全说明(v4.8.0 修复 M-1)
206///
207/// `message_log` 为有界环形缓冲(上限 [`MAX_MESSAGE_LOG`],默认 10_000 条),
208/// 修复前无界增长——未认证客户端持续发消息可耗尽内存(黑帽审计实证)。
209pub struct DefaultWebSocketHandler {
210    connections: Arc<RwLock<HashMap<String, WebSocketConnection>>>,
211    message_log: Arc<RwLock<std::collections::VecDeque<WebSocketMessage>>>,
212}
213
214/// 消息日志上限(v4.8.0 修复 M-1):超出后丢弃最旧消息
215const 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        // M-1 修复:有界日志——超出上限丢弃最旧条目,防内存无界增长
260        {
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}