use crate::request::Request;
use crate::response::Response;
use crate::{Handler, HttpError};
use dashmap::DashMap;
use flate2::read::DeflateDecoder;
use flate2::write::DeflateEncoder;
use flate2::Compression;
use json::{object, JsonValue};
use std::collections::HashSet;
use std::io::{Read, Write};
use std::sync::mpsc::{channel, Sender};
use std::sync::Mutex;
use std::{io, thread};
const MAX_FRAME_SIZE: usize = 16 * 1024 * 1024;
const MAX_DECOMPRESSED_SIZE: usize = 16 * 1024 * 1024;
const MAX_CONTROL_FRAME_PAYLOAD: usize = 125;
const MIN_COMPRESS_SIZE: usize = 64;
pub static USERS: std::sync::LazyLock<DashMap<String, Websocket>> =
std::sync::LazyLock::new(DashMap::new);
pub static WS_NOTICE: std::sync::LazyLock<Mutex<Vec<NoticeMsg>>> =
std::sync::LazyLock::new(|| Mutex::new(Vec::new()));
pub static SUBSCRIPTIONS: std::sync::LazyLock<DashMap<String, HashSet<String>>> =
std::sync::LazyLock::new(DashMap::new);
#[derive(Debug, Clone, Default)]
pub struct DeflateConfig {
pub enabled: bool,
pub server_no_context_takeover: bool,
pub client_no_context_takeover: bool,
}
impl DeflateConfig {
pub fn from_header(header_value: &str) -> Option<Self> {
if !header_value.contains("permessage-deflate") {
return None;
}
let mut config = Self {
enabled: true,
server_no_context_takeover: true,
client_no_context_takeover: false,
};
for part in header_value.split(';').map(|s| s.trim()) {
if part.starts_with("server-no-context-takeover")
|| part.starts_with("server_no_context_takeover")
{
config.server_no_context_takeover = true;
} else if part.starts_with("client-no-context-takeover")
|| part.starts_with("client_no_context_takeover")
{
config.client_no_context_takeover = true;
}
}
if !config.client_no_context_takeover {
return None;
}
Some(config)
}
pub fn to_header_value(&self) -> String {
let mut parts = vec!["permessage-deflate".to_string()];
if self.server_no_context_takeover {
parts.push("server-no-context-takeover".to_string());
}
if self.client_no_context_takeover {
parts.push("client-no-context-takeover".to_string());
}
parts.join("; ")
}
pub fn decompress(&self, data: &[u8]) -> io::Result<Vec<u8>> {
let mut input = data.to_vec();
input.extend_from_slice(&[0x00, 0x00, 0xff, 0xff]);
let decoder = DeflateDecoder::new(&input[..]);
let mut limited = decoder.take(MAX_DECOMPRESSED_SIZE as u64 + 1);
let mut decompressed = Vec::new();
limited.read_to_end(&mut decompressed)?;
if decompressed.len() > MAX_DECOMPRESSED_SIZE {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Decompressed payload exceeds maximum size",
));
}
Ok(decompressed)
}
pub fn compress(&self, data: &[u8]) -> io::Result<Vec<u8>> {
let mut encoder = DeflateEncoder::new(Vec::new(), Compression::default());
encoder.write_all(data)?;
let mut compressed = encoder.finish()?;
if compressed.len() >= 4 && compressed[compressed.len() - 4..] == [0x00, 0x00, 0xff, 0xff] {
compressed.truncate(compressed.len() - 4);
}
Ok(compressed)
}
}
#[derive(Debug, Clone)]
pub struct Websocket {
pub send: Option<Sender<Message>>,
pub key: String,
version: String,
request: Request,
response: Response,
pub deflate: DeflateConfig,
}
impl Websocket {
#[must_use]
pub fn http(request: Request, response: Response) -> Self {
Self {
send: None,
request,
key: String::new(),
version: String::new(),
response,
deflate: DeflateConfig::default(),
}
}
pub fn new(request: Request, response: Response) -> Self {
Self {
send: None,
request,
key: String::new(),
version: String::new(),
response,
deflate: DeflateConfig::default(),
}
}
pub fn send(&self, data: &JsonValue) {
let msg = Message {
mode: MessageMode::Server,
message_type: MessageType::Text,
payload: data.to_string().into_bytes(),
text: data.to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
if let Some(sender) = &self.send {
if let Err(e) = sender.send(msg) {
log::warn!("WebSocket send failed: {:?}", e);
}
} else {
log::warn!("WebSocket send channel is None");
}
}
pub fn send_binary(&self, data: &[u8]) {
let msg = Message {
mode: MessageMode::Server,
message_type: MessageType::Binary,
payload: data.to_vec(),
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
if let Some(sender) = &self.send {
if let Err(e) = sender.send(msg) {
log::warn!("WebSocket send_binary failed: {:?}", e);
}
} else {
log::warn!("WebSocket send channel is None");
}
}
pub fn ping(&self, payload: &[u8]) {
let msg = Message {
mode: MessageMode::Server,
message_type: MessageType::Ping,
payload: payload.to_vec(),
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
if let Some(sender) = &self.send {
if let Err(e) = sender.send(msg) {
log::warn!("WebSocket ping failed: {:?}", e);
}
}
}
pub fn close(&self, code: CloseCode, reason: &str) {
let msg = Message {
mode: MessageMode::Server,
message_type: MessageType::Close,
payload: reason.as_bytes().to_vec(),
text: reason.to_string(),
close: code,
error: ErrorCode::None,
};
match &self.send {
Some(sender) => {
if let Err(e) = sender.send(msg) {
log::warn!("WebSocket close failed: {:?}", e);
}
}
None => {
log::warn!("WebSocket close called but send channel is None");
}
}
}
pub fn online_users(&self) -> usize {
USERS.len()
}
pub fn is_connected(&self) -> bool {
self.send.is_some()
}
pub fn handle(&mut self) -> Result<(), HttpError> {
let (send, receive) = channel();
self.send = Some(send);
self.on_frame()?;
let mut factory = (self.response.factory)(self.clone());
USERS.insert(self.key.to_string(), self.clone());
factory.on_open()?;
let deflate = self.deflate.clone();
let key = self.key.clone();
let that = self.clone();
let split_result =
crate::stream::Scheme::split_for_websocket(&self.response.request.scheme);
match split_result {
Ok((mut reader, mut writer)) => {
let send_clone = self.send.clone();
let thr = thread::spawn(move || -> Result<(), HttpError> {
loop {
let msg = match reader.read_ws_data(&deflate) {
Ok(e) => e,
Err(_) => return Ok(()),
};
match msg.message_type {
MessageType::TimeOut => continue,
_ => match send_clone.clone() {
Some(sender) => match sender.send(msg) {
Ok(()) => continue,
Err(_) => return Ok(()),
},
None => return Ok(()),
},
}
}
});
let key_clone = key.clone();
thread::spawn(move || -> io::Result<()> {
let mut factory = (that.response.factory)(that.clone());
loop {
match receive.recv() {
Ok(msg) => match msg.message_type {
MessageType::TimeOut => continue,
MessageType::Close => {
factory.on_close(msg.close.clone(), &msg.text);
USERS.remove(&key_clone);
return Ok(());
}
MessageType::Pong => {}
MessageType::Ping => {
let pong_frame = Message::send_pong(&msg.payload);
if let Err(e) = writer.write_all(&pong_frame) {
log::warn!("发送 Pong 失败: {:?}", e);
}
}
MessageType::Binary | MessageType::Text => match msg.mode {
MessageMode::Server => {
let frame = msg.clone().send_message(&that.deflate);
if let Err(e) = writer.write_all(&frame) {
log::warn!("发送消息失败: {:?}", e);
}
}
MessageMode::Client => {
if msg.message_type == MessageType::Text {
if let Ok(parsed) = json::parse(&msg.text) {
if parsed["type"] == "ping" {
that.send(&object! {
"type": "pong",
"timestamp": parsed["timestamp"].clone()
});
continue;
}
}
}
if let Ok(()) = factory.on_message(msg) {};
}
},
MessageType::Error => continue,
_ => continue,
},
Err(_) => {
factory.on_close(CloseCode::AbnormalClosure, "连接异常断开");
USERS.remove(&key_clone);
return Ok(());
}
}
}
});
if let Err(e) = thr.join() {
log::warn!("WebSocket 线程异常退出: {:?}", e);
}
}
Err(_) => {
let scheme = self.response.request.scheme.clone();
let send_clone = self.send.clone();
let deflate_clone = deflate.clone();
let thr = thread::spawn(move || -> Result<(), HttpError> {
loop {
let mut guard = match scheme.lock() {
Ok(g) => g,
Err(_) => return Ok(()),
};
let msg = match guard.read_ws_data(&deflate_clone) {
Ok(e) => e,
Err(_) => return Ok(()),
};
drop(guard);
match msg.message_type {
MessageType::TimeOut => continue,
_ => match send_clone.clone() {
Some(sender) => match sender.send(msg) {
Ok(()) => continue,
Err(_) => return Ok(()),
},
None => return Ok(()),
},
}
}
});
let scheme = self.response.request.scheme.clone();
let key_clone = key.clone();
thread::spawn(move || -> io::Result<()> {
let mut factory = (that.response.factory)(that.clone());
loop {
match receive.recv() {
Ok(msg) => match msg.message_type {
MessageType::TimeOut => continue,
MessageType::Close => {
factory.on_close(msg.close.clone(), &msg.text);
USERS.remove(&key_clone);
return Ok(());
}
MessageType::Pong => {}
MessageType::Ping => {
let pong_frame = Message::send_pong(&msg.payload);
match scheme.lock() {
Ok(mut s) => {
if let Err(e) = s.write_all(&pong_frame) {
log::warn!("发送 Pong 失败: {:?}", e);
}
}
Err(e) => log::warn!("lock failed: {:?}", e),
}
}
MessageType::Binary | MessageType::Text => match msg.mode {
MessageMode::Server => {
let frame = msg.clone().send_message(&that.deflate);
match scheme.lock() {
Ok(mut s) => {
if let Err(e) = s.write_all(&frame) {
log::warn!("发送消息失败: {:?}", e);
}
}
Err(e) => log::warn!("lock failed: {:?}", e),
}
}
MessageMode::Client => {
if msg.message_type == MessageType::Text {
if let Ok(parsed) = json::parse(&msg.text) {
if parsed["type"] == "ping" {
that.send(&object! {
"type": "pong",
"timestamp": parsed["timestamp"].clone()
});
continue;
}
}
}
if let Ok(()) = factory.on_message(msg) {};
}
},
MessageType::Error => continue,
_ => continue,
},
Err(_) => {
factory.on_close(CloseCode::AbnormalClosure, "连接异常断开");
USERS.remove(&key_clone);
return Ok(());
}
}
}
});
if let Err(e) = thr.join() {
log::warn!("WebSocket 线程异常退出: {:?}", e);
}
}
}
Ok(())
}
}
impl Handler for Websocket {
fn on_request(&mut self, _request: Request, _response: &mut Response) {}
fn on_frame(&mut self) -> Result<(), HttpError> {
self.key = self.request.header["sec-websocket-key"]
.as_str()
.unwrap_or("")
.to_string();
self.version = self.request.header["sec-websocket-version"]
.as_str()
.unwrap_or("")
.to_string();
if self.key.is_empty() {
log::warn!("WebSocket 握手失败: 缺少 Sec-WebSocket-Key");
self.response.status(400).send()?;
return Err(HttpError::new(400, "Missing Sec-WebSocket-Key"));
}
if self.version != "13" {
log::warn!("WebSocket 版本不支持: {}", self.version);
self.response
.header("Sec-WebSocket-Version", "13")
.status(426)
.send()?;
return Err(HttpError::new(426, "Unsupported WebSocket version"));
}
let extensions = self.request.header["sec-websocket-extensions"]
.as_str()
.unwrap_or("");
if let Some(deflate_config) = DeflateConfig::from_header(extensions) {
self.deflate = deflate_config;
}
self.response.header("Upgrade", "websocket");
self.response.header("Connection", "Upgrade");
let sec_websocket_accept = br_crypto::sha1::encrypt_base64(
format!("{}258EAFA5-E914-47DA-95CA-C5AB0DC85B11", self.key).as_bytes(),
);
self.response
.header("Sec-WebSocket-Accept", sec_websocket_accept.as_str());
if self.deflate.enabled {
self.response
.header("Sec-WebSocket-Extensions", &self.deflate.to_header_value());
}
self.response.status(101).send()?;
self.response
.request
.scheme
.lock()
.map_err(|e| HttpError::new(500, &format!("lock: {}", e)))?
.flush()?;
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Message {
pub mode: MessageMode,
pub message_type: MessageType,
pub payload: Vec<u8>, pub text: String,
pub close: CloseCode,
pub error: ErrorCode,
}
impl Message {
#[must_use]
pub fn msg_error() -> Self {
Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "长度不够".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::SendingDataFailed,
}
}
pub fn parse_message(data: &mut Vec<u8>, deflate: &DeflateConfig) -> Message {
log::trace!("WebSocket parse_message: data.len()={}", data.len());
if data.len() < 2 {
return Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
let byte0 = data[0];
let byte1 = data[1];
let rsv2 = (byte0 & 0b0010_0000) != 0;
let rsv3 = (byte0 & 0b0001_0000) != 0;
if rsv2 || rsv3 {
log::warn!("WebSocket RSV2/RSV3 位非零, byte0={:#04x}", byte0);
data.clear();
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "RSV2/RSV3位必须为0".to_string(),
close: CloseCode::ProtocolError,
error: ErrorCode::SendingDataFailed,
};
}
let rsv1 = (byte0 & 0b0100_0000) != 0;
if rsv1 && !deflate.enabled {
log::warn!("WebSocket RSV1 位非零但未启用压缩");
data.clear();
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "RSV1位必须为0".to_string(),
close: CloseCode::ProtocolError,
error: ErrorCode::SendingDataFailed,
};
}
let len_flag = byte1 & 0b0111_1111;
let masked = (byte1 & 0b1000_0000) != 0;
let (ext_len_size, payload_length) = match len_flag {
0..=125 => (0usize, len_flag as usize),
126 => {
if data.len() < 4 {
return Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
(2usize, u16::from_be_bytes([data[2], data[3]]) as usize)
}
127 => {
if data.len() < 10 {
return Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
(
8usize,
u64::from_be_bytes([
data[2], data[3], data[4], data[5], data[6], data[7], data[8], data[9],
]) as usize,
)
}
_ => {
data.clear();
return Message::msg_error();
}
};
if payload_length > MAX_FRAME_SIZE {
log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
data.clear();
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "消息过大".to_string(),
close: CloseCode::MessageTooBig,
error: ErrorCode::SendingDataFailed,
};
}
let mask_len = if masked { 4 } else { 0 };
let total_len = 2 + ext_len_size + mask_len + payload_length;
if data.len() < total_len {
log::trace!(
"WebSocket 数据不足: 需要 {} 字节, 当前 {} 字节",
total_len,
data.len()
);
return Message {
mode: MessageMode::Client,
message_type: MessageType::TimeOut,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
let header = data.drain(..2).collect::<Vec<u8>>();
let rsv1 = (header[0] & 0b0100_0000) != 0;
let fin = (header[0] & 0b1000_0000) != 0;
let opcode = header[0] & 0b0000_1111;
let masked = (header[1] & 0b1000_0000) != 0;
let len_flag = header[1] & 0b0111_1111;
let mut payload_data = Vec::new();
let message_type = MessageType::from(opcode);
log::trace!(
"fin: {:#?} message_type: {:?} opcode: {} masked: {} len_flag: {} rsv1: {}",
fin,
message_type,
opcode,
masked,
len_flag,
rsv1
);
match message_type {
MessageType::Text => {
let payload_length = match len_flag {
0..=125 => len_flag as usize,
126 => {
if data.len() < 2 {
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "数据不足".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::SendingDataFailed,
};
}
let ext = data.drain(..2).collect::<Vec<u8>>();
u16::from_be_bytes([ext[0], ext[1]]) as usize
}
127 => {
if data.len() < 8 {
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "数据不足".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::SendingDataFailed,
};
}
let ext = data.drain(..8).collect::<Vec<u8>>();
u64::from_be_bytes([
ext[0], ext[1], ext[2], ext[3], ext[4], ext[5], ext[6], ext[7],
]) as usize
}
_ => {
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "数据格式错误".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::SendingDataFailed,
}
}
};
if payload_length > MAX_FRAME_SIZE {
log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "消息过大".to_string(),
close: CloseCode::MessageTooBig,
error: ErrorCode::SendingDataFailed,
};
}
if masked {
if data.len() < payload_length + 4 {
return Message {
mode: MessageMode::Client,
message_type,
payload: payload_data,
text: "继续加载".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
let mask_key = data.drain(..4).collect::<Vec<u8>>();
let payload = &data[..payload_length];
for i in 0..payload.len() {
payload_data.push(payload[i] ^ mask_key[i % 4]);
}
data.drain(..payload_length);
} else {
if data.len() < payload_length {
return Message {
mode: MessageMode::Client,
message_type,
payload: payload_data,
text: "继续加载".to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
};
}
let t = data.drain(..payload_length).collect::<Vec<u8>>();
payload_data.extend_from_slice(&t);
}
let final_payload = if rsv1 && deflate.enabled {
match deflate.decompress(&payload_data) {
Ok(decompressed) => decompressed,
Err(e) => {
log::warn!("WebSocket 解压缩失败: {:?}", e);
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "解压缩失败".to_string(),
close: CloseCode::InvalidPayloadData,
error: ErrorCode::SendingDataFailed,
};
}
}
} else {
payload_data
};
let text = String::from_utf8_lossy(&final_payload).into_owned();
Message {
mode: MessageMode::Client,
message_type,
payload: final_payload,
text: text.to_string(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
}
}
MessageType::Binary
| MessageType::Continuation
| MessageType::Close
| MessageType::Ping
| MessageType::Pong => {
let is_control_frame = matches!(
message_type,
MessageType::Close | MessageType::Ping | MessageType::Pong
);
if is_control_frame && len_flag > MAX_CONTROL_FRAME_PAYLOAD as u8 {
log::warn!("控制帧 payload 超过 125 字节: {}", len_flag);
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "控制帧过大".to_string(),
close: CloseCode::ProtocolError,
error: ErrorCode::SendingDataFailed,
};
}
let payload_length = match len_flag {
0..=125 => len_flag as usize,
126 => {
if data.len() < 2 {
return Message::msg_error();
}
let ext = data.drain(..2).collect::<Vec<u8>>();
u16::from_be_bytes([ext[0], ext[1]]) as usize
}
127 => {
if data.len() < 8 {
return Message::msg_error();
}
let ext = data.drain(..8).collect::<Vec<u8>>();
u64::from_be_bytes([
ext[0], ext[1], ext[2], ext[3], ext[4], ext[5], ext[6], ext[7],
]) as usize
}
_ => return Message::msg_error(),
};
if payload_length > MAX_FRAME_SIZE {
log::warn!("帧大小超过限制: {} > {}", payload_length, MAX_FRAME_SIZE);
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "消息过大".to_string(),
close: CloseCode::MessageTooBig,
error: ErrorCode::SendingDataFailed,
};
}
if masked {
if data.len() < payload_length + 4 {
return Message::msg_error();
}
let mask_key = data.drain(..4).collect::<Vec<u8>>();
let payload = &data[..payload_length];
for i in 0..payload.len() {
payload_data.push(payload[i] ^ mask_key[i % 4]);
}
data.drain(..payload_length);
} else if data.len() >= payload_length {
let t = data.drain(..payload_length).collect::<Vec<u8>>();
payload_data.extend_from_slice(&t);
}
let should_decompress = rsv1 && deflate.enabled && !is_control_frame;
let final_payload = if should_decompress {
match deflate.decompress(&payload_data) {
Ok(decompressed) => decompressed,
Err(e) => {
log::warn!("WebSocket 解压缩失败: {:?}", e);
return Message {
mode: MessageMode::Client,
message_type: MessageType::Error,
payload: vec![],
text: "解压缩失败".to_string(),
close: CloseCode::InvalidPayloadData,
error: ErrorCode::SendingDataFailed,
};
}
}
} else {
payload_data
};
let (text, close) = match message_type {
MessageType::Close => {
let close_code = if final_payload.len() >= 2 {
let code = u16::from_be_bytes([final_payload[0], final_payload[1]]);
CloseCode::from(code)
} else {
CloseCode::NormalClosure
};
let reason = if final_payload.len() > 2 {
String::from_utf8_lossy(&final_payload[2..]).into_owned()
} else {
"客户端关闭".to_string()
};
(reason, close_code)
}
MessageType::Ping => {
let text = String::from_utf8_lossy(&final_payload).into_owned();
(
if text.is_empty() {
"Ping".to_string()
} else {
text
},
CloseCode::NormalClosure,
)
}
MessageType::Pong => {
let text = String::from_utf8_lossy(&final_payload).into_owned();
(
if text.is_empty() {
"Pong".to_string()
} else {
text
},
CloseCode::NormalClosure,
)
}
MessageType::Binary => (String::new(), CloseCode::NormalClosure),
MessageType::Continuation => ("继续加载".to_string(), CloseCode::NormalClosure),
_ => (String::new(), CloseCode::NormalClosure),
};
Message {
mode: MessageMode::Client,
message_type,
payload: final_payload,
text,
close,
error: ErrorCode::None,
}
}
MessageType::Error => Message {
mode: MessageMode::Client,
message_type,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::Unknown,
},
MessageType::None => Message {
mode: MessageMode::Client,
message_type,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::None,
},
MessageType::TimeOut => Message {
mode: MessageMode::Client,
message_type,
payload: vec![],
text: String::new(),
close: CloseCode::NormalClosure,
error: ErrorCode::TimeOut,
},
}
}
pub fn send_message(&self, deflate: &DeflateConfig) -> Vec<u8> {
let mut frame = Vec::new();
let opcode = self.message_type.to_u8();
let should_compress = deflate.enabled
&& matches!(self.message_type, MessageType::Text | MessageType::Binary)
&& self.payload.len() > MIN_COMPRESS_SIZE;
let (payload, rsv1) = if should_compress {
match deflate.compress(&self.payload) {
Ok(compressed) if compressed.len() < self.payload.len() => (compressed, true),
_ => (self.payload.clone(), false),
}
} else {
(self.payload.clone(), false)
};
let byte1 = 0x80 | (if rsv1 { 0x40 } else { 0x00 }) | (opcode & 0x0F);
frame.push(byte1);
let payload_len = payload.len();
if payload_len < 126 {
frame.push(payload_len as u8);
} else if payload_len <= 65535 {
frame.push(126);
frame.extend_from_slice(&(payload_len as u16).to_be_bytes());
} else {
frame.push(127);
frame.extend_from_slice(&(payload_len as u64).to_be_bytes());
}
frame.extend_from_slice(&payload);
frame
}
#[must_use]
pub fn send_close(code: CloseCode, reason: &str) -> Vec<u8> {
let mut frame = Vec::new();
frame.push(0x88);
let reason_bytes = reason.as_bytes();
let reason_len = reason_bytes.len().min(123);
let payload_len = 2 + reason_len;
frame.push(payload_len as u8);
frame.extend(&code.to_u16().to_be_bytes());
frame.extend(&reason_bytes[..reason_len]);
frame
}
#[must_use]
pub fn send_pong(payload: &[u8]) -> Vec<u8> {
let mut frame = Vec::new();
frame.push(0x8A); let payload_len = payload.len().min(125);
frame.push(payload_len as u8);
frame.extend(&payload[..payload_len]);
frame
}
#[must_use]
pub fn send_ping(payload: &[u8]) -> Vec<u8> {
let mut frame = Vec::new();
frame.push(0x89); let payload_len = payload.len().min(125);
frame.push(payload_len as u8);
frame.extend(&payload[..payload_len]);
frame
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MessageType {
Text,
Continuation,
Close,
Binary,
Ping,
Pong,
None,
TimeOut,
Error,
}
impl MessageType {
#[must_use]
pub fn from(types: u8) -> Self {
match types {
0x0 => Self::Continuation,
0x1 => Self::Text,
0x2 => Self::Binary,
0x8 => Self::Close,
0x9 => Self::Ping,
0xa => Self::Pong,
_ => Self::None,
}
}
#[must_use]
pub fn to_u8(&self) -> u8 {
match self {
MessageType::Text => 0x1,
MessageType::Continuation
| MessageType::None
| MessageType::Error
| MessageType::TimeOut => 0x0,
MessageType::Close => 0x8,
MessageType::Binary => 0x2,
MessageType::Ping => 0x9,
MessageType::Pong => 0xa,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum CloseCode {
#[default]
NormalClosure,
GoingAway,
ProtocolError,
UnsupportedData,
NoStatusReceived,
AbnormalClosure,
InvalidPayloadData,
PolicyViolation,
MessageTooBig,
MandatoryExtension,
InternalError,
Other(u16),
}
impl CloseCode {
#[must_use]
pub fn from_err(err: ErrorCode) -> CloseCode {
match err {
ErrorCode::SendingDataFailed => CloseCode::InternalError,
ErrorCode::ThreadException => CloseCode::InternalError,
ErrorCode::TimeOut => CloseCode::GoingAway,
ErrorCode::Unknown | ErrorCode::None => CloseCode::NormalClosure,
}
}
#[must_use]
pub fn str(&self) -> String {
match self {
CloseCode::NormalClosure => "正常关闭",
CloseCode::GoingAway => "端点离开",
CloseCode::ProtocolError => "协议错误",
CloseCode::UnsupportedData => "不支持的数据",
CloseCode::NoStatusReceived => "未收到状态码",
CloseCode::AbnormalClosure => "异常关闭",
CloseCode::InvalidPayloadData => "无效数据",
CloseCode::PolicyViolation => "策略违规",
CloseCode::MessageTooBig => "消息过大",
CloseCode::MandatoryExtension => "缺少扩展",
CloseCode::InternalError => "内部错误",
CloseCode::Other(code) => return format!("关闭码 {code}"),
}
.to_string()
}
#[must_use]
pub fn to_u16(&self) -> u16 {
match self {
CloseCode::NormalClosure => 1000,
CloseCode::GoingAway => 1001,
CloseCode::ProtocolError => 1002,
CloseCode::UnsupportedData => 1003,
CloseCode::NoStatusReceived => 1005,
CloseCode::AbnormalClosure => 1006,
CloseCode::InvalidPayloadData => 1007,
CloseCode::PolicyViolation => 1008,
CloseCode::MessageTooBig => 1009,
CloseCode::MandatoryExtension => 1010,
CloseCode::InternalError => 1011,
CloseCode::Other(code) => *code,
}
}
#[must_use]
pub fn is_valid_for_send(&self) -> bool {
!matches!(
self,
CloseCode::NoStatusReceived | CloseCode::AbnormalClosure
)
}
}
impl From<u16> for CloseCode {
fn from(code: u16) -> Self {
match code {
1000 => CloseCode::NormalClosure,
1001 => CloseCode::GoingAway,
1002 => CloseCode::ProtocolError,
1003 => CloseCode::UnsupportedData,
1005 => CloseCode::NoStatusReceived,
1006 => CloseCode::AbnormalClosure,
1007 => CloseCode::InvalidPayloadData,
1008 => CloseCode::PolicyViolation,
1009 => CloseCode::MessageTooBig,
1010 => CloseCode::MandatoryExtension,
1011 => CloseCode::InternalError,
_ => CloseCode::Other(code),
}
}
}
#[derive(Debug, Clone, Copy)]
pub enum ErrorCode {
SendingDataFailed,
Unknown,
ThreadException,
TimeOut,
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MessageMode {
Client,
Server,
}
pub struct NoticeMsg {
pub types: Types,
pub msg: JsonValue,
pub timestamp: i64,
pub channel: String,
pub user: String,
pub org: String,
pub admin: String,
}
impl NoticeMsg {
pub fn json(&mut self) -> JsonValue {
object! {
type:"notice",
channel: self.channel.clone(),
msg: self.msg.clone(),
timestamp: self.timestamp,
}
}
fn now() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64
}
fn push(notice: NoticeMsg) {
if let Ok(mut queue) = WS_NOTICE.lock() {
queue.push(notice);
}
}
pub fn to_all(channel: &str, msg: JsonValue) {
Self::push(NoticeMsg {
types: Types::All,
msg,
timestamp: Self::now(),
channel: channel.to_string(),
user: String::new(),
org: String::new(),
admin: String::new(),
});
}
pub fn to_org(channel: &str, msg: JsonValue, org: &str) {
Self::push(NoticeMsg {
types: Types::Org,
msg,
timestamp: Self::now(),
channel: channel.to_string(),
user: String::new(),
org: org.to_string(),
admin: String::new(),
});
}
pub fn to_user(channel: &str, msg: JsonValue, user: &str) {
Self::push(NoticeMsg {
types: Types::User,
msg,
timestamp: Self::now(),
channel: channel.to_string(),
user: user.to_string(),
org: String::new(),
admin: String::new(),
});
}
pub fn to_admin(channel: &str, msg: JsonValue, admin: &str) {
Self::push(NoticeMsg {
types: Types::Admin,
msg,
timestamp: Self::now(),
channel: channel.to_string(),
user: String::new(),
org: String::new(),
admin: admin.to_string(),
});
}
pub fn to_channel(channel: &str, msg: JsonValue) {
Self::push(NoticeMsg {
types: Types::Channel,
msg,
timestamp: Self::now(),
channel: channel.to_string(),
user: String::new(),
org: String::new(),
admin: String::new(),
});
}
}
pub enum Types {
All,
User,
Org,
Admin,
Channel,
}