Files
micromesh/micromesh/protobuf.py
T
2026-07-31 23:16:45 -07:00

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