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>, 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 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 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); 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); 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 Text,
982 Continuation,
983 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 #[default]
1025 NormalClosure,
1026 GoingAway,
1028 ProtocolError,
1030 UnsupportedData,
1032 NoStatusReceived,
1034 AbnormalClosure,
1036 InvalidPayloadData,
1038 PolicyViolation,
1040 MessageTooBig,
1042 MandatoryExtension,
1044 InternalError,
1046 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 SendingDataFailed,
1130 Unknown,
1132 ThreadException,
1134 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}