Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
176 changes: 109 additions & 67 deletions src/aprotobuf_decoder.erl
Original file line number Diff line number Diff line change
Expand Up @@ -19,12 +19,15 @@
%

-module(aprotobuf_decoder).
-export([parse/2, transform_schema/1]).
-export([parse/2, parse/3, transform_schema/1, transform_schemas/1]).

transform_schema(Schema) ->
Iterator = maps:iterator(Schema),
decode_schema(maps:next(Iterator), #{}).

transform_schemas(Registry) ->
maps:map(fun(_Name, Schema) -> transform_schema(Schema) end, Registry).

decode_schema(none, Acc) ->
Acc;
decode_schema({K, {FieldNum, Type}, I}, Acc) when
Expand All @@ -43,11 +46,31 @@ decode_schema({K, {FieldNum, {repeated, ElemType}}, I}, Acc) when
is_atom(ElemType) and is_integer(FieldNum) and (FieldNum >= 0)
->
decode_schema(maps:next(I), Acc#{FieldNum => {K, {repeated, ElemType}}});
decode_schema({K, {FieldNum, {repeated, {ref, Name}}}, I}, Acc) when
is_atom(Name) and is_integer(FieldNum) and (FieldNum >= 0)
->
decode_schema(maps:next(I), Acc#{FieldNum => {K, {repeated, {ref, Name}}}});
decode_schema({K, {FieldNum, {map, KeyType, ValueType}}, I}, Acc) when
is_atom(KeyType) and is_atom(ValueType) and is_integer(FieldNum) and (FieldNum >= 0)
->
EntrySubSchema = transform_schema(#{key => {1, KeyType}, value => {2, ValueType}}),
decode_schema(maps:next(I), Acc#{FieldNum => {K, {map, EntrySubSchema}}});
decode_schema({K, {FieldNum, {map, KeyType, {ref, Name}}}, I}, Acc) when
is_atom(KeyType) and is_atom(Name) and is_integer(FieldNum) and (FieldNum >= 0)
->
EntrySubSchema = transform_schema(#{key => {1, KeyType}, value => {2, {ref, Name}}}),
decode_schema(maps:next(I), Acc#{FieldNum => {K, {map, EntrySubSchema}}});
decode_schema({K, {FieldNum, {ref, Name}}, I}, Acc) when
is_atom(Name) and is_integer(FieldNum) and (FieldNum >= 0)
->
decode_schema(maps:next(I), Acc#{FieldNum => {K, {ref, Name}}});
decode_schema({K, {oneof, InnerSchema}, I}, Acc) when is_map(InnerSchema) ->
InnerTransformed = transform_schema(InnerSchema),
Expanded = maps:map(
fun(_FN, {Variant, Type}) -> {K, {oneof_variant, Variant, Type}} end,
InnerTransformed
),
decode_schema(maps:next(I), maps:merge(Acc, Expanded));
decode_schema({K, T, _I}, _Acc) ->
error({badarg, K, T}).

Expand All @@ -61,64 +84,76 @@ transform_enum_map({K, V, I}, Acc) ->
transform_enum_map(maps:next(I), Acc#{V => K}).

parse(Bin, Schema) ->
parse(Bin, Schema, tag, #{}).
parse(Bin, root, #{root => Schema}).

parse(<<>>, _Schema, tag, Acc) ->
parse(Bin, EntryName, Registry) ->
Schema = maps:get(EntryName, Registry),
parse_state(Bin, Schema, Registry, tag, #{}).

parse_state(<<>>, _Schema, _Registry, tag, Acc) ->
Acc;
parse(Bin, Schema, What, Acc) ->
parse_state(Bin, Schema, Registry, What, Acc) ->
case What of
tag ->
parse_varint(Bin, 0, 0, value, Schema, Acc);
parse_varint(Bin, 0, 0, value, Schema, Registry, Acc);
value ->
[Tag | _Built] = Acc,
WireType = Tag band 7,
case WireType of
0 -> parse_varint(Bin, 0, 0, end_of_field, Schema, Acc);
1 -> parse_fixed64(Bin, end_of_field, Schema, Acc);
2 -> parse_varint(Bin, 0, 0, len_field_value, Schema, Acc);
0 -> parse_varint(Bin, 0, 0, end_of_field, Schema, Registry, Acc);
1 -> parse_fixed64(Bin, end_of_field, Schema, Registry, Acc);
2 -> parse_varint(Bin, 0, 0, len_field_value, Schema, Registry, Acc);
3 -> {error, unsupported_group};
4 -> {error, invalid};
5 -> parse_fixed32(Bin, end_of_field, Schema, Acc);
5 -> parse_fixed32(Bin, end_of_field, Schema, Registry, Acc);
6 -> {error, unsupported_feature};
7 -> {error, unsupported_feature}
end;
end_of_field ->
[Value, Tag | Built] = Acc,
FieldNum = Tag bsr 3,
{Key, Type} = maps:get(FieldNum, Schema, {x, undefined}),
NewAcc = put_value(Built, Key, Type, Value),
parse(Bin, Schema, tag, NewAcc);
NewAcc = put_value(Built, Key, Type, Value, Registry),
parse_state(Bin, Schema, Registry, tag, NewAcc);
len_field_value ->
[Len, Tag | Built] = Acc,
<<SubBin:Len/binary, Rest/binary>> = Bin,
FieldNum = Tag bsr 3,
{Key, Type} = maps:get(FieldNum, Schema, {x, undefined}),
NewAcc = put_len_value(Built, Key, Type, SubBin),
parse(Rest, Schema, tag, NewAcc)
NewAcc = put_len_value(Built, Key, Type, SubBin, Registry),
parse_state(Rest, Schema, Registry, tag, NewAcc)
end.

put_value(Built, Key, {repeated, ElemType}, Value) ->
put_value(Built, Key, {repeated, ElemType}, Value, Registry) ->
Existing = maps:get(Key, Built, []),
maps:put(Key, Existing ++ [cast(Value, ElemType)], Built);
put_value(Built, Key, Type, Value) ->
maps:put(Key, cast(Value, Type), Built).
maps:put(Key, Existing ++ [cast(Value, ElemType, Registry)], Built);
put_value(Built, Key, {oneof_variant, Variant, Type}, Value, Registry) ->
maps:put(Key, {Variant, cast(Value, Type, Registry)}, Built);
put_value(Built, Key, Type, Value, Registry) ->
maps:put(Key, cast(Value, Type, Registry), Built).

put_len_value(Built, Key, {repeated, ElemType}, Bin) ->
put_len_value(Built, Key, {repeated, {ref, Name}}, Bin, Registry) ->
Schema = maps:get(Name, Registry),
Existing = maps:get(Key, Built, []),
maps:put(Key, Existing ++ [cast(Bin, Schema, Registry)], Built);
put_len_value(Built, Key, {repeated, ElemType}, Bin, Registry) ->
Existing = maps:get(Key, Built, []),
NewElems =
case is_packable(ElemType) of
true -> parse_packed(Bin, ElemType, []);
false -> [cast(Bin, ElemType)]
true -> parse_packed(Bin, ElemType, Registry, []);
false -> [cast(Bin, ElemType, Registry)]
end,
maps:put(Key, Existing ++ NewElems, Built);
put_len_value(Built, Key, {map, EntrySubSchema}, Bin) ->
EntryMap = cast(Bin, EntrySubSchema),
put_len_value(Built, Key, {map, EntrySubSchema}, Bin, Registry) ->
EntryMap = cast(Bin, EntrySubSchema, Registry),
K0 = maps:get(key, EntryMap),
V0 = maps:get(value, EntryMap),
Existing = maps:get(Key, Built, #{}),
maps:put(Key, Existing#{K0 => V0}, Built);
put_len_value(Built, Key, Type, Bin) ->
maps:put(Key, cast(Bin, Type), Built).
put_len_value(Built, Key, {oneof_variant, Variant, Type}, Bin, Registry) ->
maps:put(Key, {Variant, cast(Bin, Type, Registry)}, Built);
put_len_value(Built, Key, Type, Bin, Registry) ->
maps:put(Key, cast(Bin, Type, Registry), Built).

is_packable(int32) -> true;
is_packable(int64) -> true;
Expand All @@ -135,86 +170,93 @@ is_packable(float) -> true;
is_packable(double) -> true;
is_packable(_) -> false.

parse_varint(<<0:1, IntValue:7, Rest/binary>>, IntAcc, Bytes, Next, Schema, Acc) when Bytes =< 9 ->
parse_varint(<<0:1, IntValue:7, Rest/binary>>, IntAcc, Bytes, Next, Schema, Registry, Acc) when
Bytes =< 9
->
VarInt = (IntValue bsl 7 * Bytes) bor IntAcc,
parse(Rest, Schema, Next, [VarInt | Acc]);
parse_varint(<<1:1, IntValue:7, Rest/binary>>, IntAcc, Bytes, Next, Schema, Acc) when Bytes =< 9 ->
parse_state(Rest, Schema, Registry, Next, [VarInt | Acc]);
parse_varint(<<1:1, IntValue:7, Rest/binary>>, IntAcc, Bytes, Next, Schema, Registry, Acc) when
Bytes =< 9
->
VarInt = (IntValue bsl 7 * Bytes) bor IntAcc,
parse_varint(Rest, VarInt, Bytes + 1, Next, Schema, Acc);
parse_varint(Bin, _IntAcc, _Bytes, _Next, _Schema, Acc) ->
parse_varint(Rest, VarInt, Bytes + 1, Next, Schema, Registry, Acc);
parse_varint(Bin, _IntAcc, _Bytes, _Next, _Schema, _Registry, Acc) ->
{invalid, Bin, Acc}.

parse_fixed32(<<Fixed32:4/binary, Rest/binary>>, Next, Schema, Acc) ->
parse(Rest, Schema, Next, [Fixed32 | Acc]).
parse_fixed32(<<Fixed32:4/binary, Rest/binary>>, Next, Schema, Registry, Acc) ->
parse_state(Rest, Schema, Registry, Next, [Fixed32 | Acc]).

parse_fixed64(<<Fixed64:8/binary, Rest/binary>>, Next, Schema, Acc) ->
parse(Rest, Schema, Next, [Fixed64 | Acc]).
parse_fixed64(<<Fixed64:8/binary, Rest/binary>>, Next, Schema, Registry, Acc) ->
parse_state(Rest, Schema, Registry, Next, [Fixed64 | Acc]).

cast(Value, int32) ->
cast(Value, int32, _Registry) ->
case Value bsr 63 of
0 -> Value;
_ -> Value - (1 bsl 64)
end;
cast(Value, int64) ->
cast(Value, int64, _Registry) ->
case Value bsr 63 of
0 -> Value;
_ -> Value - (1 bsl 64)
end;
cast(Value, uint32) ->
cast(Value, uint32, _Registry) ->
Value;
cast(Value, uint64) ->
cast(Value, uint64, _Registry) ->
Value;
cast(Value, sint32) ->
cast(Value, sint32, _Registry) ->
(Value bsr 1) bxor -(Value band 1);
cast(Value, sint64) ->
cast(Value, sint64, _Registry) ->
(Value bsr 1) bxor -(Value band 1);
cast(Value, {enum, IntToLabels}) ->
cast(Value, {enum, IntToLabels}, _Registry) ->
case maps:find(Value, IntToLabels) of
{ok, Label} -> Label;
error -> Value
end;
cast(<<Value:32/integer-little-unsigned>>, fixed32) ->
cast(<<Value:32/integer-little-unsigned>>, fixed32, _Registry) ->
Value;
cast(<<Value:32/integer-little-signed>>, sfixed32) ->
cast(<<Value:32/integer-little-signed>>, sfixed32, _Registry) ->
Value;
cast(<<Value:64/integer-little-unsigned>>, fixed64) ->
cast(<<Value:64/integer-little-unsigned>>, fixed64, _Registry) ->
Value;
cast(<<Value:64/integer-little-signed>>, sfixed64) ->
cast(<<Value:64/integer-little-signed>>, sfixed64, _Registry) ->
Value;
cast(<<X:32/integer-little-unsigned>>, float) when X =:= 16#7F800000 ->
cast(<<X:32/integer-little-unsigned>>, float, _Registry) when X =:= 16#7F800000 ->
infinity;
cast(<<X:32/integer-little-unsigned>>, float) when X =:= 16#FF800000 ->
cast(<<X:32/integer-little-unsigned>>, float, _Registry) when X =:= 16#FF800000 ->
'-infinity';
cast(<<X:32/integer-little-unsigned>>, float) when (X band 16#7F800000) =:= 16#7F800000 ->
cast(<<X:32/integer-little-unsigned>>, float, _Registry) when
(X band 16#7F800000) =:= 16#7F800000
->
nan;
cast(<<Value:32/float-little>>, float) ->
cast(<<Value:32/float-little>>, float, _Registry) ->
Value;
cast(<<X:64/integer-little-unsigned>>, double) when X =:= 16#7FF0000000000000 ->
cast(<<X:64/integer-little-unsigned>>, double, _Registry) when X =:= 16#7FF0000000000000 ->
infinity;
cast(<<X:64/integer-little-unsigned>>, double) when X =:= 16#FFF0000000000000 ->
cast(<<X:64/integer-little-unsigned>>, double, _Registry) when X =:= 16#FFF0000000000000 ->
'-infinity';
cast(<<X:64/integer-little-unsigned>>, double) when
cast(<<X:64/integer-little-unsigned>>, double, _Registry) when
(X band 16#7FF0000000000000) =:= 16#7FF0000000000000
->
nan;
cast(<<Value:64/float-little>>, double) ->
cast(<<Value:64/float-little>>, double, _Registry) ->
Value;
cast(Value, undefined) ->
cast(Value, undefined, _Registry) ->
Value;
cast(Value, bytes) ->
cast(Value, bytes, _Registry) ->
Value;
cast(Value, string) ->
cast(Value, string, _Registry) ->
Value;
cast(Value, bool) ->
cast(Value, bool, _Registry) ->
Value =/= 0;
cast(Bin, {repeated, ElemType}) when is_binary(Bin) ->
parse_packed(Bin, ElemType, []);
cast(Value, Proto) when is_map(Proto) ->
parse(Value, Proto).
cast(Value, {ref, Name}, Registry) ->
Schema = maps:get(Name, Registry),
cast(Value, Schema, Registry);
cast(Value, Proto, Registry) when is_map(Proto) ->
parse_state(Value, Proto, Registry, tag, #{}).

parse_packed(<<>>, _ElemType, Acc) ->
parse_packed(<<>>, _ElemType, _Registry, Acc) ->
lists:reverse(Acc);
parse_packed(Bin, ElemType, Acc) when
parse_packed(Bin, ElemType, Registry, Acc) when
ElemType =:= int32;
ElemType =:= int64;
ElemType =:= uint32;
Expand All @@ -224,15 +266,15 @@ parse_packed(Bin, ElemType, Acc) when
ElemType =:= bool
->
{V, Rest} = parse_packed_varint(Bin, 0, 0),
parse_packed(Rest, ElemType, [cast(V, ElemType) | Acc]);
parse_packed(<<B:4/binary, Rest/binary>>, ElemType, Acc) when
parse_packed(Rest, ElemType, Registry, [cast(V, ElemType, Registry) | Acc]);
parse_packed(<<B:4/binary, Rest/binary>>, ElemType, Registry, Acc) when
ElemType =:= fixed32; ElemType =:= sfixed32; ElemType =:= float
->
parse_packed(Rest, ElemType, [cast(B, ElemType) | Acc]);
parse_packed(<<B:8/binary, Rest/binary>>, ElemType, Acc) when
parse_packed(Rest, ElemType, Registry, [cast(B, ElemType, Registry) | Acc]);
parse_packed(<<B:8/binary, Rest/binary>>, ElemType, Registry, Acc) when
ElemType =:= fixed64; ElemType =:= sfixed64; ElemType =:= double
->
parse_packed(Rest, ElemType, [cast(B, ElemType) | Acc]).
parse_packed(Rest, ElemType, Registry, [cast(B, ElemType, Registry) | Acc]).

parse_packed_varint(<<0:1, V:7, Rest/binary>>, Acc, Bytes) when Bytes =< 9 ->
{(V bsl 7 * Bytes) bor Acc, Rest};
Expand Down
Loading
Loading