Files
aistream-test/python/WebSocketFrameHeader.py
T
2025-12-01 23:32:34 +08:00

294 lines
13 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, 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~80b001~0b111)范围内,当前传入:{serialization.value}")
if not (0b001 <= compression.value <= 0b111):
raise ValueError(f"压缩方式必须在1~80b001~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]}...")