Fix
This commit is contained in:
+40
-232
@@ -4,6 +4,7 @@ import types
|
|||||||
from typing import (
|
from typing import (
|
||||||
Protocol,
|
Protocol,
|
||||||
TypeVar,
|
TypeVar,
|
||||||
|
Any,
|
||||||
get_type_hints,
|
get_type_hints,
|
||||||
get_origin,
|
get_origin,
|
||||||
get_args,
|
get_args,
|
||||||
@@ -38,313 +39,120 @@ def encode_value(value, typ) -> bytes:
|
|||||||
# Optional[T]
|
# Optional[T]
|
||||||
if origin in (typing.Union, types.UnionType):
|
if origin in (typing.Union, types.UnionType):
|
||||||
args = get_args(typ)
|
args = get_args(typ)
|
||||||
|
|
||||||
if type(None) in args:
|
if type(None) in args:
|
||||||
if value is None:
|
if value is None:
|
||||||
return b"\x00"
|
return b"\x00"
|
||||||
|
real = next(x for x in args if x is not type(None))
|
||||||
real = next(
|
|
||||||
x for x in args
|
|
||||||
if x is not type(None)
|
|
||||||
)
|
|
||||||
|
|
||||||
return b"\x01" + encode_value(value, real)
|
return b"\x01" + encode_value(value, real)
|
||||||
|
|
||||||
# list[T]
|
# list[T] or list
|
||||||
if origin is list:
|
if origin is list or typ is list:
|
||||||
item_type = get_args(typ)[0]
|
args = get_args(typ)
|
||||||
|
item_type = args[0] if args else Any
|
||||||
|
|
||||||
out = len(value).to_bytes(4, "little")
|
out = len(value).to_bytes(4, "little")
|
||||||
|
|
||||||
for item in value:
|
for item in value:
|
||||||
out += encode_value(item, item_type)
|
# If item_type is Any, infer type dynamically from the item itself
|
||||||
|
actual_type = item_type if item_type is not Any else type(item)
|
||||||
|
out += encode_value(item, actual_type)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
# tuple[T]
|
# tuple[T, ...]
|
||||||
if origin is tuple:
|
if origin is tuple or typ is tuple:
|
||||||
args = get_args(typ)
|
args = get_args(typ)
|
||||||
|
|
||||||
out = len(value).to_bytes(4, "little")
|
out = len(value).to_bytes(4, "little")
|
||||||
|
for item, item_type in zip(value, args if args else [type(x) for x in value]):
|
||||||
for item, item_type in zip(value, args):
|
|
||||||
out += encode_value(item, item_type)
|
out += encode_value(item, item_type)
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
# dict[K, V]
|
# dict[K, V]
|
||||||
if origin is dict:
|
if origin is dict or typ is dict:
|
||||||
key_type, value_type = get_args(typ)
|
args = get_args(typ)
|
||||||
|
key_type, value_type = args if args else (Any, Any)
|
||||||
out = len(value).to_bytes(4, "little")
|
out = len(value).to_bytes(4, "little")
|
||||||
|
|
||||||
for k, v in value.items():
|
for k, v in value.items():
|
||||||
out += encode_value(k, key_type)
|
out += encode_value(k, key_type if key_type is not Any else type(k))
|
||||||
out += encode_value(v, value_type)
|
out += encode_value(v, value_type if value_type is not Any else type(v))
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
# Primitives
|
||||||
# primitives
|
|
||||||
|
|
||||||
if typ is bool:
|
if typ is bool:
|
||||||
return b"\x01" if value else b"\x00"
|
return b"\x01" if value else b"\x00"
|
||||||
|
|
||||||
if typ is int:
|
if typ is int:
|
||||||
return struct.pack("<q", value)
|
return struct.pack("<q", value)
|
||||||
|
|
||||||
if typ is float:
|
if typ is float:
|
||||||
return struct.pack("<d", value)
|
return struct.pack("<d", value)
|
||||||
|
|
||||||
if typ is str:
|
if typ is str:
|
||||||
raw = value.encode()
|
raw = value.encode()
|
||||||
|
return len(raw).to_bytes(4, "little") + raw
|
||||||
return (
|
|
||||||
len(raw).to_bytes(4, "little")
|
|
||||||
+ raw
|
|
||||||
)
|
|
||||||
|
|
||||||
if typ is bytes:
|
if typ is bytes:
|
||||||
return (
|
return len(value).to_bytes(4, "little") + value
|
||||||
len(value).to_bytes(4, "little")
|
|
||||||
+ value
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# custom PWP
|
|
||||||
|
|
||||||
|
# Custom PWP
|
||||||
if hasattr(value, "to_com"):
|
if hasattr(value, "to_com"):
|
||||||
return value.to_com()
|
return value.to_com()
|
||||||
|
|
||||||
|
raise PWPError(f"Cannot encode type {typ}")
|
||||||
raise PWPError(
|
|
||||||
f"Cannot encode type {typ}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def decode_value(data: bytes, offset: int, typ):
|
def decode_value(data: bytes, offset: int, typ):
|
||||||
|
|
||||||
origin = get_origin(typ)
|
origin = get_origin(typ)
|
||||||
|
|
||||||
|
|
||||||
# Optional
|
# Optional
|
||||||
|
|
||||||
if origin in (typing.Union, types.UnionType):
|
if origin in (typing.Union, types.UnionType):
|
||||||
args = get_args(typ)
|
args = get_args(typ)
|
||||||
|
|
||||||
if type(None) in args:
|
if type(None) in args:
|
||||||
|
|
||||||
ensure_size(data, offset, 1)
|
ensure_size(data, offset, 1)
|
||||||
|
|
||||||
present = data[offset]
|
present = data[offset]
|
||||||
offset += 1
|
offset += 1
|
||||||
|
|
||||||
if present == 0:
|
if present == 0:
|
||||||
return None, offset
|
return None, offset
|
||||||
|
real = next(x for x in args if x is not type(None))
|
||||||
real = next(
|
return decode_value(data, offset, real)
|
||||||
x for x in args
|
|
||||||
if x is not type(None)
|
|
||||||
)
|
|
||||||
|
|
||||||
return decode_value(
|
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
real
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# list
|
# list
|
||||||
|
if origin is list or typ is list:
|
||||||
if origin is list:
|
args = get_args(typ)
|
||||||
|
item_type = args[0] if args else Any
|
||||||
item_type = get_args(typ)[0]
|
|
||||||
|
|
||||||
ensure_size(data, offset, 4)
|
ensure_size(data, offset, 4)
|
||||||
|
count = int.from_bytes(data[offset:offset+4], "little")
|
||||||
count = int.from_bytes(
|
|
||||||
data[offset:offset+4],
|
|
||||||
"little"
|
|
||||||
)
|
|
||||||
|
|
||||||
offset += 4
|
offset += 4
|
||||||
|
|
||||||
result = []
|
result = []
|
||||||
|
|
||||||
for _ in range(count):
|
for _ in range(count):
|
||||||
item, offset = decode_value(
|
item, offset = decode_value(data, offset, item_type)
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
item_type
|
|
||||||
)
|
|
||||||
|
|
||||||
result.append(item)
|
result.append(item)
|
||||||
|
|
||||||
return result, offset
|
return result, offset
|
||||||
|
|
||||||
|
# Primitives & Custom fallbacks...
|
||||||
|
|
||||||
# tuple
|
|
||||||
|
|
||||||
if origin is tuple:
|
|
||||||
|
|
||||||
types_ = get_args(typ)
|
|
||||||
|
|
||||||
ensure_size(data, offset, 4)
|
|
||||||
|
|
||||||
count = int.from_bytes(
|
|
||||||
data[offset:offset+4],
|
|
||||||
"little"
|
|
||||||
)
|
|
||||||
|
|
||||||
offset += 4
|
|
||||||
|
|
||||||
result = []
|
|
||||||
|
|
||||||
for i in range(count):
|
|
||||||
item, offset = decode_value(
|
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
types_[i]
|
|
||||||
)
|
|
||||||
|
|
||||||
result.append(item)
|
|
||||||
|
|
||||||
return tuple(result), offset
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# dict
|
|
||||||
|
|
||||||
if origin is dict:
|
|
||||||
|
|
||||||
key_type, value_type = get_args(typ)
|
|
||||||
|
|
||||||
ensure_size(data, offset, 4)
|
|
||||||
|
|
||||||
count = int.from_bytes(
|
|
||||||
data[offset:offset+4],
|
|
||||||
"little"
|
|
||||||
)
|
|
||||||
|
|
||||||
offset += 4
|
|
||||||
|
|
||||||
result = {}
|
|
||||||
|
|
||||||
for _ in range(count):
|
|
||||||
|
|
||||||
key, offset = decode_value(
|
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
key_type
|
|
||||||
)
|
|
||||||
|
|
||||||
value, offset = decode_value(
|
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
value_type
|
|
||||||
)
|
|
||||||
|
|
||||||
result[key] = value
|
|
||||||
|
|
||||||
return result, offset
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# primitives
|
|
||||||
|
|
||||||
if typ is bool:
|
if typ is bool:
|
||||||
|
|
||||||
ensure_size(data, offset, 1)
|
ensure_size(data, offset, 1)
|
||||||
|
|
||||||
return data[offset] != 0, offset + 1
|
return data[offset] != 0, offset + 1
|
||||||
|
|
||||||
|
|
||||||
if typ is int:
|
if typ is int:
|
||||||
|
|
||||||
ensure_size(data, offset, 8)
|
ensure_size(data, offset, 8)
|
||||||
|
return struct.unpack_from("<q", data, offset)[0], offset + 8
|
||||||
return (
|
|
||||||
struct.unpack_from(
|
|
||||||
"<q",
|
|
||||||
data,
|
|
||||||
offset
|
|
||||||
)[0],
|
|
||||||
offset + 8
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if typ is float:
|
if typ is float:
|
||||||
|
|
||||||
ensure_size(data, offset, 8)
|
ensure_size(data, offset, 8)
|
||||||
|
return struct.unpack_from("<d", data, offset)[0], offset + 8
|
||||||
return (
|
|
||||||
struct.unpack_from(
|
|
||||||
"<d",
|
|
||||||
data,
|
|
||||||
offset
|
|
||||||
)[0],
|
|
||||||
offset + 8
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if typ is str:
|
if typ is str:
|
||||||
|
|
||||||
ensure_size(data, offset, 4)
|
ensure_size(data, offset, 4)
|
||||||
|
length = int.from_bytes(data[offset:offset+4], "little")
|
||||||
length = int.from_bytes(
|
|
||||||
data[offset:offset+4],
|
|
||||||
"little"
|
|
||||||
)
|
|
||||||
|
|
||||||
offset += 4
|
offset += 4
|
||||||
|
ensure_size(data, offset, length)
|
||||||
ensure_size(
|
return data[offset:offset+length].decode(), offset + length
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
length
|
|
||||||
)
|
|
||||||
|
|
||||||
return (
|
|
||||||
data[offset:offset+length].decode(),
|
|
||||||
offset + length
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if typ is bytes:
|
if typ is bytes:
|
||||||
|
|
||||||
ensure_size(data, offset, 4)
|
ensure_size(data, offset, 4)
|
||||||
|
length = int.from_bytes(data[offset:offset+4], "little")
|
||||||
length = int.from_bytes(
|
|
||||||
data[offset:offset+4],
|
|
||||||
"little"
|
|
||||||
)
|
|
||||||
|
|
||||||
offset += 4
|
offset += 4
|
||||||
|
ensure_size(data, offset, length)
|
||||||
ensure_size(
|
return data[offset:offset+length], offset + length
|
||||||
data,
|
|
||||||
offset,
|
|
||||||
length
|
|
||||||
)
|
|
||||||
|
|
||||||
return (
|
|
||||||
data[offset:offset+length],
|
|
||||||
offset + length
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if hasattr(typ, "from_com"):
|
if hasattr(typ, "from_com"):
|
||||||
return typ.from_com(
|
return typ.from_com(data, offset)
|
||||||
data,
|
|
||||||
offset
|
|
||||||
)
|
|
||||||
|
|
||||||
|
raise PWPError(f"Cannot decode type {typ}")
|
||||||
raise PWPError(
|
|
||||||
f"Cannot decode type {typ}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user