Skip to main content

br_web_server/
websocket.rs

1use crate::request::Request;
2use crate::response::Response;
3use crate::{Handler, HttpError};
4use dashmap::DashMap;
5use flate2::read::DeflateDecoder;
6use flate2::write::DeflateEncoder;
7use flate2::Compression;
8use json::{object, JsonValue};
9use std::collections::HashSet;
10use std::io::{Read, Write};
11use std::sync::mpsc::{channel, Sender};
12use std::sync::Mutex;
13use std::{io, thread};
14
15const MAX_FRAME_SIZE: usize = 16 * 1024 * 1024;
16const MAX_DECOMPRESSED_SIZE: usize = 16 * 1024 * 1024;
17const MAX_CONTROL_FRAME_PAYLOAD: usize = 125;
18const MIN_COMPRESS_SIZE: usize = 64;
19
20pub static USERS: std::sync::LazyLock<DashMap<String, Websocket>> =
21    std::sync::LazyLock::new(DashMap::new);
22pub static WS_NOTICE: std::sync::LazyLock<Mutex<Vec<NoticeMsg>>> =
23    std::sync::LazyLock::new(|| Mutex::new(Vec::new()));
24pub static SUBSCRIPTIONS: std::sync::LazyLock<DashMap<String, HashSet<String>>> =
25    std::sync::LazyLock::new(DashMap::new);
26
27#[derive(Debug, Clone, Default)]
28pub struct DeflateConfig {
29    pub enabled: bool,
30    pub server_no_context_takeover: bool,
31    pub client_no_context_takeover: bool,
32}
33
34impl DeflateConfig {
35    pub fn from_header(header_value: &str) -> Option<Self> {
36        if !header_value.contains("permessage-deflate") {
37            return None;
38        }
39        let mut config = Self {
40            enabled: true,
41            server_no_context_takeover: true,
42            client_no_context_takeover: false,
43        };
44        for part in header_value.split(';').map(|s| s.trim()) {
45            if part.starts_with("server-no-context-takeover")
46                || part.starts_with("server_no_context_takeover")
47            {
48                config.server_no_context_takeover = true;
49            } else if part.starts_with("client-no-context-takeover")
50                || part.starts_with("client_no_context_takeover")
51            {
52                config.client_no_context_takeover = true;
53            }
54        }
55        if !config.client_no_context_takeover {
56            return None;
57        }
58        Some(config)
59    }
60
61    pub fn to_header_value(&self) -> String {
62        let mut parts = vec!["permessage-deflate".to_string()];
63        if self.server_no_context_takeover {
64            parts.push("server-no-context-takeover".to_string());
65        }
66        if self.client_no_context_takeover {
67            parts.push("client-no-context-takeover".to_string());
68        }
69        parts.join("; ")
70    }
71
72    pub fn decompress(&self, data: &[u8]) -> io::Result<Vec<u8>> {
73        let mut input = data.to_vec();
74        input.extend_from_slice(&[0x00, 0x00, 0xff, 0xff]);
75        let decoder = DeflateDecoder::new(&input[..]);
76        let mut limited = decoder.take(MAX_DECOMPRESSED_SIZE as u64 + 1);
77        let mut decompressed = Vec::new();
78        limited.read_to_end(&mut decompressed)?;
79        if decompressed.len() > MAX_DECOMPRESSED_SIZE {
80            return Err(io::Error::new(
81                io::ErrorKind::InvalidData,
82                "Decompressed payload exceeds maximum size",
83            ));
84        }
85        Ok(decompressed)
86    }
87
88    pub fn compress(&self, data: &[u8]) -> io::Result<Vec<u8>> {
89        let mut encoder = DeflateEncoder::new(Vec::new(), Compression::default());
90        encoder.write_all(data)?;
91        let mut compressed = encoder.finish()?;
92        if compressed.len() >= 4 && compressed[compressed.len() - 4..] == [0x00, 0x00, 0xff, 0xff] {
93            compressed.truncate(compressed.len() - 4);
94        }
95        Ok(compressed)
96    }
97}
98
99#[derive(Debug, Clone)]
100pub struct Websocket {
101    pub send: Option<Sender<Message>>,
102    pub key: String,
103    version: String,
104    request: Request,
105    response: Response,
106    pub deflate: DeflateConfig,
107}
108
109impl Websocket {
110    #[must_use]
111    pub fn http(request: Request, response: Response) -> Self {
112        Self {
113            send: None,
114            request,
115            key: String::new(),
116            version: String::new(),
117            response,
118            deflate: DeflateConfig::default(),
119        }
120    }
121    pub fn new(request: Request, response: Response) -> Self {
122        Self {
123            send: None,
124            request,
125            key: String::new(),
126            version: String::new(),
127            response,
128            deflate: DeflateConfig::default(),
129        }
130    }
131    pub fn send(&self, data: &JsonValue) {
132        let msg = Message {
133            mode: MessageMode::Server,
134            message_type: MessageType::Text,
135            payload: data.to_string().into_bytes(),
136            text: data.to_string(),
137            close: CloseCode::NormalClosure,
138            error: ErrorCode::None,
139        };
140        if let Some(sender) = &self.send {
141            if let Err(e) = sender.send(msg) {
142                log::warn!("WebSocket send failed: {:?}", e);
143            }
144        } else {
145            log::warn!("WebSocket send channel is None");
146        }
147    }
148
149    pub fn send_binary(&self, data: &[u8]) {
150        let msg = Message {
151            mode: MessageMode::Server,
152            message_type: MessageType::Binary,
153            payload: data.to_vec(),
154            text: String::new(),
155            close: CloseCode::NormalClosure,
156            error: ErrorCode::None,
157        };
158        if let Some(sender) = &self.send {
159            if let Err(e) = sender.send(msg) {
160                log::warn!("WebSocket send_binary failed: {:?}", e);
161            }
162        } else {
163            log::warn!("WebSocket send channel is None");
164        }
165    }
166
167    pub fn ping(&self, payload: &[u8]) {
168        let msg = Message {
169            mode: MessageMode::Server,
170            message_type: MessageType::Ping,
171            payload: payload.to_vec(),
172            text: String::new(),
173            close: CloseCode::NormalClosure,
174            error: ErrorCode::None,
175        };
176        if let Some(sender) = &self.send {
177            if let Err(e) = sender.send(msg) {
178                log::warn!("WebSocket ping failed: {:?}", e);
179            }
180        }
181    }
182
183    pub fn close(&self, code: CloseCode, reason: &str) {
184        let msg = Message {
185            mode: MessageMode::Server,
186            message_type: MessageType::Close,
187            payload: reason.as_bytes().to_vec(),
188            text: reason.to_string(),
189            close: code,
190            error: ErrorCode::None,
191        };
192        match &self.send {
193            Some(sender) => {
194                if let Err(e) = sender.send(msg) {
195                    log::warn!("WebSocket close failed: {:?}", e);
196                }
197            }
198            None => {
199                log::warn!("WebSocket close called but send channel is None");
200            }
201        }
202    }
203
204    pub fn online_users(&self) -> usize {
205        USERS.len()
206    }
207
208    pub fn is_connected(&self) -> bool {
209        self.send.is_some()
210    }
211    pub fn handle(&mut self) -> Result<(), HttpError> {
212        let (send, receive) = channel();
213        self.send = Some(send);
214        self.on_frame()?;
215        let mut factory = (self.response.factory)(self.clone());
216        USERS.insert(self.key.to_string(), self.clone());
217        factory.on_open()?;
218
219        let deflate = self.deflate.clone();
220        let key = self.key.clone();
221        let that = self.clone();
222
223        let split_result =
224            crate::stream::Scheme::split_for_websocket(&self.response.request.scheme);
225
226        match split_result {
227            Ok((mut reader, mut writer)) => {
228                let send_clone = self.send.clone();
229                let thr = thread::spawn(move || -> Result<(), HttpError> {
230                    loop {
231                        let msg = match reader.read_ws_data(&deflate) {
232                            Ok(e) => e,
233                            Err(_) => return Ok(()),
234                        };
235                        match msg.message_type {
236                            MessageType::TimeOut => continue,
237                            _ => match send_clone.clone() {
238                                Some(sender) => match sender.send(msg) {
239                                    Ok(()) => continue,
240                                    Err(_) => return Ok(()),
241                                },
242                                None => return Ok(()),
243                            },
244                        }
245                    }
246                });
247
248                let key_clone = key.clone();
249                thread::spawn(move || -> io::Result<()> {
250                    let mut factory = (that.response.factory)(that.clone());
251                    loop {
252                        match receive.recv() {
253                            Ok(msg) => match msg.message_type {
254                                MessageType::TimeOut => continue,
255                                MessageType::Close => {
256                                    factory.on_close(msg.close.clone(), &msg.text);
257                                    USERS.remove(&key_clone);
258                                    return Ok(());
259                                }
260                                MessageType::Pong => {}
261                                MessageType::Ping => {
262                                    let pong_frame = Message::send_pong(&msg.payload);
263                                    if let Err(e) = writer.write_all(&pong_frame) {
264                                        log::warn!("发送 Pong 失败: {:?}", e);
265                                    }
266                                }
267                                MessageType::Binary | MessageType::Text => match msg.mode {
268                                    MessageMode::Server => {
269                                        let frame = msg.clone().send_message(&that.deflate);
270                                        if let Err(e) = writer.write_all(&frame) {
271                                            log::warn!("发送消息失败: {:?}", e);
272                                        }
273                                    }
274                                    MessageMode::Client => {
275                                        if msg.message_type == MessageType::Text {
276                                            if let Ok(parsed) = json::parse(&msg.text) {
277                                                if parsed["type"] == "ping" {
278                                                    that.send(&object! {
279                                                        "type": "pong",
280                                                        "timestamp": parsed["timestamp"].clone()
281                                                    });
282                                                    continue;
283                                                }
284                                            }
285                                        }
286                                        if let Ok(()) = factory.on_message(msg) {};
287                                    }
288                                },
289                                MessageType::Error => continue,
290                                _ => continue,
291                            },
292                            Err(_) => {
293                                factory.on_close(CloseCode::AbnormalClosure, "连接异常断开");
294                                USERS.remove(&key_clone);
295                                return Ok(());
296                            }
297                        }
298                    }
299                });
300
301                if let Err(e) = thr.join() {
302                    log::warn!("WebSocket 线程异常退出: {:?}", e);
303                }
304            }
305            Err(_) => {
306                let scheme = self.response.request.scheme.clone();
307                let send_clone = self.send.clone();
308                let deflate_clone = deflate.clone();
309
310                let thr = thread::spawn(move || -> Result<(), HttpError> {
311                    loop {
312                        let mut guard = match scheme.lock() {
313                            Ok(g) => g,
314                            Err(_) => return Ok(()),
315                        };
316                        let msg = match guard.read_ws_data(&deflate_clone) {
317                            Ok(e) => e,
318                            Err(_) => return Ok(()),
319                        };
320                        drop(guard);
321                        match msg.message_type {
322                            MessageType::TimeOut => continue,
323                            _ => match send_clone.clone() {
324                                Some(sender) => match sender.send(msg) {
325                                    Ok(()) => continue,
326                                    Err(_) => return Ok(()),
327                                },
328                                None => return Ok(()),
329                            },
330                        }
331                    }
332                });
333
334                let scheme = self.response.request.scheme.clone();
335                let key_clone = key.clone();
336                thread::spawn(move || -> io::Result<()> {
337                    let mut factory = (that.response.factory)(that.clone());
338                    loop {
339                        match receive.recv() {
340                            Ok(msg) => match msg.message_type {
341                                MessageType::TimeOut => continue,
342                                MessageType::Close => {
343                                    factory.on_close(msg.close.clone(), &msg.text);
344                                    USERS.remove(&key_clone);
345                                    return Ok(());
346                                }
347                                MessageType::Pong => {}
348                                MessageType::Ping => {
349                                    let pong_frame = Message::send_pong(&msg.payload);
350                                    match scheme.lock() {
351                                        Ok(mut s) => {
352                                            if let Err(e) = s.write_all(&pong_frame) {
353                                                log::warn!("发送 Pong 失败: {:?}", e);
354                                            }
355                                        }
356                                        Err(e) => log::warn!("lock failed: {:?}", e),
357                                    }
358                                }
359                                MessageType::Binary | MessageType::Text => match msg.mode {
360                                    MessageMode::Server => {
361                                        let frame = msg.clone().send_message(&that.deflate);
362                                        match scheme.lock() {
363                                            Ok(mut s) => {
364                                                if let Err(e) = s.write_all(&frame) {
365                                                    log::warn!("发送消息失败: {:?}", e);
366                                                }
367                                            }
368                                            Err(e) => log::warn!("lock failed: {:?}", e),
369                                        }
370                                    }
371                                    MessageMode::Client => {
372                                        if msg.message_type == MessageType::Text {
373                                            if let Ok(parsed) = json::parse(&msg.text) {
374                                                if parsed["type"] == "ping" {
375                                                    that.send(&object! {
376                                                        "type": "pong",
377                                                        "timestamp": parsed["timestamp"].clone()
378                                                    });
379                                                    continue;
380                                                }
381                                            }
382                                        }
383                                        if let Ok(()) = factory.on_message(msg) {};
384                                    }
385                                },
386                                MessageType::Error => continue,
387                                _ => continue,
388                            },
389                            Err(_) => {
390                                factory.on_close(CloseCode::AbnormalClosure, "连接异常断开");
391                                USERS.remove(&key_clone);
392                                return Ok(());
393                            }
394                        }
395                    }
396                });
397
398                if let Err(e) = thr.join() {
399                    log::warn!("WebSocket 线程异常退出: {:?}", e);
400                }
401            }
402        }
403
404        Ok(())
405    }
406}
407impl Handler for Websocket {
408    fn on_request(&mut self, _request: Request, _response: &mut Response) {}
409    fn on_frame(&mut self) -> Result<(), HttpError> {
410        self.key = self.request.header["sec-websocket-key"]
411            .as_str()
412            .unwrap_or("")
413            .to_string();
414        self.version = self.request.header["sec-websocket-version"]
415            .as_str()
416            .unwrap_or("")
417            .to_string();
418
419        if self.key.is_empty() {
420            log::warn!("WebSocket 握手失败: 缺少 Sec-WebSocket-Key");
421            self.response.status(400).send()?;
422            return Err(HttpError::new(400, "Missing Sec-WebSocket-Key"));
423        }
424
425        if self.version != "13" {
426            log::warn!("WebSocket 版本不支持: {}", self.version);
427            self.response
428                .header("Sec-WebSocket-Version", "13")
429                .status(426)
430                .send()?;
431            return Err(HttpError::new(426, "Unsupported WebSocket version"));
432        }
433
434        let extensions = self.request.header["sec-websocket-extensions"]
435            .as_str()
436            .unwrap_or("");
437        if let Some(deflate_config) = DeflateConfig::from_header(extensions) {
438            self.deflate = deflate_config;
439        }
440
441        self.response.header("Upgrade", "websocket");
442        self.response.header("Connection", "Upgrade");
443        let sec_websocket_accept = br_crypto::sha1::encrypt_base64(
444            format!("{}258EAFA5-E914-47DA-95CA-C5AB0DC85B11", self.key).as_bytes(),
445        );
446        self.response
447            .header("Sec-WebSocket-Accept", sec_websocket_accept.as_str());
448
449        if self.deflate.enabled {
450            self.response
451                .header("Sec-WebSocket-Extensions", &self.deflate.to_header_value());
452        }
453
454        self.response.status(101).send()?;
455        self.response
456            .request
457            .scheme
458            .lock()
459            .map_err(|e| HttpError::new(500, &format!("lock: {}", e)))?
460            .flush()?;
461        Ok(())
462    }
463}
464#[derive(Debug, Clone)]
465pub struct Message {
466    pub mode: MessageMode,
467    pub message_type: MessageType,
468    pub payload: Vec<u8>, // 消息载荷,以字节向量形式表示
469    pub text: String,
470    pub close: CloseCode,
471    pub error: ErrorCode,
472}
473
474impl Message {
475    #[must_use]
476    pub fn msg_error() -> Self {
477        Message {
478            mode: MessageMode::Client,
479            message_type: MessageType::Error,
480            payload: vec![],
481            text: "长度不够".to_string(),
482            close: CloseCode::NormalClosure,
483            error: ErrorCode::SendingDataFailed,
484        }
485    }
486    // 解析WebSocket消息
487    pub fn parse_message(data: &mut Vec<u8>, deflate: &DeflateConfig) -> Message {
488        log::trace!("WebSocket parse_message: data.len()={}", data.len());
489
490        if data.len() < 2 {
491            return Message {
492                mode: MessageMode::Client,
493                message_type: MessageType::TimeOut,
494                payload: vec![],
495                text: String::new(),
496                close: CloseCode::NormalClosure,
497                error: ErrorCode::None,
498            };
499        }
500
501        let byte0 = data[0];
502        let byte1 = data[1];
503        let rsv2 = (byte0 & 0b0010_0000) != 0;
504        let rsv3 = (byte0 & 0b0001_0000) != 0;
505
506        if rsv2 || rsv3 {
507            log::warn!("WebSocket RSV2/RSV3 位非零, byte0={:#04x}", byte0);
508            data.clear();
509            return Message {
510                mode: MessageMode::Client,
511                message_type: MessageType::Error,
512                payload: vec![],
513                text: "RSV2/RSV3位必须为0".to_string(),
514                close: CloseCode::ProtocolError,
515                error: ErrorCode::SendingDataFailed,
516            };
517        }
518
519        let rsv1 = (byte0 & 0b0100_0000) != 0;
520        if rsv1 && !deflate.enabled {
521            log::warn!("WebSocket RSV1 位非零但未启用压缩");
522            data.clear();
523            return Message {
524                mode: MessageMode::Client,
525                message_type: MessageType::Error,
526                payload: vec![],
527                text: "RSV1位必须为0".to_string(),
528                close: CloseCode::ProtocolError,
529                error: ErrorCode::SendingDataFailed,
530            };
531        }
532
533        let len_flag = byte1 & 0b0111_1111;
534        let masked = (byte1 & 0b1000_0000) != 0;
535
536        let (ext_len_size, payload_length) = match len_flag {
537            0..=125 => (0usize, len_flag as usize),
538            126 => {
539                if data.len() < 4 {
540                    return Message {
541                        mode: MessageMode::Client,
542                        message_type: MessageType::TimeOut,
543                        payload: vec![],
544                        text: String::new(),
545                        close: CloseCode::NormalClosure,
546                        error: ErrorCode::None,
547                    };
548                }
549                (2usize, u16::from_be_bytes([data[2], data[3]]) as usize)
550            }
551            127 => {
552                if data.len() < 10 {
553                    return Message {
554                        mode: MessageMode::Client,
555                        message_type: MessageType::TimeOut,
556                        payload: vec![],
557                        text: String::new(),
558                        close: CloseCode::NormalClosure,
559                        error: ErrorCode::None,
560                    };
561                }
562                (
563                    8usize,
564                    u64::from_be_bytes([
565                        data[2], data[3], data[4], data[5], data[6], data[7], data[8], data[9],
566                    ]) as usize,
567                )
568            }
569            _ => {
570                data.clear();
571                return Message::msg_error();
572            }
573        };
574
575        if payload_length > MAX_FRAME_SIZE {
576            log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
577            data.clear();
578            return Message {
579                mode: MessageMode::Client,
580                message_type: MessageType::Error,
581                payload: vec![],
582                text: "消息过大".to_string(),
583                close: CloseCode::MessageTooBig,
584                error: ErrorCode::SendingDataFailed,
585            };
586        }
587
588        let mask_len = if masked { 4 } else { 0 };
589        let total_len = 2 + ext_len_size + mask_len + payload_length;
590
591        if data.len() < total_len {
592            log::trace!(
593                "WebSocket 数据不足: 需要 {} 字节, 当前 {} 字节",
594                total_len,
595                data.len()
596            );
597            return Message {
598                mode: MessageMode::Client,
599                message_type: MessageType::TimeOut,
600                payload: vec![],
601                text: String::new(),
602                close: CloseCode::NormalClosure,
603                error: ErrorCode::None,
604            };
605        }
606
607        let header = data.drain(..2).collect::<Vec<u8>>();
608
609        let rsv1 = (header[0] & 0b0100_0000) != 0;
610
611        let fin = (header[0] & 0b1000_0000) != 0;
612        let opcode = header[0] & 0b0000_1111;
613        let masked = (header[1] & 0b1000_0000) != 0;
614        let len_flag = header[1] & 0b0111_1111;
615        let mut payload_data = Vec::new();
616        let message_type = MessageType::from(opcode);
617        log::trace!(
618            "fin: {:#?} message_type: {:?} opcode: {} masked: {} len_flag: {} rsv1: {}",
619            fin,
620            message_type,
621            opcode,
622            masked,
623            len_flag,
624            rsv1
625        );
626        match message_type {
627            MessageType::Text => {
628                let payload_length = match len_flag {
629                    0..=125 => len_flag as usize,
630                    126 => {
631                        if data.len() < 2 {
632                            return Message {
633                                mode: MessageMode::Client,
634                                message_type: MessageType::Error,
635                                payload: vec![],
636                                text: "数据不足".to_string(),
637                                close: CloseCode::NormalClosure,
638                                error: ErrorCode::SendingDataFailed,
639                            };
640                        }
641                        let ext = data.drain(..2).collect::<Vec<u8>>();
642                        u16::from_be_bytes([ext[0], ext[1]]) as usize
643                    }
644                    127 => {
645                        if data.len() < 8 {
646                            return Message {
647                                mode: MessageMode::Client,
648                                message_type: MessageType::Error,
649                                payload: vec![],
650                                text: "数据不足".to_string(),
651                                close: CloseCode::NormalClosure,
652                                error: ErrorCode::SendingDataFailed,
653                            };
654                        }
655                        let ext = data.drain(..8).collect::<Vec<u8>>();
656                        u64::from_be_bytes([
657                            ext[0], ext[1], ext[2], ext[3], ext[4], ext[5], ext[6], ext[7],
658                        ]) as usize
659                    }
660                    _ => {
661                        return Message {
662                            mode: MessageMode::Client,
663                            message_type: MessageType::Error,
664                            payload: vec![],
665                            text: "数据格式错误".to_string(),
666                            close: CloseCode::NormalClosure,
667                            error: ErrorCode::SendingDataFailed,
668                        }
669                    }
670                };
671
672                if payload_length > MAX_FRAME_SIZE {
673                    log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
674                    return Message {
675                        mode: MessageMode::Client,
676                        message_type: MessageType::Error,
677                        payload: vec![],
678                        text: "消息过大".to_string(),
679                        close: CloseCode::MessageTooBig,
680                        error: ErrorCode::SendingDataFailed,
681                    };
682                }
683
684                if masked {
685                    // Need 4 bytes for mask key + payload_length bytes for payload
686                    if data.len() < payload_length + 4 {
687                        return Message {
688                            mode: MessageMode::Client,
689                            message_type,
690                            payload: payload_data,
691                            text: "继续加载".to_string(),
692                            close: CloseCode::NormalClosure,
693                            error: ErrorCode::None,
694                        };
695                    }
696                    let mask_key = data.drain(..4).collect::<Vec<u8>>();
697                    let payload = &data[..payload_length];
698                    for i in 0..payload.len() {
699                        payload_data.push(payload[i] ^ mask_key[i % 4]);
700                    }
701                    data.drain(..payload_length);
702                } else {
703                    if data.len() < payload_length {
704                        return Message {
705                            mode: MessageMode::Client,
706                            message_type,
707                            payload: payload_data,
708                            text: "继续加载".to_string(),
709                            close: CloseCode::NormalClosure,
710                            error: ErrorCode::None,
711                        };
712                    }
713                    let t = data.drain(..payload_length).collect::<Vec<u8>>();
714                    payload_data.extend_from_slice(&t);
715                }
716
717                let final_payload = if rsv1 && deflate.enabled {
718                    match deflate.decompress(&payload_data) {
719                        Ok(decompressed) => decompressed,
720                        Err(e) => {
721                            log::warn!("WebSocket 解压缩失败: {:?}", e);
722                            return Message {
723                                mode: MessageMode::Client,
724                                message_type: MessageType::Error,
725                                payload: vec![],
726                                text: "解压缩失败".to_string(),
727                                close: CloseCode::InvalidPayloadData,
728                                error: ErrorCode::SendingDataFailed,
729                            };
730                        }
731                    }
732                } else {
733                    payload_data
734                };
735
736                let text = String::from_utf8_lossy(&final_payload).into_owned();
737                Message {
738                    mode: MessageMode::Client,
739                    message_type,
740                    payload: final_payload,
741                    text: text.to_string(),
742                    close: CloseCode::NormalClosure,
743                    error: ErrorCode::None,
744                }
745            }
746            MessageType::Binary
747            | MessageType::Continuation
748            | MessageType::Close
749            | MessageType::Ping
750            | MessageType::Pong => {
751                let is_control_frame = matches!(
752                    message_type,
753                    MessageType::Close | MessageType::Ping | MessageType::Pong
754                );
755
756                if is_control_frame && len_flag > MAX_CONTROL_FRAME_PAYLOAD as u8 {
757                    log::warn!("控制帧 payload 超过 125 字节: {}", len_flag);
758                    return Message {
759                        mode: MessageMode::Client,
760                        message_type: MessageType::Error,
761                        payload: vec![],
762                        text: "控制帧过大".to_string(),
763                        close: CloseCode::ProtocolError,
764                        error: ErrorCode::SendingDataFailed,
765                    };
766                }
767
768                let payload_length = match len_flag {
769                    0..=125 => len_flag as usize,
770                    126 => {
771                        if data.len() < 2 {
772                            return Message::msg_error();
773                        }
774                        let ext = data.drain(..2).collect::<Vec<u8>>();
775                        u16::from_be_bytes([ext[0], ext[1]]) as usize
776                    }
777                    127 => {
778                        if data.len() < 8 {
779                            return Message::msg_error();
780                        }
781                        let ext = data.drain(..8).collect::<Vec<u8>>();
782                        u64::from_be_bytes([
783                            ext[0], ext[1], ext[2], ext[3], ext[4], ext[5], ext[6], ext[7],
784                        ]) as usize
785                    }
786                    _ => return Message::msg_error(),
787                };
788
789                if payload_length > MAX_FRAME_SIZE {
790                    log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
791                    return Message {
792                        mode: MessageMode::Client,
793                        message_type: MessageType::Error,
794                        payload: vec![],
795                        text: "消息过大".to_string(),
796                        close: CloseCode::MessageTooBig,
797                        error: ErrorCode::SendingDataFailed,
798                    };
799                }
800
801                if masked {
802                    if data.len() < payload_length + 4 {
803                        return Message::msg_error();
804                    }
805                    let mask_key = data.drain(..4).collect::<Vec<u8>>();
806                    let payload = &data[..payload_length];
807                    for i in 0..payload.len() {
808                        payload_data.push(payload[i] ^ mask_key[i % 4]);
809                    }
810                    data.drain(..payload_length);
811                } else if data.len() >= payload_length {
812                    let t = data.drain(..payload_length).collect::<Vec<u8>>();
813                    payload_data.extend_from_slice(&t);
814                }
815
816                let should_decompress = rsv1 && deflate.enabled && !is_control_frame;
817                let final_payload = if should_decompress {
818                    match deflate.decompress(&payload_data) {
819                        Ok(decompressed) => decompressed,
820                        Err(e) => {
821                            log::warn!("WebSocket 解压缩失败: {:?}", e);
822                            return Message {
823                                mode: MessageMode::Client,
824                                message_type: MessageType::Error,
825                                payload: vec![],
826                                text: "解压缩失败".to_string(),
827                                close: CloseCode::InvalidPayloadData,
828                                error: ErrorCode::SendingDataFailed,
829                            };
830                        }
831                    }
832                } else {
833                    payload_data
834                };
835
836                let (text, close) = match message_type {
837                    MessageType::Close => {
838                        let close_code = if final_payload.len() >= 2 {
839                            let code = u16::from_be_bytes([final_payload[0], final_payload[1]]);
840                            CloseCode::from(code)
841                        } else {
842                            CloseCode::NormalClosure
843                        };
844                        let reason = if final_payload.len() > 2 {
845                            String::from_utf8_lossy(&final_payload[2..]).into_owned()
846                        } else {
847                            "客户端关闭".to_string()
848                        };
849                        (reason, close_code)
850                    }
851                    MessageType::Ping => {
852                        let text = String::from_utf8_lossy(&final_payload).into_owned();
853                        (
854                            if text.is_empty() {
855                                "Ping".to_string()
856                            } else {
857                                text
858                            },
859                            CloseCode::NormalClosure,
860                        )
861                    }
862                    MessageType::Pong => {
863                        let text = String::from_utf8_lossy(&final_payload).into_owned();
864                        (
865                            if text.is_empty() {
866                                "Pong".to_string()
867                            } else {
868                                text
869                            },
870                            CloseCode::NormalClosure,
871                        )
872                    }
873                    MessageType::Binary => (String::new(), CloseCode::NormalClosure),
874                    MessageType::Continuation => ("继续加载".to_string(), CloseCode::NormalClosure),
875                    _ => (String::new(), CloseCode::NormalClosure),
876                };
877
878                Message {
879                    mode: MessageMode::Client,
880                    message_type,
881                    payload: final_payload,
882                    text,
883                    close,
884                    error: ErrorCode::None,
885                }
886            }
887            MessageType::Error => Message {
888                mode: MessageMode::Client,
889                message_type,
890                payload: vec![],
891                text: String::new(),
892                close: CloseCode::NormalClosure,
893                error: ErrorCode::Unknown,
894            },
895            MessageType::None => Message {
896                mode: MessageMode::Client,
897                message_type,
898                payload: vec![],
899                text: String::new(),
900                close: CloseCode::NormalClosure,
901                error: ErrorCode::None,
902            },
903            MessageType::TimeOut => Message {
904                mode: MessageMode::Client,
905                message_type,
906                payload: vec![],
907                text: String::new(),
908                close: CloseCode::NormalClosure,
909                error: ErrorCode::TimeOut,
910            },
911        }
912    }
913    pub fn send_message(&self, deflate: &DeflateConfig) -> Vec<u8> {
914        let mut frame = Vec::new();
915
916        let opcode = self.message_type.to_u8();
917
918        let should_compress = deflate.enabled
919            && matches!(self.message_type, MessageType::Text | MessageType::Binary)
920            && self.payload.len() > MIN_COMPRESS_SIZE;
921
922        let (payload, rsv1) = if should_compress {
923            match deflate.compress(&self.payload) {
924                Ok(compressed) if compressed.len() < self.payload.len() => (compressed, true),
925                _ => (self.payload.clone(), false),
926            }
927        } else {
928            (self.payload.clone(), false)
929        };
930
931        let byte1 = 0x80 | (if rsv1 { 0x40 } else { 0x00 }) | (opcode & 0x0F);
932        frame.push(byte1);
933
934        let payload_len = payload.len();
935        if payload_len < 126 {
936            frame.push(payload_len as u8);
937        } else if payload_len <= 65535 {
938            frame.push(126);
939            frame.extend_from_slice(&(payload_len as u16).to_be_bytes());
940        } else {
941            frame.push(127);
942            frame.extend_from_slice(&(payload_len as u64).to_be_bytes());
943        }
944        frame.extend_from_slice(&payload);
945        frame
946    }
947    #[must_use]
948    pub fn send_close(code: CloseCode, reason: &str) -> Vec<u8> {
949        let mut frame = Vec::new();
950        frame.push(0x88);
951        let reason_bytes = reason.as_bytes();
952        let reason_len = reason_bytes.len().min(123);
953        let payload_len = 2 + reason_len;
954        frame.push(payload_len as u8);
955        frame.extend(&code.to_u16().to_be_bytes());
956        frame.extend(&reason_bytes[..reason_len]);
957        frame
958    }
959    #[must_use]
960    pub fn send_pong(payload: &[u8]) -> Vec<u8> {
961        let mut frame = Vec::new();
962        frame.push(0x8A); // FIN=1, opcode=0xA (Pong)
963        let payload_len = payload.len().min(125);
964        frame.push(payload_len as u8);
965        frame.extend(&payload[..payload_len]);
966        frame
967    }
968    #[must_use]
969    pub fn send_ping(payload: &[u8]) -> Vec<u8> {
970        let mut frame = Vec::new();
971        frame.push(0x89); // FIN=1, opcode=0x9 (Ping)
972        let payload_len = payload.len().min(125);
973        frame.push(payload_len as u8);
974        frame.extend(&payload[..payload_len]);
975        frame
976    }
977}
978#[derive(Debug, Clone, Copy, PartialEq, Eq)]
979pub enum MessageType {
980    /// 文本
981    Text,
982    Continuation,
983    /// 客户端关闭
984    Close,
985    Binary,
986    Ping,
987    Pong,
988    None,
989    TimeOut,
990    Error,
991}
992
993impl MessageType {
994    #[must_use]
995    pub fn from(types: u8) -> Self {
996        match types {
997            0x0 => Self::Continuation,
998            0x1 => Self::Text,
999            0x2 => Self::Binary,
1000            0x8 => Self::Close,
1001            0x9 => Self::Ping,
1002            0xa => Self::Pong,
1003            _ => Self::None,
1004        }
1005    }
1006    #[must_use]
1007    pub fn to_u8(&self) -> u8 {
1008        match self {
1009            MessageType::Text => 0x1,
1010            MessageType::Continuation
1011            | MessageType::None
1012            | MessageType::Error
1013            | MessageType::TimeOut => 0x0,
1014            MessageType::Close => 0x8,
1015            MessageType::Binary => 0x2,
1016            MessageType::Ping => 0x9,
1017            MessageType::Pong => 0xa,
1018        }
1019    }
1020}
1021#[derive(Debug, Clone, PartialEq, Eq, Default)]
1022pub enum CloseCode {
1023    /// 1000: 正常关闭
1024    #[default]
1025    NormalClosure,
1026    /// 1001: 端点离开 (如服务器关闭或浏览器导航离开)
1027    GoingAway,
1028    /// 1002: 协议错误
1029    ProtocolError,
1030    /// 1003: 不支持的数据类型
1031    UnsupportedData,
1032    /// 1005: 未收到状态码 (保留,不应发送)
1033    NoStatusReceived,
1034    /// 1006: 异常关闭 (保留,不应发送)
1035    AbnormalClosure,
1036    /// 1007: 无效的帧 payload 数据 (如非 UTF-8 文本)
1037    InvalidPayloadData,
1038    /// 1008: 策略违规
1039    PolicyViolation,
1040    /// 1009: 消息过大
1041    MessageTooBig,
1042    /// 1010: 缺少必需的扩展
1043    MandatoryExtension,
1044    /// 1011: 内部服务器错误
1045    InternalError,
1046    /// 其它关闭码
1047    Other(u16),
1048}
1049
1050impl CloseCode {
1051    #[must_use]
1052    pub fn from_err(err: ErrorCode) -> CloseCode {
1053        match err {
1054            ErrorCode::SendingDataFailed => CloseCode::InternalError,
1055            ErrorCode::ThreadException => CloseCode::InternalError,
1056            ErrorCode::TimeOut => CloseCode::GoingAway,
1057            ErrorCode::Unknown | ErrorCode::None => CloseCode::NormalClosure,
1058        }
1059    }
1060
1061    #[must_use]
1062    pub fn str(&self) -> String {
1063        match self {
1064            CloseCode::NormalClosure => "正常关闭",
1065            CloseCode::GoingAway => "端点离开",
1066            CloseCode::ProtocolError => "协议错误",
1067            CloseCode::UnsupportedData => "不支持的数据",
1068            CloseCode::NoStatusReceived => "未收到状态码",
1069            CloseCode::AbnormalClosure => "异常关闭",
1070            CloseCode::InvalidPayloadData => "无效数据",
1071            CloseCode::PolicyViolation => "策略违规",
1072            CloseCode::MessageTooBig => "消息过大",
1073            CloseCode::MandatoryExtension => "缺少扩展",
1074            CloseCode::InternalError => "内部错误",
1075            CloseCode::Other(code) => return format!("关闭码 {code}"),
1076        }
1077        .to_string()
1078    }
1079
1080    #[must_use]
1081    pub fn to_u16(&self) -> u16 {
1082        match self {
1083            CloseCode::NormalClosure => 1000,
1084            CloseCode::GoingAway => 1001,
1085            CloseCode::ProtocolError => 1002,
1086            CloseCode::UnsupportedData => 1003,
1087            CloseCode::NoStatusReceived => 1005,
1088            CloseCode::AbnormalClosure => 1006,
1089            CloseCode::InvalidPayloadData => 1007,
1090            CloseCode::PolicyViolation => 1008,
1091            CloseCode::MessageTooBig => 1009,
1092            CloseCode::MandatoryExtension => 1010,
1093            CloseCode::InternalError => 1011,
1094            CloseCode::Other(code) => *code,
1095        }
1096    }
1097
1098    #[must_use]
1099    pub fn is_valid_for_send(&self) -> bool {
1100        !matches!(
1101            self,
1102            CloseCode::NoStatusReceived | CloseCode::AbnormalClosure
1103        )
1104    }
1105}
1106
1107impl From<u16> for CloseCode {
1108    fn from(code: u16) -> Self {
1109        match code {
1110            1000 => CloseCode::NormalClosure,
1111            1001 => CloseCode::GoingAway,
1112            1002 => CloseCode::ProtocolError,
1113            1003 => CloseCode::UnsupportedData,
1114            1005 => CloseCode::NoStatusReceived,
1115            1006 => CloseCode::AbnormalClosure,
1116            1007 => CloseCode::InvalidPayloadData,
1117            1008 => CloseCode::PolicyViolation,
1118            1009 => CloseCode::MessageTooBig,
1119            1010 => CloseCode::MandatoryExtension,
1120            1011 => CloseCode::InternalError,
1121            _ => CloseCode::Other(code),
1122        }
1123    }
1124}
1125
1126#[derive(Debug, Clone, Copy)]
1127pub enum ErrorCode {
1128    /// 发送数据失败
1129    SendingDataFailed,
1130    /// Unknown request error
1131    Unknown,
1132    /// 线程异常
1133    ThreadException,
1134    /// 超时
1135    TimeOut,
1136    None,
1137}
1138#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1139pub enum MessageMode {
1140    Client,
1141    Server,
1142}
1143
1144pub struct NoticeMsg {
1145    pub types: Types,
1146    pub msg: JsonValue,
1147    pub timestamp: i64,
1148    pub channel: String,
1149    pub user: String,
1150    pub org: String,
1151    pub admin: String,
1152}
1153
1154impl NoticeMsg {
1155    pub fn json(&mut self) -> JsonValue {
1156        object! {
1157            type:"notice",
1158            channel: self.channel.clone(),
1159            msg: self.msg.clone(),
1160            timestamp: self.timestamp,
1161        }
1162    }
1163
1164    fn now() -> i64 {
1165        std::time::SystemTime::now()
1166            .duration_since(std::time::UNIX_EPOCH)
1167            .unwrap_or_default()
1168            .as_secs() as i64
1169    }
1170
1171    fn push(notice: NoticeMsg) {
1172        if let Ok(mut queue) = WS_NOTICE.lock() {
1173            queue.push(notice);
1174        }
1175    }
1176
1177    pub fn to_all(channel: &str, msg: JsonValue) {
1178        Self::push(NoticeMsg {
1179            types: Types::All,
1180            msg,
1181            timestamp: Self::now(),
1182            channel: channel.to_string(),
1183            user: String::new(),
1184            org: String::new(),
1185            admin: String::new(),
1186        });
1187    }
1188
1189    pub fn to_org(channel: &str, msg: JsonValue, org: &str) {
1190        Self::push(NoticeMsg {
1191            types: Types::Org,
1192            msg,
1193            timestamp: Self::now(),
1194            channel: channel.to_string(),
1195            user: String::new(),
1196            org: org.to_string(),
1197            admin: String::new(),
1198        });
1199    }
1200
1201    pub fn to_user(channel: &str, msg: JsonValue, user: &str) {
1202        Self::push(NoticeMsg {
1203            types: Types::User,
1204            msg,
1205            timestamp: Self::now(),
1206            channel: channel.to_string(),
1207            user: user.to_string(),
1208            org: String::new(),
1209            admin: String::new(),
1210        });
1211    }
1212
1213    pub fn to_admin(channel: &str, msg: JsonValue, admin: &str) {
1214        Self::push(NoticeMsg {
1215            types: Types::Admin,
1216            msg,
1217            timestamp: Self::now(),
1218            channel: channel.to_string(),
1219            user: String::new(),
1220            org: String::new(),
1221            admin: admin.to_string(),
1222        });
1223    }
1224
1225    pub fn to_channel(channel: &str, msg: JsonValue) {
1226        Self::push(NoticeMsg {
1227            types: Types::Channel,
1228            msg,
1229            timestamp: Self::now(),
1230            channel: channel.to_string(),
1231            user: String::new(),
1232            org: String::new(),
1233            admin: String::new(),
1234        });
1235    }
1236}
1237
1238pub enum Types {
1239    All,
1240    User,
1241    Org,
1242    Admin,
1243    Channel,
1244}