x
This commit is contained in:
@@ -0,0 +1,223 @@
|
||||
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}")
|
||||
Reference in New Issue
Block a user