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