294 lines
13 KiB
Python
294 lines
13 KiB
Python
from enum import IntEnum
|
||
import struct
|
||
import json
|
||
import gzip
|
||
from typing import Optional, Union, Dict, Any, List, Tuple
|
||
|
||
|
||
# -------------------------- 协议常量定义(与JS完全一致) --------------------------
|
||
class ProtocolConst:
|
||
PROTOCOL_VERSION = 0b0001 # 协议版本(4位,字节0低4位,0~15)
|
||
HEADER_SIZE = 8 # 头部固定字节数(字节0~7)
|
||
MAX_BODY_SIZE = 0xFFFFFF # 最大包体大小(24位,≈16MB,字节4~6存储)
|
||
STRING_ENCODING = "utf-8" # 字符串默认编码
|
||
|
||
|
||
# -------------------------- 枚举定义(与JS码值完全对齐) --------------------------
|
||
class MessageType(IntEnum):
|
||
"""消息类型(4位,字节0高4位,0~15)"""
|
||
PING = 0b0001 # 心跳(支持空包体)
|
||
AUDIO_DATA = 0b0010 # 纯音频数据(pcm)
|
||
TEXT_MESSAGE = 0b0011 # 纯文本消息
|
||
CONTROL_CMD = 0b0100 # 控制指令(暂停/继续等)
|
||
# 预留12种类型(0b0100 ~ 0b1111)
|
||
|
||
|
||
class SerializationType(IntEnum):
|
||
"""序列化方式(3位,字节1高3位,1~8)"""
|
||
RAW = 0b001 # 原始二进制(1)
|
||
JSON = 0b010 # JSON格式(2)
|
||
STRING = 0b011 # 直接字符串(3)
|
||
# 预留5种方式(0b100 ~ 0b111)
|
||
|
||
|
||
class CompressionType(IntEnum):
|
||
"""压缩方式(3位,字节1中3位,1~8)"""
|
||
NONE = 0b001 # 无压缩(默认值1)
|
||
GZIP = 0b010 # gzip压缩(2)
|
||
# 预留6种方式(0b011 ~ 0b111)
|
||
|
||
|
||
class ControlCommand(IntEnum):
|
||
"""控制指令类型(配合MessageType.CONTROL_CMD使用)"""
|
||
HEARTBEAT = 0b0001 # 心跳响应
|
||
PAUSE = 0b0010 # 暂停
|
||
RESUME = 0b0011 # 继续
|
||
STOP = 0b0100 # 停止
|
||
|
||
|
||
# -------------------------- 协议工具类(与JS协议结构一致) --------------------------
|
||
class ProtocolCodec:
|
||
@staticmethod
|
||
def pack(
|
||
msg_type: MessageType,
|
||
body: Union[bytes, str, Dict[str, Any], List[Any], None] = None,
|
||
sequence: int = 0,
|
||
serialization: Optional[SerializationType] = None,
|
||
compression: CompressionType = CompressionType.NONE
|
||
) -> bytes:
|
||
"""
|
||
封装协议包(与JS pack 方法完全兼容)
|
||
:param msg_type: 消息类型
|
||
:param body: 包体数据(PING消息可传None/空,其他类型必填)
|
||
:param sequence: 消息顺序号(0~65535,默认0)
|
||
:param serialization: 序列化方式(None时自动推导)
|
||
:param compression: 压缩方式(默认无压缩)
|
||
:return: 完整协议包(bytes)
|
||
"""
|
||
# 校验顺序号范围(0~65535)
|
||
if not isinstance(sequence, int) or not (0 <= sequence <= 0xFFFF):
|
||
raise ValueError(f"消息顺序号必须是0~65535的整数,当前传入:{sequence}")
|
||
|
||
# 特殊处理:PING消息支持空包体
|
||
if msg_type == MessageType.PING:
|
||
# PING消息强制RAW序列化
|
||
serialization = SerializationType.RAW
|
||
# 空包体转为空bytes
|
||
if body is None:
|
||
body = b""
|
||
if not isinstance(body, bytes):
|
||
raise TypeError(f"PING消息仅支持空包体或bytes类型,当前传入:{type(body)}")
|
||
else:
|
||
# 非PING消息包体不能为空
|
||
if body is None:
|
||
raise ValueError(f"非PING消息({msg_type.name})包体不能为空")
|
||
|
||
# 1. 自动推导序列化方式(非PING消息)
|
||
if serialization is None and msg_type != MessageType.PING:
|
||
if msg_type == MessageType.AUDIO_DATA:
|
||
serialization = SerializationType.RAW
|
||
elif msg_type == MessageType.TEXT_MESSAGE:
|
||
serialization = SerializationType.STRING
|
||
elif msg_type == MessageType.CONTROL_CMD:
|
||
serialization = SerializationType.JSON
|
||
else:
|
||
raise ValueError(f"不支持的消息类型:{msg_type}")
|
||
|
||
# 2. 校验枚举值范围
|
||
if not (0b001 <= serialization.value <= 0b111):
|
||
raise ValueError(f"序列化方式必须在1~8(0b001~0b111)范围内,当前传入:{serialization.value}")
|
||
if not (0b001 <= compression.value <= 0b111):
|
||
raise ValueError(f"压缩方式必须在1~8(0b001~0b111)范围内,当前传入:{compression.value}")
|
||
|
||
# 3. 序列化包体
|
||
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)
|
||
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}")
|
||
|
||
# 4. 压缩包体
|
||
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}")
|
||
|
||
# 5. 校验包体大小(24位最大支持0xFFFFFF字节)
|
||
body_len = len(compressed_body)
|
||
if body_len > ProtocolConst.MAX_BODY_SIZE:
|
||
raise OverflowError(
|
||
f"包体过大({body_len}字节),最大支持{ProtocolConst.MAX_BODY_SIZE}字节(≈16MB)"
|
||
)
|
||
|
||
# 6. 构造头部(与JS头部结构完全一致)
|
||
# 字节0:消息类型(高4位) + 协议版本(低4位)
|
||
byte0 = ((msg_type.value & 0x0F) << 4) | (ProtocolConst.PROTOCOL_VERSION & 0x0F)
|
||
# 字节1:序列化方式(高3位) + 压缩方式(中3位) + 保留位(低2位)
|
||
byte1 = ((serialization.value & 0x07) << 5) | ((compression.value & 0x07) << 2) | 0x00
|
||
# 字节2~3:消息顺序号(16位大端序)
|
||
byte2_3 = struct.pack(">H", sequence)
|
||
# 字节4~6:包体长度(24位大端序)
|
||
byte4 = (body_len >> 16) & 0xFF
|
||
byte5 = (body_len >> 8) & 0xFF
|
||
byte6 = body_len & 0xFF
|
||
# 字节7:保留位(固定0x00)
|
||
byte7 = 0x00
|
||
|
||
# 拼接头部
|
||
header = (
|
||
bytes([byte0, byte1]) +
|
||
byte2_3 +
|
||
bytes([byte4, byte5, byte6, byte7])
|
||
)
|
||
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, int, Any]:
|
||
"""
|
||
解析协议包(与JS unpack 方法完全兼容)
|
||
: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:]
|
||
|
||
# 字节0:消息类型(高4位) + 协议版本(低4位)
|
||
byte0 = header[0]
|
||
msg_type = MessageType((byte0 >> 4) & 0x0F)
|
||
version = byte0 & 0x0F
|
||
|
||
# 字节1:序列化方式(高3位) + 压缩方式(中3位) + 保留位(低2位)
|
||
byte1 = header[1]
|
||
serialization = SerializationType((byte1 >> 5) & 0x07)
|
||
compression = CompressionType((byte1 >> 2) & 0x07)
|
||
|
||
# 字节2~3:消息顺序号(16位大端序)
|
||
sequence = struct.unpack(">H", header[2:4])[0]
|
||
|
||
# 字节4~6:包体长度(24位大端序),字节7:保留位(忽略)
|
||
body_len = (header[4] << 16) | (header[5] << 8) | header[6]
|
||
|
||
# 校验版本和包体长度
|
||
if version != ProtocolConst.PROTOCOL_VERSION:
|
||
raise ValueError(
|
||
f"协议版本不匹配:收到v{version}(0b{version:04b}),支持v{ProtocolConst.PROTOCOL_VERSION}(0b{ProtocolConst.PROTOCOL_VERSION:04b})"
|
||
)
|
||
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. 反序列化包体(空包体返回None)
|
||
original_body: Any
|
||
if len(decompressed_body) == 0:
|
||
original_body = None
|
||
elif serialization == SerializationType.RAW:
|
||
original_body = decompressed_body
|
||
elif serialization == SerializationType.STRING:
|
||
try:
|
||
original_body = decompressed_body.decode(ProtocolConst.STRING_ENCODING)
|
||
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 as e:
|
||
raise ValueError(f"JSON反序列化失败:格式错误({str(e)})")
|
||
else:
|
||
raise ValueError(f"不支持的序列化方式:{serialization}")
|
||
|
||
return msg_type, serialization, compression, sequence, original_body
|
||
|
||
|
||
# -------------------------- 使用示例(验证与JS兼容性) --------------------------
|
||
if __name__ == "__main__":
|
||
# 示例1:PING消息(空包体,默认顺序号0)
|
||
ping_packet = ProtocolCodec.pack(MessageType.PING)
|
||
print(f"PING消息包长度:{len(ping_packet)}字节(仅头部)")
|
||
msg_type1, ser1, comp1, seq1, body1 = ProtocolCodec.unpack(ping_packet)
|
||
print(f"PING解析结果:类型={msg_type1.name},序列化={ser1.name},压缩={comp1.name},顺序号={seq1},包体={body1}\n")
|
||
|
||
# 示例2:纯文本消息(STRING序列化,指定顺序号)
|
||
text_body = "Python与JS协议兼容测试(纯字符串)"
|
||
text_packet = ProtocolCodec.pack(
|
||
msg_type=MessageType.TEXT_MESSAGE,
|
||
body=text_body,
|
||
sequence=1001,
|
||
compression=CompressionType.NONE
|
||
)
|
||
print(f"文本消息包长度:{len(text_packet)}字节")
|
||
msg_type2, ser2, comp2, seq2, body2 = ProtocolCodec.unpack(text_packet)
|
||
print(f"文本解析结果:类型={msg_type2.name},序列化={ser2.name},顺序号={seq2},内容={body2}\n")
|
||
|
||
# 示例3:控制指令(JSON序列化)
|
||
control_body = {"cmd": ControlCommand.PAUSE.value, "reason": "用户主动暂停"}
|
||
control_packet = ProtocolCodec.pack(
|
||
msg_type=MessageType.CONTROL_CMD,
|
||
body=control_body,
|
||
sequence=1002
|
||
)
|
||
print(f"控制指令包长度:{len(control_packet)}字节")
|
||
msg_type3, ser3, comp3, seq3, body3 = ProtocolCodec.unpack(control_packet)
|
||
print(f"控制指令解析结果:类型={msg_type3.name},序列化={ser3.name},顺序号={seq3},内容={body3}\n")
|
||
|
||
# 示例4:音频数据(RAW序列化)
|
||
audio_body = b"\x00\x01\x02\x03\x04\x05" * 100 # 模拟PCM数据
|
||
audio_packet = ProtocolCodec.pack(
|
||
msg_type=MessageType.AUDIO_DATA,
|
||
body=audio_body,
|
||
sequence=1003
|
||
)
|
||
print(f"音频数据包长度:{len(audio_packet)}字节")
|
||
msg_type4, ser4, comp4, seq4, body4 = ProtocolCodec.unpack(audio_packet)
|
||
print(f"音频解析结果:类型={msg_type4.name},序列化={ser4.name},顺序号={seq4},数据长度={len(body4)}字节\n")
|
||
|
||
# 示例5:GZIP压缩测试(需JS端启用GZIP解压)
|
||
long_text_body = "这是一段很长的文本,用于测试GZIP压缩效果" * 100
|
||
gzip_packet = ProtocolCodec.pack(
|
||
msg_type=MessageType.TEXT_MESSAGE,
|
||
body=long_text_body,
|
||
sequence=1004,
|
||
compression=CompressionType.GZIP
|
||
)
|
||
print(f"GZIP压缩后包长度:{len(gzip_packet)}字节(原始文本长度:{len(long_text_body.encode())}字节)")
|
||
msg_type5, ser5, comp5, seq5, body5 = ProtocolCodec.unpack(gzip_packet)
|
||
print(f"GZIP解析结果:类型={msg_type5.name},压缩={comp5.name},内容前50字:{body5[:50]}...") |