Files
aistream-test/python/WebSocketFrameHeader.py
T
2025-12-01 21:00:54 +08:00

223 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from enum import IntEnum
import struct
import json
import gzip
from typing import Optional, Union, Dict, Any, List
# -------------------------- 协议常量定义 --------------------------
class ProtocolConst:
PROTOCOL_VERSION = 0b0001 # 协议版本v1
HEADER_SIZE = 8 # 头部固定字节数(字节1~字节8
MAX_BODY_SIZE = 1024 * 1024 * 10 # 最大包体大小(10MB,防止内存溢出)
STRING_ENCODING = "utf-8" # 字符串默认编码(统一编码格式,避免乱码)
# -------------------------- 枚举定义(业务约定) --------------------------
class MessageType(IntEnum):
"""消息类型(4位,字节1高4位)"""
PING = 0b0000 # 心跳
AUDIO_DATA = 0b0001 # 纯音频数据(pcm)
TEXT_MESSAGE = 0b0010 # 纯文本消息
AUDIO_TEXT_MIX = 0b0011 # 音频+文本混合数据
CONTROL_CMD = 0b0100 # 控制指令(暂停/继续等)
class SerializationType(IntEnum):
"""序列化方式(4位,字节2高4位)"""
RAW = 0b0000 # 原始二进制(音频等)
JSON = 0b0001 # JSON格式(文本/指令)
STRING = 0b0010 # 直接字符串(纯文本,UTF-8编码,无JSON包装)🔴新增类型
class CompressionType(IntEnum):
"""压缩方式(4位,字节2低4位)"""
NONE = 0b0000 # 无压缩
GZIP = 0b0001 # gzip压缩
class ControlCommand(IntEnum):
"""控制指令类型(配合MessageType.CONTROL_CMD使用)"""
HEARTBEAT = 0b0001 # 心跳响应
PAUSE = 0b0010 # 暂停
RESUME = 0b0011 # 继续
STOP = 0b0100 # 停止
# -------------------------- 协议工具类 --------------------------
class ProtocolCodec:
@staticmethod
def pack(
msg_type: MessageType,
body: Union[bytes, str, Dict[str, Any], List[Any]],
serialization: Optional[SerializationType] = None,
compression: CompressionType = CompressionType.NONE
) -> bytes:
"""
封装协议包(头部 + 包体)
:param msg_type: 消息类型
:param body: 包体数据(bytes/str/dict/list
:param serialization: 序列化方式(None时自动推导)
:param compression: 压缩方式
:return: 完整协议包(bytes
"""
# 1. 自动推导序列化方式(优化:纯文本消息自动用STRING,复杂结构用JSON)🔴修改推导逻辑
if serialization is None:
if msg_type == MessageType.AUDIO_DATA:
serialization = SerializationType.RAW
elif msg_type == MessageType.TEXT_MESSAGE: # 纯文本消息→STRING
serialization = SerializationType.STRING
elif msg_type in (MessageType.CONTROL_CMD, MessageType.AUDIO_TEXT_MIX): # 复杂结构→JSON
serialization = SerializationType.JSON
else:
raise ValueError(f"不支持的消息类型:{msg_type}")
# 2. 序列化包体(新增STRING类型处理)🔴新增分支
serialized_body: bytes
if serialization == SerializationType.RAW:
if not isinstance(body, bytes):
raise TypeError("RAW序列化要求body必须是bytes类型")
serialized_body = body
elif serialization == SerializationType.STRING:
if not isinstance(body, str):
raise TypeError("STRING序列化要求body必须是str类型")
serialized_body = body.encode(ProtocolConst.STRING_ENCODING) # 直接UTF-8编码,无JSON包装
elif serialization == SerializationType.JSON:
if isinstance(body, str):
serialized_body = body.encode(ProtocolConst.STRING_ENCODING)
elif isinstance(body, (dict, list)):
serialized_body = json.dumps(body, ensure_ascii=False, separators=(',', ':')).encode(ProtocolConst.STRING_ENCODING)
else:
raise TypeError("JSON序列化要求body必须是str/dict/list类型")
else:
raise ValueError(f"不支持的序列化方式:{serialization}")
# 3. 压缩包体(不变)
compressed_body: bytes
if compression == CompressionType.GZIP:
compressed_body = gzip.compress(serialized_body)
elif compression == CompressionType.NONE:
compressed_body = serialized_body
else:
raise ValueError(f"不支持的压缩方式:{compression}")
# 4. 校验包体大小(不变)
body_len = len(compressed_body)
if body_len > ProtocolConst.MAX_BODY_SIZE:
raise OverflowError(f"包体过大({body_len}字节),最大支持{ProtocolConst.MAX_BODY_SIZE}字节")
# 5. 构造头部(不变)
byte1 = (msg_type.value << 4) | 0x00 # 消息类型(4位) + 保留位1(4位)
byte2 = (serialization.value << 4) | (compression.value & 0x0F) # 序列化(4位) + 压缩(4位)
byte3_4 = struct.pack(">H", ProtocolConst.PROTOCOL_VERSION) # 协议版本(16位)
byte5_8 = struct.pack(">I", body_len) # 包体长度(32位)
header = bytes([byte1, byte2]) + byte3_4 + byte5_8
assert len(header) == ProtocolConst.HEADER_SIZE, f"头部长度错误:实际{len(header)}字节,预期{ProtocolConst.HEADER_SIZE}字节"
return header + compressed_body
@staticmethod
def unpack(packet: bytes) -> tuple[MessageType, SerializationType, CompressionType, Any]:
"""
解析协议包
:param packet: 完整协议包(头部 + 包体)
:return: (消息类型, 序列化方式, 压缩方式, 原始包体数据)
"""
# 1. 校验包长度(不变)
if len(packet) < ProtocolConst.HEADER_SIZE:
raise ValueError(f"包长度过短({len(packet)}字节),至少需要{ProtocolConst.HEADER_SIZE}字节头部")
# 2. 解析头部(不变)
header = packet[:ProtocolConst.HEADER_SIZE]
body = packet[ProtocolConst.HEADER_SIZE:]
byte1 = header[0]
byte2 = header[1]
msg_type = MessageType((byte1 >> 4) & 0x0F)
serialization = SerializationType((byte2 >> 4) & 0x0F)
compression = CompressionType(byte2 & 0x0F)
version = struct.unpack(">H", header[2:4])[0]
body_len = struct.unpack(">I", header[4:8])[0]
if version != ProtocolConst.PROTOCOL_VERSION:
raise ValueError(f"协议版本不匹配:收到v{version},支持v{ProtocolConst.PROTOCOL_VERSION}")
if len(body) != body_len:
raise ValueError(f"包体长度不匹配:头部声明{body_len}字节,实际{len(body)}字节")
# 3. 解压包体(不变)
decompressed_body: bytes
if compression == CompressionType.GZIP:
try:
decompressed_body = gzip.decompress(body)
except Exception as e:
raise ValueError(f"GZIP解压失败:{str(e)}")
elif compression == CompressionType.NONE:
decompressed_body = body
else:
raise ValueError(f"不支持的压缩方式:{compression}")
# 4. 反序列化包体(新增STRING类型处理)🔴新增分支
original_body: Any
if serialization == SerializationType.RAW:
original_body = decompressed_body
elif serialization == SerializationType.STRING:
try:
original_body = decompressed_body.decode(ProtocolConst.STRING_ENCODING) # 直接UTF-8解码,无JSON解析
except UnicodeDecodeError:
raise ValueError(f"STRING反序列化失败:{ProtocolConst.STRING_ENCODING}解码错误")
elif serialization == SerializationType.JSON:
try:
original_body = json.loads(decompressed_body.decode(ProtocolConst.STRING_ENCODING))
except UnicodeDecodeError:
raise ValueError(f"JSON反序列化失败:{ProtocolConst.STRING_ENCODING}解码错误")
except json.JSONDecodeError:
raise ValueError("JSON反序列化失败:格式错误")
else:
raise ValueError(f"不支持的序列化方式:{serialization}")
return msg_type, serialization, compression, original_body
# -------------------------- 使用示例(新增纯字符串消息测试) --------------------------
if __name__ == "__main__":
# 示例1:纯文本消息(自动用STRING序列化,无JSON包装)🔴测试新增类型
text_body = "你好,这是纯字符串消息(无JSON)"
text_packet = ProtocolCodec.pack(
msg_type=MessageType.TEXT_MESSAGE,
body=text_body,
compression=CompressionType.NONE
)
print(f"纯文本消息包长度:{len(text_packet)}字节")
msg_type1, ser1, comp1, body1 = ProtocolCodec.unpack(text_packet)
print(f"解析结果:类型={msg_type1.name},序列化={ser1.name},压缩={comp1.name},内容={body1}\n")
# 示例2:控制指令(JSON序列化,复杂结构)
control_body = {"cmd": ControlCommand.HEARTBEAT.value, "timestamp": 1699999999}
control_packet = ProtocolCodec.pack(
msg_type=MessageType.CONTROL_CMD,
body=control_body,
compression=CompressionType.NONE
)
print(f"控制指令包长度:{len(control_packet)}字节")
msg_type2, ser2, comp2, body2 = ProtocolCodec.unpack(control_packet)
print(f"解析结果:类型={msg_type2.name},序列化={ser2.name},压缩={comp2.name},内容={body2}\n")
# 示例3:音频数据(RAW序列化)
audio_body = b"\x00\x01\x02\x03\x04\x05" * 100 # 模拟PCM数据
audio_packet = ProtocolCodec.pack(
msg_type=MessageType.AUDIO_DATA,
body=audio_body,
compression=CompressionType.NONE
)
print(f"音频数据包长度:{len(audio_packet)}字节")
msg_type3, ser3, comp3, body3 = ProtocolCodec.unpack(audio_packet)
print(f"解析结果:类型={msg_type3.name},序列化={ser3.name},压缩={comp3.name},数据长度={len(body3)}字节\n")
# 示例4:混合数据(JSON序列化)
mix_body = {
"audio_data": list(audio_body[:10]),
"text_content": "混合数据仍用JSON"
}
mix_packet = ProtocolCodec.pack(
msg_type=MessageType.AUDIO_TEXT_MIX,
body=mix_body
)
print(f"混合数据包长度:{len(mix_packet)}字节")
msg_type4, ser4, comp4, body4 = ProtocolCodec.unpack(mix_packet)
print(f"解析结果:类型={msg_type4.name},序列化={ser4.name},压缩={comp4.name},内容={body4}")