Files
live-hub-py/core/huya/taf_protocol.py
T

478 lines
16 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.
"""
虎牙 TAF 协议 Python 实现
基于前端 TAF/WUP 实现逆向
类型标签:
0x00 INT8 0x01 INT16 0x02 INT32 0x03 INT64
0x04 FLOAT 0x05 DOUBLE 0x06 STRING1 0x07 STRING4
0x08 MAP 0x09 LIST 0x0a STRUCT_BEGIN 0x0b STRUCT_END
0x0c ZERO 0x0d SIMPLE_LIST
"""
import struct
import io
from typing import Any, Dict, List, Optional, Tuple
class TafType:
INT8 = 0x00
INT16 = 0x01
INT32 = 0x02
INT64 = 0x03
FLOAT = 0x04
DOUBLE = 0x05
STRING1 = 0x06
STRING4 = 0x07
MAP = 0x08
LIST = 0x09
STRUCT_BEGIN = 0x0a
STRUCT_END = 0x0b
ZERO = 0x0c
SIMPLE_LIST = 0x0d
# ============================================================
# 输出流(编码)
# ============================================================
class TafOutputStream:
"""TAF 编码输出流"""
def __init__(self):
self.buf = io.BytesIO()
def get_bytes(self) -> bytes:
return self.buf.getvalue()
# ---- head ----
def write_head(self, tag: int, data_type: int):
if tag < 15:
self.buf.write(struct.pack('B', (tag << 4) | data_type))
else:
self.buf.write(struct.pack('BB', 0xF0 | data_type, tag))
# ---- 整数(带自动优化) ----
def write_int8(self, tag: int, value: int):
if value == 0:
self.write_head(tag, TafType.ZERO)
else:
self.write_head(tag, TafType.INT8)
self.buf.write(struct.pack('b', value))
def write_int16(self, tag: int, value: int):
if -128 <= value <= 127:
self.write_int8(tag, value)
else:
self.write_head(tag, TafType.INT16)
self.buf.write(struct.pack('>h', value))
def write_int32(self, tag: int, value: int):
if -32768 <= value <= 32767:
self.write_int16(tag, value)
else:
self.write_head(tag, TafType.INT32)
self.buf.write(struct.pack('>i', value))
def write_int64(self, tag: int, value: int):
if -2147483648 <= value <= 2147483647:
self.write_int32(tag, value)
else:
self.write_head(tag, TafType.INT64)
self.buf.write(struct.pack('>q', value))
def write_uint64(self, tag: int, value: int):
"""uint64:超过 int32 范围用 INT64"""
if value <= 2147483647:
self.write_int32(tag, value)
else:
self.write_head(tag, TafType.INT64)
self.buf.write(struct.pack('>Q', value))
# ---- 浮点 ----
def write_float(self, tag: int, value: float):
self.write_head(tag, TafType.FLOAT)
self.buf.write(struct.pack('>f', value))
def write_double(self, tag: int, value: float):
self.write_head(tag, TafType.DOUBLE)
self.buf.write(struct.pack('>d', value))
# ---- 字符串 ----
def write_string(self, tag: int, value: str):
encoded = value.encode('utf-8')
length = len(encoded)
if length > 255:
self.write_head(tag, TafType.STRING4)
self.buf.write(struct.pack('>I', length))
else:
self.write_head(tag, TafType.STRING1)
self.buf.write(struct.pack('B', length))
self.buf.write(encoded)
# ---- 字节数组 ----
def write_bytes(self, tag: int, value: bytes):
self.write_head(tag, TafType.SIMPLE_LIST)
self.write_head(0, TafType.INT8) # 元素类型固定 INT8
self.write_int32(0, len(value))
self.buf.write(value)
# ---- 布尔 ----
def write_boolean(self, tag: int, value: bool):
self.write_int8(tag, 1 if value else 0)
# ---- 结构体 ----
def write_struct_begin(self, tag: int):
self.write_head(tag, TafType.STRUCT_BEGIN)
def write_struct_end(self):
self.write_head(0, TafType.STRUCT_END)
def write_struct(self, tag: int, struct_obj):
"""写入结构体对象(需实现 write_to"""
self.write_struct_begin(tag)
struct_obj.write_to(self)
self.write_struct_end()
# ---- Map ----
def write_map(self, tag: int, value: Dict[Any, Any],
key_writer=None, val_writer=None):
self.write_head(tag, TafType.MAP)
self.write_int32(0, len(value))
for k, v in value.items():
if key_writer:
key_writer(self, 0, k)
else:
self._write_any(0, k)
if val_writer:
val_writer(self, 1, v)
else:
self._write_any(1, v)
# ---- List ----
def write_list(self, tag: int, value: List[Any], item_writer=None):
self.write_head(tag, TafType.LIST)
self.write_int32(0, len(value))
for item in value:
if item_writer:
item_writer(self, 0, item)
else:
self._write_any(0, item)
def _write_any(self, tag: int, value: Any):
if isinstance(value, bool):
self.write_boolean(tag, value)
elif isinstance(value, int):
self.write_int64(tag, value)
elif isinstance(value, float):
self.write_double(tag, value)
elif isinstance(value, str):
self.write_string(tag, value)
elif isinstance(value, bytes):
self.write_bytes(tag, value)
elif isinstance(value, dict):
self.write_map(tag, value)
elif isinstance(value, (list, tuple)):
self.write_list(tag, list(value))
elif hasattr(value, 'write_to'):
self.write_struct(tag, value)
else:
raise TypeError(f"不支持的类型: {type(value)}")
# ============================================================
# 输入流(解码)—— 完整实现,支持所有类型
# ============================================================
class TafInputStream:
"""TAF 解码输入流"""
def __init__(self, data: bytes):
self.buf = io.BytesIO(data)
def peek_head(self) -> Tuple[int, int]:
"""读取 head 但不消费(用于探测)"""
pos = self.buf.tell()
try:
return self.read_head()
finally:
self.buf.seek(pos)
def read_head(self) -> Tuple[int, int]:
"""返回 (tag, type)"""
data = self.buf.read(1)
if not data:
raise EOFError("读取到文件末尾")
b = struct.unpack('B', data)[0]
tag = (b >> 4) & 0x0F
data_type = b & 0x0F
if tag == 15:
data = self.buf.read(1)
if not data:
raise EOFError("读取 tag 扩展字节失败")
tag = struct.unpack('B', data)[0]
return tag, data_type
# ---- 跳过 ----
def skip_field(self, data_type: int):
if data_type == TafType.ZERO:
return
if data_type == TafType.INT8:
self.buf.read(1)
elif data_type == TafType.INT16:
self.buf.read(2)
elif data_type == TafType.INT32:
self.buf.read(4)
elif data_type == TafType.INT64:
self.buf.read(8)
elif data_type == TafType.FLOAT:
self.buf.read(4)
elif data_type == TafType.DOUBLE:
self.buf.read(8)
elif data_type == TafType.STRING1:
length = struct.unpack('B', self.buf.read(1))[0]
self.buf.read(length)
elif data_type == TafType.STRING4:
length = struct.unpack('>I', self.buf.read(4))[0]
self.buf.read(length)
elif data_type == TafType.MAP:
self._skip_map()
elif data_type == TafType.LIST:
self._skip_list()
elif data_type == TafType.SIMPLE_LIST:
# [head(0,INT8)] [int32 length] [bytes]
self.read_head() # 元素类型
length = self._read_int_len()
self.buf.read(length)
elif data_type == TafType.STRUCT_BEGIN:
self._skip_struct()
elif data_type == TafType.STRUCT_END:
pass
else:
raise ValueError(f"未知类型 0x{data_type:02x}")
def _read_int_len(self) -> int:
"""读 map/list 长度(int32 带优化)"""
tag, dtype = self.read_head()
return self._read_int_value(dtype)
def _read_int_value(self, dtype: int) -> int:
if dtype == TafType.ZERO:
return 0
if dtype == TafType.INT8:
return struct.unpack('b', self.buf.read(1))[0]
if dtype == TafType.INT16:
return struct.unpack('>h', self.buf.read(2))[0]
if dtype == TafType.INT32:
return struct.unpack('>i', self.buf.read(4))[0]
if dtype == TafType.INT64:
return struct.unpack('>q', self.buf.read(8))[0]
raise ValueError(f"期望整数, 实际 0x{dtype:02x}")
def _skip_struct(self):
while True:
tag, dtype = self.read_head()
if dtype == TafType.STRUCT_END:
break
self.skip_field(dtype)
def _skip_map(self):
count = self._read_int_len()
for _ in range(count):
_, kt = self.read_head()
self.skip_field(kt)
_, vt = self.read_head()
self.skip_field(vt)
def _skip_list(self):
count = self._read_int_len()
for _ in range(count):
_, it = self.read_head()
self.skip_field(it)
# ---- 带跳过策略的字段读取:找到 tag,否则返回默认 ----
def _find_tag(self, target_tag: int, required: bool) -> Optional[Tuple[int, int]]:
"""逐个读 headtag 相等则返回,tag 超过则回退并返回 None"""
while True:
pos = self.buf.tell()
tag, dtype = self.read_head()
if tag == target_tag:
return (tag, dtype)
if dtype == TafType.STRUCT_END:
# 回退,让上层处理 STRUCT_END
self.buf.seek(pos)
if required:
raise ValueError(f"未找到 tag={target_tag} (遇到 STRUCT_END)")
return None
if tag > target_tag:
# 超过目标 tag,回退(字段不存在)
self.buf.seek(pos)
if required:
raise ValueError(f"未找到 tag={target_tag} (遇到 tag={tag})")
return None
# tag < target_tag,跳过此字段
self.skip_field(dtype)
# ---- 基本类型读取 ----
def read_int8(self, tag: int, required: bool = False, default: int = 0) -> int:
found = self._find_tag(tag, required)
if not found:
return default
return self._read_int_value(found[1])
def read_int16(self, tag: int, required: bool = False, default: int = 0) -> int:
return self.read_int8(tag, required, default)
def read_int32(self, tag: int, required: bool = False, default: int = 0) -> int:
return self.read_int8(tag, required, default)
def read_int64(self, tag: int, required: bool = False, default: int = 0) -> int:
return self.read_int8(tag, required, default)
def read_uint64(self, tag: int, required: bool = False, default: int = 0) -> int:
found = self._find_tag(tag, required)
if not found:
return default
dtype = found[1]
if dtype == TafType.ZERO:
return 0
if dtype == TafType.INT8:
return struct.unpack('B', self.buf.read(1))[0]
if dtype == TafType.INT16:
return struct.unpack('>H', self.buf.read(2))[0]
if dtype == TafType.INT32:
return struct.unpack('>I', self.buf.read(4))[0]
if dtype == TafType.INT64:
return struct.unpack('>Q', self.buf.read(8))[0]
raise ValueError(f"期望 uint, 实际 0x{dtype:02x}")
def read_boolean(self, tag: int, required: bool = False, default: bool = False) -> bool:
return bool(self.read_int8(tag, required, 1 if default else 0))
def read_float(self, tag: int, required: bool = False, default: float = 0.0) -> float:
found = self._find_tag(tag, required)
if not found:
return default
dtype = found[1]
if dtype == TafType.ZERO:
return 0.0
if dtype == TafType.FLOAT:
return struct.unpack('>f', self.buf.read(4))[0]
if dtype == TafType.DOUBLE:
return struct.unpack('>d', self.buf.read(8))[0]
return float(self._read_int_value(dtype))
def read_double(self, tag: int, required: bool = False, default: float = 0.0) -> float:
return self.read_float(tag, required, default)
def read_string(self, tag: int, required: bool = False, default: str = "") -> str:
found = self._find_tag(tag, required)
if not found:
return default
dtype = found[1]
if dtype == TafType.STRING1:
length = struct.unpack('B', self.buf.read(1))[0]
elif dtype == TafType.STRING4:
length = struct.unpack('>I', self.buf.read(4))[0]
else:
raise ValueError(f"期望 string, 实际 0x{dtype:02x}")
return self.buf.read(length).decode('utf-8', errors='replace')
def read_bytes(self, tag: int, required: bool = False, default: bytes = b'') -> bytes:
found = self._find_tag(tag, required)
if not found:
return default
dtype = found[1]
if dtype != TafType.SIMPLE_LIST:
raise ValueError(f"期望 bytes, 实际 0x{dtype:02x}")
self.read_head() # 元素类型 INT8
length = self._read_int_len()
return self.buf.read(length)
# ---- 复合类型 ----
def read_map(self, tag: int, required: bool = False,
key_reader=None, val_reader=None) -> Dict:
found = self._find_tag(tag, required)
if not found:
return {}
if found[1] != TafType.MAP:
raise ValueError(f"期望 map, 实际 0x{found[1]:02x}")
count = self._read_int_len()
result = {}
for _ in range(count):
_, kt = self.read_head()
k = self._read_value(kt, key_reader)
_, vt = self.read_head()
v = self._read_value(vt, val_reader)
result[k] = v
return result
def read_list(self, tag: int, required: bool = False,
item_reader=None) -> List:
found = self._find_tag(tag, required)
if not found:
return []
if found[1] != TafType.LIST:
raise ValueError(f"期望 list, 实际 0x{found[1]:02x}")
count = self._read_int_len()
result = []
for _ in range(count):
_, it = self.read_head()
result.append(self._read_value(it, item_reader))
return result
def read_struct(self, tag: int, struct_class, required: bool = False):
"""读取结构体,struct_class 需有无参构造 + read_from"""
found = self._find_tag(tag, required)
if not found:
return None
if found[1] != TafType.STRUCT_BEGIN:
raise ValueError(f"期望 struct, 实际 0x{found[1]:02x}")
obj = struct_class()
obj.read_from(self)
# 消费 STRUCT_END
t, dt = self.read_head()
if dt != TafType.STRUCT_END:
raise ValueError(f"期望 STRUCT_END, 实际 0x{dt:02x}")
return obj
def _read_value(self, dtype: int, reader=None):
if reader:
return reader(self, 0)
# 自动推断
if dtype == TafType.STRING1:
length = struct.unpack('B', self.buf.read(1))[0]
return self.buf.read(length).decode('utf-8', errors='replace')
if dtype == TafType.STRING4:
length = struct.unpack('>I', self.buf.read(4))[0]
return self.buf.read(length).decode('utf-8', errors='replace')
if dtype in (TafType.ZERO, TafType.INT8, TafType.INT16,
TafType.INT32, TafType.INT64):
return self._read_int_value(dtype)
if dtype == TafType.STRUCT_BEGIN:
# 未知 struct,跳过
self._skip_struct()
return None
self.skip_field(dtype)
return None
# ============================================================
# 结构体基类
# ============================================================
class TafStruct:
"""TAF 结构体基类:子类实现 write_to / read_from"""
def write_to(self, os: TafOutputStream):
raise NotImplementedError
def read_from(self, ins: TafInputStream):
raise NotImplementedError
def to_dict(self) -> Dict[str, Any]:
"""调试用:转字典"""
return {k: v for k, v in self.__dict__.items()
if not k.startswith('_')}
def __repr__(self):
return f"{self.__class__.__name__}({self.to_dict()})"