From 76e9ae755c43035025207a58952bafb78c29fd0a Mon Sep 17 00:00:00 2001 From: Gavin Morrow Date: Sun, 21 Sep 2025 20:54:31 -0400 Subject: [PATCH] Pass position to `ParseError` --- src/protobuf_decode_gleam.gleam | 88 ++++++++++++++++++++++----------- 1 file changed, 58 insertions(+), 30 deletions(-) diff --git a/src/protobuf_decode_gleam.gleam b/src/protobuf_decode_gleam.gleam index 69871ff..7418b99 100644 --- a/src/protobuf_decode_gleam.gleam +++ b/src/protobuf_decode_gleam.gleam @@ -14,19 +14,23 @@ pub fn parse( from bits: BitArray, using decoder: Decoder(t), ) -> Result(t, ParseError) { - use data <- result.try(read_fields(bits, [])) + use data <- result.try(read_fields(bits, [], 0)) decode.run(data, decoder) |> result.map_error(UnableToDecode) } pub type ParseError { - UnknownWireType(Int) - InvalidVarInt(leftover_bits: BitArray, acc: BitArray) - InvalidFixed(size: Int, bits: BitArray) - InvalidLen(len: Int, value_bits: BitArray) + UnknownWireType(Int, pos: BytePos) + InvalidVarInt(leftover_bits: BitArray, acc: BitArray, pos: BytePos) + InvalidFixed(size: Int, bits: BitArray, pos: BytePos) + InvalidLen(len: Int, value_bits: BitArray, pos: BytePos) UnableToDecode(List(decode.DecodeError)) } -fn read_fields(bits: BitArray, acc: List(Field)) -> Result(Dynamic, ParseError) { +fn read_fields( + bits: BitArray, + acc: List(Field), + pos: BytePos, +) -> Result(Dynamic, ParseError) { case bits { <<>> -> acc @@ -35,8 +39,11 @@ fn read_fields(bits: BitArray, acc: List(Field)) -> Result(Dynamic, ParseError) |> dynamic.properties |> Ok bits -> { - use Parsed(value: prop, rest: bits) <- result.try(read_field(bits)) - read_fields(bits, [prop, ..acc]) + use Parsed(value: prop, rest: bits, pos:) <- result.try(read_field( + bits, + pos, + )) + read_fields(bits, [prop, ..acc], pos) } } } @@ -72,20 +79,23 @@ fn repeated_to_list(fields: List(Field)) -> dict.Dict(Dynamic, Dynamic) { type DecodeResult(t) = Result(t, ParseError) +type BytePos = + Int + type Parsed(t) { - Parsed(value: t, rest: BitArray) + Parsed(value: t, rest: BitArray, pos: BytePos) } fn parsed_map(of parsed: Parsed(t), with fun: fn(t) -> u) -> Parsed(u) { - let Parsed(value:, rest:) = parsed - Parsed(value: fun(value), rest:) + let Parsed(value:, rest:, pos:) = parsed + Parsed(value: fun(value), rest:, pos:) } type Field { Field(key: Dynamic, value: Dynamic) } -fn wire_type_read_fn(ty: WireType) -> fn(BitArray) -> ValueResult { +fn wire_type_read_fn(ty: WireType) -> fn(BitArray, BytePos) -> ValueResult { case ty { wire_type.VarInt -> read_varint wire_type.I64 -> read_fixed(64) @@ -94,21 +104,31 @@ fn wire_type_read_fn(ty: WireType) -> fn(BitArray) -> ValueResult { } } -fn read_field(bits: BitArray) -> DecodeResult(Parsed(Field)) { - use Parsed(value: tag, rest: bits) <- result.try(read_varint(bits)) +fn read_field(bits: BitArray, tag_pos: BytePos) -> DecodeResult(Parsed(Field)) { + use Parsed(value: tag, rest: bits, pos:) <- result.try(read_varint( + bits, + tag_pos, + )) let tag = util.bit_array_to_uint(tag) let field_id = tag |> int.bitwise_shift_right(3) let wire_type = tag |> int.bitwise_and(0b111) + case wire_type { + 6 -> { + echo tag as "tag" + Nil + } + _ -> Nil + } use wire_type <- result.try(option.to_result( wire_type.parse(wire_type), - UnknownWireType(wire_type), + UnknownWireType(wire_type, pos: tag_pos), )) let read_fn = wire_type_read_fn(wire_type) use value: Parsed(Dynamic) <- result.try({ - use value <- result.map(read_fn(bits)) + use value <- result.map(read_fn(bits, pos)) parsed_map(value, dynamic.bit_array) }) @@ -121,35 +141,41 @@ fn read_field(bits: BitArray) -> DecodeResult(Parsed(Field)) { type ValueResult = DecodeResult(Parsed(BitArray)) -fn read_varint(bits: BitArray) -> ValueResult { - read_varint_acc(bits, <<>>) +fn read_varint(bits: BitArray, pos: BytePos) -> ValueResult { + read_varint_acc(bits, <<>>, pos) } -fn read_varint_acc(bits: BitArray, acc: BitArray) -> ValueResult { +fn read_varint_acc(bits: BitArray, acc: BitArray, pos: BytePos) -> ValueResult { case bits { <<0:size(1), n:bits-size(7), rest:bytes>> -> { let acc = bit_array.concat([n, acc]) - Ok(Parsed(value: acc, rest:)) + Ok(Parsed(value: acc, rest:, pos: pos + 1)) } <<1:size(1), n:bits-size(7), rest:bytes>> -> - read_varint_acc(rest, bit_array.concat([n, acc])) - bits -> Error(InvalidVarInt(leftover_bits: bits, acc:)) + read_varint_acc(rest, bit_array.concat([n, acc]), pos + 1) + bits -> Error(InvalidVarInt(leftover_bits: bits, acc:, pos:)) } } -fn read_fixed(size: Int) -> fn(BitArray) -> ValueResult { - fn(bits: BitArray) -> ValueResult { +/// Size must be a multiple of 8. +fn read_fixed(size: Int) -> fn(BitArray, BytePos) -> ValueResult { + assert size % 8 == 0 + fn(bits: BitArray, pos: BytePos) -> ValueResult { case bits { - <> -> Ok(Parsed(value: num, rest:)) - bits -> Error(InvalidFixed(size:, bits:)) + <> -> + Ok(Parsed(value: num, rest:, pos: pos + size / 8)) + bits -> Error(InvalidFixed(size:, bits:, pos:)) } } } -fn read_len(bits: BitArray) -> ValueResult { +fn read_len(bits: BitArray, len_pos: BytePos) -> ValueResult { // First, read the length of the value // It is encoded as a varint immediately after the tag - use Parsed(value: len, rest: bits) <- result.try(read_varint(bits)) + use Parsed(value: len, rest: bits, pos:) <- result.try(read_varint( + bits, + len_pos, + )) let len = util.bit_array_to_uint(len) // Just decoded a uint, so should be safe @@ -157,11 +183,13 @@ fn read_len(bits: BitArray) -> ValueResult { use value <- result.try( bit_array.slice(from: bits, at: 0, take: len) - |> result.map_error(fn(_) { InvalidLen(len:, value_bits: bits) }), + |> result.map_error(fn(_) { + InvalidLen(len:, value_bits: bits, pos: len_pos) + }), ) // Assert b/c if the len was too long, it would've errored in the prev slice let assert Ok(rest) = bit_array.slice(from: bits, at: len, take: bit_array.byte_size(bits) - len) - Ok(Parsed(value:, rest:)) + Ok(Parsed(value:, rest:, pos: pos + len)) } -- 2.51.2