311 lines
10 KiB
Python
311 lines
10 KiB
Python
"""Tiny proto3 codec used by MicroMesh.
|
|
|
|
It deliberately implements only wire types used by the public messages below.
|
|
There are no descriptors and no dependency on ``google.protobuf``.
|
|
"""
|
|
|
|
try:
|
|
import struct
|
|
except ImportError: # pragma: no cover - all supported ports normally have it
|
|
import ustruct as struct
|
|
|
|
|
|
VARINT = 0
|
|
FIXED64 = 1
|
|
BYTES = 2
|
|
FIXED32 = 5
|
|
|
|
|
|
class DecodeError(ValueError):
|
|
pass
|
|
|
|
|
|
def encode_varint(value):
|
|
value = int(value)
|
|
if value < 0:
|
|
value &= 0xFFFFFFFFFFFFFFFF
|
|
out = bytearray()
|
|
while value > 0x7F:
|
|
out.append((value & 0x7F) | 0x80)
|
|
value >>= 7
|
|
out.append(value)
|
|
return bytes(out)
|
|
|
|
|
|
def decode_varint(buf, offset=0):
|
|
value = 0
|
|
shift = 0
|
|
length = len(buf)
|
|
while offset < length and shift < 70:
|
|
byte = buf[offset]
|
|
offset += 1
|
|
value |= (byte & 0x7F) << shift
|
|
if not byte & 0x80:
|
|
return value, offset
|
|
shift += 7
|
|
raise DecodeError("truncated or invalid varint")
|
|
|
|
|
|
def _skip(buf, offset, wire):
|
|
if wire == VARINT:
|
|
_, offset = decode_varint(buf, offset)
|
|
return offset
|
|
if wire == FIXED64:
|
|
offset += 8
|
|
elif wire == BYTES:
|
|
size, offset = decode_varint(buf, offset)
|
|
offset += size
|
|
elif wire == FIXED32:
|
|
offset += 4
|
|
else:
|
|
raise DecodeError("unsupported protobuf wire type: %d" % wire)
|
|
if offset > len(buf):
|
|
raise DecodeError("truncated protobuf field")
|
|
return offset
|
|
|
|
|
|
class Field:
|
|
def __init__(self, number, kind="uint", message=None, repeated=False, packed=False, optional=False):
|
|
self.number = number
|
|
self.kind = kind
|
|
self.message = message
|
|
self.repeated = repeated
|
|
self.packed = packed
|
|
self.optional = optional
|
|
|
|
@property
|
|
def wire(self):
|
|
if self.kind in ("bytes", "string", "message"):
|
|
return BYTES
|
|
if self.kind in ("fixed32", "sfixed32", "float"):
|
|
return FIXED32
|
|
if self.kind in ("fixed64", "sfixed64", "double"):
|
|
return FIXED64
|
|
return VARINT
|
|
|
|
|
|
class ProtoMessage:
|
|
FIELDS = {}
|
|
ONEOFS = {}
|
|
_field_numbers = None
|
|
|
|
def __init__(self, **kwargs):
|
|
object.__setattr__(self, "_values", {})
|
|
object.__setattr__(self, "_unknown", [])
|
|
for name, value in kwargs.items():
|
|
setattr(self, name, value)
|
|
|
|
@classmethod
|
|
def _by_number(cls):
|
|
result = cls._field_numbers
|
|
if result is None:
|
|
result = {}
|
|
for name, field in cls.FIELDS.items():
|
|
result[field.number] = (name, field)
|
|
cls._field_numbers = result
|
|
return result
|
|
|
|
def __getattr__(self, name):
|
|
fields = type(self).FIELDS
|
|
if name.endswith("_") and name[:-1] in fields:
|
|
name = name[:-1]
|
|
field = fields.get(name)
|
|
if field is None:
|
|
raise AttributeError(name)
|
|
if name in self._values:
|
|
return self._values[name]
|
|
if field.repeated:
|
|
value = []
|
|
self._values[name] = value
|
|
return value
|
|
if field.kind == "message":
|
|
value = field.message()
|
|
for members in type(self).ONEOFS.values():
|
|
if name in members:
|
|
for other in members:
|
|
if other != name:
|
|
self._values.pop(other, None)
|
|
self._values[name] = value
|
|
return value
|
|
if field.kind == "bytes":
|
|
return b""
|
|
if field.kind == "string":
|
|
return ""
|
|
if field.kind in ("float", "double"):
|
|
return 0.0
|
|
return False if field.kind == "bool" else 0
|
|
|
|
def __setattr__(self, name, value):
|
|
fields = type(self).FIELDS
|
|
if name.endswith("_") and name[:-1] in fields:
|
|
name = name[:-1]
|
|
field = fields.get(name)
|
|
if field is None:
|
|
object.__setattr__(self, name, value)
|
|
return
|
|
if field.kind == "message" and isinstance(value, dict):
|
|
value = field.message(**value)
|
|
for members in type(self).ONEOFS.values():
|
|
if name in members:
|
|
for other in members:
|
|
if other != name:
|
|
self._values.pop(other, None)
|
|
self._values[name] = value
|
|
|
|
def HasField(self, name):
|
|
return name in self._values
|
|
|
|
def ClearField(self, name):
|
|
self._values.pop(name, None)
|
|
|
|
def WhichOneof(self, name):
|
|
for field_name in self.ONEOFS.get(name, ()):
|
|
if field_name in self._values:
|
|
return field_name
|
|
return None
|
|
|
|
def CopyFrom(self, other):
|
|
self.ParseFromString(other.SerializeToString())
|
|
|
|
def SerializeToString(self):
|
|
out = bytearray()
|
|
for name, field in sorted(type(self).FIELDS.items(), key=lambda item: item[1].number):
|
|
if name not in self._values:
|
|
continue
|
|
value = self._values[name]
|
|
in_oneof = any(name in members for members in type(self).ONEOFS.values())
|
|
if not in_oneof and not field.optional and not field.repeated and field.kind != "message":
|
|
if value in (0, False, b"", "", 0.0):
|
|
continue
|
|
values = value if field.repeated else (value,)
|
|
if field.packed and values:
|
|
body = bytearray()
|
|
for item in values:
|
|
body.extend(self._encode_scalar(field, item, include_tag=False))
|
|
out.extend(encode_varint((field.number << 3) | BYTES))
|
|
out.extend(encode_varint(len(body)))
|
|
out.extend(body)
|
|
else:
|
|
for item in values:
|
|
out.extend(encode_varint((field.number << 3) | field.wire))
|
|
out.extend(self._encode_scalar(field, item, include_tag=False))
|
|
for raw in self._unknown:
|
|
out.extend(raw)
|
|
return bytes(out)
|
|
|
|
def _encode_scalar(self, field, value, include_tag=False):
|
|
kind = field.kind
|
|
if kind == "message":
|
|
data = value.SerializeToString()
|
|
return encode_varint(len(data)) + data
|
|
if kind == "string":
|
|
value = value.encode("utf-8")
|
|
if kind in ("bytes", "string"):
|
|
value = bytes(value)
|
|
return encode_varint(len(value)) + value
|
|
if kind == "float":
|
|
return struct.pack("<f", value)
|
|
if kind == "double":
|
|
return struct.pack("<d", value)
|
|
if kind in ("fixed32", "sfixed32"):
|
|
return struct.pack("<I", int(value) & 0xFFFFFFFF)
|
|
if kind in ("fixed64", "sfixed64"):
|
|
return struct.pack("<Q", int(value) & 0xFFFFFFFFFFFFFFFF)
|
|
if kind in ("sint", "sint64"):
|
|
value = (int(value) << 1) ^ (int(value) >> (63 if kind == "sint64" else 31))
|
|
return encode_varint(value)
|
|
|
|
def ParseFromString(self, data):
|
|
object.__setattr__(self, "_values", {})
|
|
object.__setattr__(self, "_unknown", [])
|
|
data = bytes(data)
|
|
offset = 0
|
|
by_number = type(self)._by_number()
|
|
while offset < len(data):
|
|
start = offset
|
|
tag, offset = decode_varint(data, offset)
|
|
number, wire = tag >> 3, tag & 7
|
|
entry = by_number.get(number)
|
|
if entry is None:
|
|
offset = _skip(data, offset, wire)
|
|
self._unknown.append(data[start:offset])
|
|
continue
|
|
name, field = entry
|
|
if wire == BYTES:
|
|
size, offset = decode_varint(data, offset)
|
|
end = offset + size
|
|
if end > len(data):
|
|
raise DecodeError("truncated length-delimited field")
|
|
raw = data[offset:end]
|
|
offset = end
|
|
if field.packed and field.wire != BYTES:
|
|
values = self._values.setdefault(name, [])
|
|
inner = 0
|
|
while inner < len(raw):
|
|
item, inner = self._decode_scalar(field, raw, inner, field.wire)
|
|
values.append(item)
|
|
continue
|
|
value = self._decode_bytes(field, raw)
|
|
else:
|
|
if wire != field.wire:
|
|
offset = _skip(data, offset, wire)
|
|
self._unknown.append(data[start:offset])
|
|
continue
|
|
value, offset = self._decode_scalar(field, data, offset, wire)
|
|
if field.repeated:
|
|
self._values.setdefault(name, []).append(value)
|
|
else:
|
|
setattr(self, name, value)
|
|
return self
|
|
|
|
def _decode_bytes(self, field, raw):
|
|
if field.kind == "message":
|
|
return field.message().ParseFromString(raw)
|
|
if field.kind == "string":
|
|
return raw.decode("utf-8")
|
|
return raw
|
|
|
|
def _decode_scalar(self, field, data, offset, wire):
|
|
kind = field.kind
|
|
if wire == VARINT:
|
|
value, offset = decode_varint(data, offset)
|
|
if kind in ("sint", "sint64"):
|
|
value = (value >> 1) ^ -(value & 1)
|
|
elif kind in ("int", "int64"):
|
|
bits = 64
|
|
if value & (1 << (bits - 1)):
|
|
value -= 1 << bits
|
|
elif kind == "bool":
|
|
value = bool(value)
|
|
return value, offset
|
|
size = 4 if wire == FIXED32 else 8
|
|
if offset + size > len(data):
|
|
raise DecodeError("truncated fixed-width field")
|
|
raw = data[offset:offset + size]
|
|
if kind == "float":
|
|
value = struct.unpack("<f", raw)[0]
|
|
elif kind == "double":
|
|
value = struct.unpack("<d", raw)[0]
|
|
elif kind in ("sfixed32", "sfixed64"):
|
|
value = struct.unpack("<i" if size == 4 else "<q", raw)[0]
|
|
else:
|
|
value = struct.unpack("<I" if size == 4 else "<Q", raw)[0]
|
|
return value, offset + size
|
|
|
|
def to_dict(self):
|
|
result = {}
|
|
for name, value in self._values.items():
|
|
if isinstance(value, ProtoMessage):
|
|
value = value.to_dict()
|
|
elif isinstance(value, list):
|
|
value = [item.to_dict() if isinstance(item, ProtoMessage) else item for item in value]
|
|
result[name] = value
|
|
return result
|
|
|
|
def __repr__(self):
|
|
return "%s(%r)" % (type(self).__name__, self.to_dict())
|
|
|
|
|
|
def fields(**definitions):
|
|
return definitions
|