90 lines
2.8 KiB
Python
90 lines
2.8 KiB
Python
from __future__ import annotations
|
|
|
|
from collections.abc import Iterator
|
|
from dataclasses import dataclass
|
|
|
|
|
|
class ProtobufWireError(ValueError):
|
|
"""Raised when a bounded protobuf wire parse fails."""
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class ProtoField:
|
|
number: int
|
|
wire_type: int
|
|
value: int | bytes
|
|
|
|
|
|
def read_varint(data: bytes, offset: int) -> tuple[int, int]:
|
|
"""Read one protobuf unsigned varint, bounded to 64 bits."""
|
|
value = 0
|
|
for shift in range(0, 70, 7):
|
|
if offset >= len(data):
|
|
raise ProtobufWireError("truncated varint")
|
|
octet = data[offset]
|
|
offset += 1
|
|
if shift == 63 and octet > 1:
|
|
raise ProtobufWireError("varint exceeds 64 bits")
|
|
value |= (octet & 0x7F) << shift
|
|
if not octet & 0x80:
|
|
return value, offset
|
|
raise ProtobufWireError("varint exceeds 10 bytes")
|
|
|
|
|
|
def decode_zigzag64(value: int) -> int:
|
|
"""Decode protobuf sint64 ZigZag representation."""
|
|
if value < 0 or value > 0xFFFFFFFFFFFFFFFF:
|
|
raise ProtobufWireError("ZigZag input is outside uint64")
|
|
return (value >> 1) ^ -(value & 1)
|
|
|
|
|
|
def iter_fields(data: bytes, *, max_fields: int = 1_000_000) -> Iterator[ProtoField]:
|
|
"""Iterate supported protobuf fields without recursion or unbounded allocation."""
|
|
if max_fields < 1:
|
|
raise ValueError("max_fields must be positive")
|
|
|
|
offset = 0
|
|
field_count = 0
|
|
while offset < len(data):
|
|
field_count += 1
|
|
if field_count > max_fields:
|
|
raise ProtobufWireError(f"message exceeds {max_fields} fields")
|
|
|
|
key, offset = read_varint(data, offset)
|
|
number = key >> 3
|
|
wire_type = key & 0x07
|
|
if number == 0:
|
|
raise ProtobufWireError("protobuf field number zero is invalid")
|
|
|
|
if wire_type == 0:
|
|
value, offset = read_varint(data, offset)
|
|
yield ProtoField(number, wire_type, value)
|
|
continue
|
|
|
|
if wire_type == 1:
|
|
end = offset + 8
|
|
if end > len(data):
|
|
raise ProtobufWireError("truncated fixed64 field")
|
|
yield ProtoField(number, wire_type, data[offset:end])
|
|
offset = end
|
|
continue
|
|
|
|
if wire_type == 2:
|
|
length, offset = read_varint(data, offset)
|
|
end = offset + length
|
|
if end > len(data):
|
|
raise ProtobufWireError("truncated length-delimited field")
|
|
yield ProtoField(number, wire_type, data[offset:end])
|
|
offset = end
|
|
continue
|
|
|
|
if wire_type == 5:
|
|
end = offset + 4
|
|
if end > len(data):
|
|
raise ProtobufWireError("truncated fixed32 field")
|
|
yield ProtoField(number, wire_type, data[offset:end])
|
|
offset = end
|
|
continue
|
|
|
|
raise ProtobufWireError(f"unsupported protobuf wire type {wire_type}")
|