From c965d0eb39e8eea0c4acf68addd6144f7a1addcf Mon Sep 17 00:00:00 2001 From: Tristen Wen Date: Mon, 31 Aug 2026 20:30:04 +0800 Subject: [PATCH] Generate specialized encoder functions at compile time Protobuf.encode/1 interpreted MessageProps on every call: per-field dispatch on metadata, presence checks, and IO.iodata_length/1 walks to write length prefixes. For large messages this interpretation overhead dominates encoding time. The DSL now compiles each message's field metadata into a specialized __encode_sized__/1, resolving dispatch, presence checks, field tags, and wire types at compile time. Encoders return {iodata, byte_size} so parents write length prefixes without re-walking the iodata, and map entries are encoded straight from key/value pairs via __encode_entry_sized__/2 without building entry structs. The interpreter remains for messages with a transform module and for values the generated clauses reject, and failed messages are replayed through it so Protobuf.EncodeError still names the failing field. Varint.encode/1 now builds each varint as a single binary in one bit-syntax instruction and rejects integers that don't fit in 64 bits instead of silently truncating them. Encoding is 3-6x faster depending on message shape, with about half the allocations. Wire output is byte-identical; the conformance suite passes unchanged. --- CHANGELOG.md | 6 + lib/protobuf/dsl.ex | 2 + lib/protobuf/dsl/encoder.ex | 663 +++++++++++++++++++++++++++++ lib/protobuf/encoder.ex | 159 ++++++- lib/protobuf/wire/varint.ex | 56 ++- test/protobuf/dsl/encoder_test.exs | 129 ++++++ test/protobuf/wire/varint_test.exs | 10 + test/support/test_msg.ex | 18 + 8 files changed, 1016 insertions(+), 27 deletions(-) create mode 100644 lib/protobuf/dsl/encoder.ex create mode 100644 test/protobuf/dsl/encoder_test.exs diff --git a/CHANGELOG.md b/CHANGELOG.md index e8f1d5e6..96c71c37 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,11 @@ # Changelog +## Unreleased + +### Enhancements + + * Generate specialized encoder functions at compile time, making `Protobuf.encode/1` 3-6x faster. + ## v0.17.0 ### Enhancements diff --git a/lib/protobuf/dsl.ex b/lib/protobuf/dsl.ex index bcd1baf0..574062a0 100644 --- a/lib/protobuf/dsl.ex +++ b/lib/protobuf/dsl.ex @@ -189,6 +189,8 @@ defmodule Protobuf.DSL do unquote(msg_props.enum? && Protobuf.DSL.Enum.quoted_enum_functions(msg_props)) + unquote(Protobuf.DSL.Encoder.quoted_encode_functions(msg_props, transform_module_ast)) + if unquote(Macro.escape(extension_props)) != nil do def __protobuf_info__(:extension_props) do unquote(Macro.escape(extension_props)) diff --git a/lib/protobuf/dsl/encoder.ex b/lib/protobuf/dsl/encoder.ex new file mode 100644 index 00000000..559e57a4 --- /dev/null +++ b/lib/protobuf/dsl/encoder.ex @@ -0,0 +1,663 @@ +defmodule Protobuf.DSL.Encoder do + @moduledoc false + + # Compile-time generation of specialized encoders. + # + # `Protobuf.Encoder` is an interpreter: for every field of every message it + # walks `__message_props__/0`, looks up the field props, dispatches on the + # type and re-evaluates the presence rules. All of that is decided by the + # schema, which is fully known when `use Protobuf` expands, so it can be done + # once at compile time instead — which is what protoc does for Go/Java/Scala. + # + # Each message module gets: + # + # * `__encode_sized__(message)` - returns `{iodata, byte_size}` (or `:skip`), + # so that a parent message can write its length prefix without the extra + # `IO.iodata_length/1` traversal that the interpreter needs. + # + # * `__encode_entry_sized__(key, value)` - only for map entry messages, to + # encode an entry without materializing the entry struct. + # + # Messages with a `transform_module/0` fall back to the interpreted encoder, + # and so do values that aren't a struct of the expected module (implicit + # casts). When the generated encoder raises, `Protobuf.Encoder` replays the + # message through the interpreter so that error messages keep naming the + # field that failed. + + alias Protobuf.{FieldProps, MessageProps} + + @varint_types %{ + int32: {-0x80000000, 0x7FFFFFFF}, + int64: {-0x8000000000000000, 0x7FFFFFFFFFFFFFFF}, + uint32: {0, 0xFFFFFFFF}, + uint64: {0, 0xFFFFFFFFFFFFFFFF} + } + + @zigzag_types %{ + sint32: {-0x80000000, 0x7FFFFFFF}, + sint64: {-0x8000000000000000, 0x7FFFFFFFFFFFFFFF} + } + + # Types for which Protobuf.Wire.encode/2 returns a binary of a known length. + @fixed_types [:fixed32, :sfixed32, :float, :fixed64, :sfixed64, :double] + + @spec quoted_encode_functions(MessageProps.t(), Macro.t() | nil) :: Macro.t() | nil + def quoted_encode_functions(%MessageProps{enum?: true}, _transform_module_ast), do: nil + + def quoted_encode_functions(%MessageProps{} = props, nil) do + block([ + quoted_no_warn_undefined(props), + quoted_message_encoder(props), + quoted_entry_encoder(props) + ]) + end + + def quoted_encode_functions(%MessageProps{} = props, _transform_module_ast) do + block([quoted_fallback_encoder(), quoted_fallback_entry_encoder(props)]) + end + + # The generated code calls child message and enum modules directly, and a + # message compiled on its own (as the protoc generator tests do) may not have + # them at all. + defp quoted_no_warn_undefined(%MessageProps{field_props: field_props}) do + modules = + field_props + |> Map.values() + |> Enum.flat_map(&referenced_modules/1) + |> Enum.uniq() + + if modules != [] do + quote do + @compile {:no_warn_undefined, unquote(modules)} + end + end + end + + defp referenced_modules(%FieldProps{embedded?: true, type: mod}) when is_atom(mod), do: [mod] + defp referenced_modules(%FieldProps{type: {:enum, mod}}), do: [mod] + defp referenced_modules(%FieldProps{}), do: [] + + ## Message encoder + + defp quoted_message_encoder(%MessageProps{} = props) do + fields = ordered_fields(props) + + quote do + @doc false + @spec __encode_sized__(term()) :: {iodata(), non_neg_integer()} | :skip + def __encode_sized__(%__MODULE__{unquote_splicing(quoted_destructure(fields))} = msg) do + unquote(quoted_encode_body(fields, props)) + rescue + error -> + Protobuf.Encoder.reraise_generated_error(msg, __MODULE__, error, __STACKTRACE__) + end + + def __encode_sized__(other) do + Protobuf.Encoder.encode_child_sized(other, __MODULE__) + end + + unquote(block(Enum.map(fields, "ed_field_helper(&1, props)))) + end + end + + defp quoted_fallback_encoder do + quote do + @doc false + @spec __encode_sized__(term()) :: {iodata(), non_neg_integer()} | :skip + def __encode_sized__(msg) do + Protobuf.Encoder.encode_child_sized(msg, __MODULE__) + end + end + end + + defp quoted_entry_encoder(%MessageProps{map?: true} = props) do + key_prop = Map.fetch!(props.field_props, 1) + value_prop = Map.fetch!(props.field_props, 2) + + quote do + @doc false + @spec __encode_entry_sized__(term(), term()) :: {iodata(), non_neg_integer()} | :skip + def __encode_entry_sized__(key, value) do + acc = {[], 0} + unquote(quoted_field(key_prop, props, var(:key))) + unquote(quoted_field(value_prop, props, var(:value))) + acc + rescue + error -> + Protobuf.Encoder.reraise_generated_entry_error( + {key, value}, + __MODULE__, + error, + __STACKTRACE__ + ) + end + + unquote(quoted_field_helper(key_prop, props)) + unquote(quoted_field_helper(value_prop, props)) + end + end + + defp quoted_entry_encoder(%MessageProps{}), do: nil + + defp quoted_fallback_entry_encoder(%MessageProps{map?: true}) do + quote do + @doc false + @spec __encode_entry_sized__(term(), term()) :: {iodata(), non_neg_integer()} | :skip + def __encode_entry_sized__(key, value) do + Protobuf.Encoder.encode_map_entry_sized({key, value}, __MODULE__) + end + end + end + + defp quoted_fallback_entry_encoder(%MessageProps{}), do: nil + + defp quoted_encode_body(fields, %MessageProps{} = props) do + quote do + acc = {[], 0} + unquote(quoted_oneof_values(props)) + unquote_splicing(Enum.map(fields, "ed_field(&1, props, field_var(&1, props)))) + unquote(quoted_unknown_fields()) + unquote(quoted_extensions(props)) + acc + end + end + + # Fields are encoded in tag order because some conformance tests expect it. + defp ordered_fields(%MessageProps{ordered_tags: tags, field_props: field_props}) do + Enum.map(tags, &Map.fetch!(field_props, &1)) + end + + defp quoted_destructure(fields) do + for %FieldProps{oneof: nil} = fp <- fields, do: {fp.name_atom, var(:"v_#{fp.name_atom}")} + end + + # oneof_actual_vals/2 also validates that each {field, value} tuple belongs + # to the oneof group it is stored in. + defp quoted_oneof_values(%MessageProps{oneof: []}), do: nil + + defp quoted_oneof_values(%MessageProps{}) do + quote do + oneofs = Protobuf.Encoder.oneof_actual_vals(__MODULE__.__message_props__(), msg) + end + end + + defp field_var(%FieldProps{oneof: nil} = fp, _props), do: var(:"v_#{fp.name_atom}") + + defp field_var(%FieldProps{name_atom: name}, _props) do + quote do: Map.get(oneofs, unquote(name)) + end + + ## Trailing fields + + defp quoted_unknown_fields do + quote do + acc = + case msg.__unknown_fields__ do + [] -> + acc + + unknown -> + {io, size} = acc + encoded = Protobuf.Encoder.encode_unknown_fields_iodata(unknown) + {[io | encoded], size + IO.iodata_length(encoded)} + end + end + end + + defp quoted_extensions(%MessageProps{syntax: :proto2, extension_range: ranges}) + when not is_nil(ranges) do + quote do + acc = + case msg.__pb_extensions__ do + extensions when extensions == %{} -> + acc + + extensions -> + {io, size} = acc + encoded = Protobuf.Encoder.encode_extensions_iodata(__MODULE__, extensions) + {[io | encoded], size + IO.iodata_length(encoded)} + end + end + end + + defp quoted_extensions(%MessageProps{}), do: nil + + ## Fields + + # Repeated, packed and map fields loop in a generated private function so that + # the loop is a direct local call instead of an anonymous function apply. + defp quoted_field_helper(%FieldProps{} = fp, %MessageProps{} = props) do + cond do + fp.map? -> quoted_map_helper(fp, props) + fp.repeated? and fp.packed? -> quoted_packed_helper(fp, props) + fp.repeated? -> quoted_repeated_helper(fp, props) + true -> nil + end + end + + defp quoted_field(%FieldProps{} = fp, %MessageProps{} = props, val) do + cond do + fp.map? -> quoted_map_field(fp, props, val) + fp.repeated? and fp.packed? -> quoted_packed_field(fp, props, val) + fp.repeated? -> quoted_repeated_field(fp, props, val) + true -> quoted_singular_field(fp, props, val) + end + end + + defp quoted_singular_field(%FieldProps{} = fp, %MessageProps{} = props, val) do + quote do + acc = unquote(case_ast(val, skipped_then_emit_clauses(fp, props, var(:value)))) + end + end + + defp skipped_then_emit_clauses(%FieldProps{embedded?: true} = fp, %MessageProps{} = props, val) do + embedded_clauses(fp, props, val) + end + + # Skip clauses come first, so any emit clause matching a value that is already + # skipped (a proto3 `false`, a proto2 field holding its declared default) would + # be unreachable and is dropped. + defp skipped_then_emit_clauses(%FieldProps{} = fp, %MessageProps{} = props, val) do + patterns = skip_patterns(fp, props) + + emit = + Enum.reject( + emit_clauses(fp, implicit_presence?(fp, props), val), + &covered_by?(&1, patterns) + ) + + Enum.map(patterns, &clause(&1, quote(do: acc))) ++ emit + end + + defp covered_by?({:->, _meta, [[pattern], _body]}, patterns) do + literal_pattern?(pattern) and pattern in patterns + end + + defp literal_pattern?(pattern) do + is_atom(pattern) or is_number(pattern) or is_binary(pattern) + end + + ## Embedded fields + + # A value that looks like a proto3 default (0, "", false) is absent when the + # child module doesn't transform, but present when its transform module turns + # it into a message — only known at runtime, so those values get their own + # clauses. The common nil/struct values never consult the transform module. + defp embedded_clauses(%FieldProps{type: type} = fp, %MessageProps{} = props, val) do + default_clauses = + for default <- embedded_default_patterns(fp, props) do + encoded = + quote(do: Protobuf.Encoder.encode_child_default_sized(unquote(default), unquote(type))) + + clause(default, quoted_embedded_emit(fp, encoded)) + end + + [clause(nil, quote(do: acc)), clause([], quote(do: acc))] ++ + default_clauses ++ + [ + clause( + val, + quoted_embedded_emit(fp, quote(do: unquote(type).__encode_sized__(unquote(val)))) + ) + ] + end + + defp embedded_default_patterns(%FieldProps{} = fp, %MessageProps{} = props) do + if implicit_presence?(fp, props) do + [0, quoted_positive_zero(), "", false] + else + [] + end + end + + defp quoted_embedded_emit(%FieldProps{encoded_fnum: key}, encoded_call) do + quote do + case unquote(encoded_call) do + :skip -> + acc + + {encoded, encoded_size} -> + {io, size} = acc + prefix = Protobuf.Wire.Varint.encode(encoded_size) + + {[io, unquote(key), prefix | encoded], + size + unquote(byte_size(key)) + byte_size(prefix) + encoded_size} + end + end + end + + ## Repeated fields + + defp quoted_repeated_field(%FieldProps{} = fp, %MessageProps{} = props, val) do + fun = helper_name(fp) + emit = quote(do: unquote(fun)(list, acc)) + clauses = skip_clauses(fp, props) ++ [clause(var(:list), emit)] + + quote do + acc = unquote(case_ast(val, clauses)) + end + end + + defp quoted_repeated_helper(%FieldProps{} = fp, %MessageProps{} = props) do + fun = helper_name(fp) + element_clauses = element_clauses(fp, props) + + quote do + defp unquote(fun)([], acc), do: acc + + defp unquote(fun)([element | rest], acc) do + acc = unquote(case_ast(var(:element), element_clauses)) + unquote(fun)(rest, acc) + end + + defp unquote(fun)(other, acc), do: unquote(fun)(Enum.to_list(other), acc) + end + end + + # Embedded elements keep the interpreter's per-element presence check; + # scalar elements are always emitted, the whole list is checked instead. That + # includes the zero value of an enum, which is why elements never have implicit + # presence. + defp element_clauses(%FieldProps{embedded?: true} = fp, %MessageProps{} = props) do + skipped_then_emit_clauses(fp, props, var(:element)) + end + + defp element_clauses(%FieldProps{} = fp, %MessageProps{}) do + emit_clauses(fp, _implicit_presence? = false, var(:element)) + end + + ## Packed fields + + defp quoted_packed_field(%FieldProps{} = fp, %MessageProps{} = props, val) do + fun = helper_name(fp) + key = fp.encoded_fnum + + emit = + quote do + {io, size} = acc + payload = unquote(fun)(list, <<>>) + prefix = Protobuf.Wire.Varint.encode(byte_size(payload)) + + {[io, unquote(key), prefix | payload], + size + unquote(byte_size(key)) + byte_size(prefix) + byte_size(payload)} + end + + clauses = skip_clauses(fp, props) ++ [clause(var(:list), emit)] + + quote do + acc = unquote(case_ast(val, clauses)) + end + end + + # Packed elements are contiguous on the wire, so the whole field becomes one + # appended binary with an O(1) byte_size for its length prefix. + defp quoted_packed_helper(%FieldProps{type: type} = fp, %MessageProps{}) do + fun = helper_name(fp) + + quote do + defp unquote(fun)([], payload), do: payload + + defp unquote(fun)([element | rest], payload) do + encoded = Protobuf.Encoder.encode_wire_binary(unquote(Macro.escape(type)), element) + unquote(fun)(rest, <>) + end + + defp unquote(fun)(other, payload), do: unquote(fun)(Enum.to_list(other), payload) + end + end + + ## Map fields + + defp quoted_map_field(%FieldProps{} = fp, %MessageProps{} = props, val) do + fun = helper_name(fp) + + empty_map_clause = + guarded_clause( + var(:map), + quote(do: is_map(map) and map_size(map) == 0), + quote(do: acc) + ) + + clauses = + skip_clauses(fp, props) ++ + [empty_map_clause, clause(var(:map), quote(do: unquote(fun)(map, acc)))] + + quote do + acc = unquote(case_ast(val, clauses)) + end + end + + defp quoted_map_helper(%FieldProps{type: entry_mod} = fp, %MessageProps{}) do + fun = helper_name(fp) + key = fp.encoded_fnum + + quote do + defp unquote(fun)(map, acc) do + Enum.reduce(map, acc, fn {key, value}, acc -> + case Protobuf.Encoder.encode_entry_sized(unquote(entry_mod), key, value) do + :skip -> + acc + + {encoded, encoded_size} -> + {io, size} = acc + prefix = Protobuf.Wire.Varint.encode(encoded_size) + + {[io, unquote(key), prefix | encoded], + size + unquote(byte_size(key)) + byte_size(prefix) + encoded_size} + end + end) + end + end + end + + ## Presence + + # Mirrors Protobuf.Presence: the patterns are the values for which the + # interpreter's skip_field?/3 returns true, in the same pattern form (so that + # +0.0 and 0.0 keep matching exactly like they do there). + defp skip_clauses(%FieldProps{} = fp, %MessageProps{} = props) do + for pattern <- skip_patterns(fp, props), do: clause(pattern, quote(do: acc)) + end + + defp skip_patterns(%FieldProps{} = fp, %MessageProps{syntax: syntax}) do + cond do + # Embedded values (or a whole repeated/map field of them) are only checked + # for emptiness here, see embedded_clauses/3 for the per-value rules. + fp.embedded? -> + [nil, []] + + not is_nil(fp.oneof) or fp.proto3_optional? -> + [nil, []] + + # Required proto2 fields are emitted even when they hold their default. + syntax == :proto2 and fp.required? -> + [] + + syntax == :proto2 -> + Enum.uniq([nil, [] | List.wrap(fp.default && quoted_default_pattern(fp.default))]) + + true -> + [nil, 0, quoted_positive_zero(), "", false, []] + end + end + + defp implicit_presence?(%FieldProps{proto3_optional?: true}, _props), do: false + defp implicit_presence?(%FieldProps{oneof: oneof}, _props) when not is_nil(oneof), do: false + defp implicit_presence?(_fp, %MessageProps{syntax: :proto3}), do: true + defp implicit_presence?(_fp, %MessageProps{}), do: false + + # Written as a unary plus to avoid the "pattern matching on 0.0" warning. + defp quoted_positive_zero, do: {:+, [], [0.0]} + + # Float zero defaults keep their sign explicit in the pattern: a bare 0.0 + # matches only +0.0 from Erlang/OTP 27 on and warns at compile time. + defp quoted_default_pattern(default) when is_float(default) and default == 0.0 do + case <> do + <<0::1, _::63>> -> quoted_positive_zero() + <<1::1, _::63>> -> {:-, [], [0.0]} + end + end + + defp quoted_default_pattern(default), do: Macro.escape(default) + + ## Value emission + + defp emit_clauses(%FieldProps{type: type} = fp, _implicit_presence?, val) + when type in [:string, :bytes] do + fast = + quote do + unquote(if type == :string, do: quoted_validate_utf8(val)) + {io, size} = acc + length = byte_size(unquote(val)) + prefix = Protobuf.Wire.Varint.encode(length) + + {[io, unquote(fp.encoded_fnum), prefix | unquote(val)], + size + unquote(byte_size(fp.encoded_fnum)) + byte_size(prefix) + length} + end + + [ + guarded_clause(val, quote(do: is_binary(unquote(val))), fast), + clause(val, quoted_wire_emit(fp, val)) + ] + end + + defp emit_clauses(%FieldProps{type: :bool} = fp, _implicit_presence?, val) do + [ + clause(true, quoted_binary_emit(fp, <<1>>)), + clause(false, quoted_binary_emit(fp, <<0>>)), + clause(val, quoted_wire_emit(fp, val)) + ] + end + + defp emit_clauses(%FieldProps{type: {:enum, enum_mod}} = fp, implicit_presence?, val) do + number = var(:number) + + # In proto3, enum fields with implicit presence skip the zero value, which + # can only be known once the atom key has been resolved to its number. + known_clauses = + if implicit_presence? do + [clause(0, quote(do: acc)), clause(number, quoted_varint_emit(fp, number))] + else + [clause(number, quoted_varint_emit(fp, number))] + end + + resolve = + quote do + unquote(case_ast(quote(do: unquote(enum_mod).value(unquote(val))), known_clauses)) + end + + [ + guarded_clause(val, quote(do: is_atom(unquote(val))), resolve), + guarded_clause(val, quote(do: is_integer(unquote(val))), quoted_varint_emit(fp, val)), + clause(val, quoted_wire_emit(fp, val)) + ] + end + + defp emit_clauses(%FieldProps{type: type} = fp, _implicit_presence?, val) + when is_map_key(@varint_types, type) do + {min, max} = Map.fetch!(@varint_types, type) + + [ + guarded_clause( + val, + quote( + do: + is_integer(unquote(val)) and unquote(val) >= unquote(min) and + unquote(val) <= unquote(max) + ), + quoted_varint_emit(fp, val) + ), + clause(val, quoted_wire_emit(fp, val)) + ] + end + + defp emit_clauses(%FieldProps{type: type} = fp, _implicit_presence?, val) + when is_map_key(@zigzag_types, type) do + {min, max} = Map.fetch!(@zigzag_types, type) + zigzagged = quote(do: Protobuf.Wire.Zigzag.encode(unquote(val))) + + [ + guarded_clause( + val, + quote( + do: + is_integer(unquote(val)) and unquote(val) >= unquote(min) and + unquote(val) <= unquote(max) + ), + quoted_varint_emit(fp, zigzagged) + ), + clause(val, quoted_wire_emit(fp, val)) + ] + end + + defp emit_clauses(%FieldProps{type: type} = fp, _implicit_presence?, val) + when type in @fixed_types do + emit = + quote do + {io, size} = acc + encoded = Protobuf.Wire.encode(unquote(type), unquote(val)) + + {[io, unquote(fp.encoded_fnum) | encoded], + size + unquote(byte_size(fp.encoded_fnum)) + byte_size(encoded)} + end + + [clause(val, emit)] + end + + # Unsupported types (groups, and anything a future protoc adds) go through + # Protobuf.Wire, which raises the same error the interpreter would. + defp emit_clauses(%FieldProps{} = fp, _implicit_presence?, val) do + [clause(val, quoted_wire_emit(fp, val))] + end + + defp quoted_varint_emit(%FieldProps{encoded_fnum: key}, number) do + quote do + {io, size} = acc + encoded = Protobuf.Wire.Varint.encode(unquote(number)) + + {[io, unquote(key) | encoded], size + unquote(byte_size(key)) + byte_size(encoded)} + end + end + + defp quoted_binary_emit(%FieldProps{encoded_fnum: key}, binary) do + quote do + {io, size} = acc + + {[io, unquote(key) | unquote(binary)], size + unquote(byte_size(key) + byte_size(binary))} + end + end + + defp quoted_wire_emit(%FieldProps{encoded_fnum: key, type: type}, val) do + quote do + {io, size} = acc + encoded = Protobuf.Encoder.encode_wire(unquote(Macro.escape(type)), unquote(val)) + + {[io, unquote(key) | encoded], size + unquote(byte_size(key)) + IO.iodata_length(encoded)} + end + end + + defp quoted_validate_utf8(val) do + quote do + if not String.valid?(unquote(val)) do + raise Protobuf.EncodeError, + message: "invalid UTF-8 data for type string: #{inspect(unquote(val))}" + end + end + end + + ## AST helpers + + defp helper_name(%FieldProps{fnum: fnum, name_atom: name}), + do: :"__encode_field_#{fnum}_#{name}__" + + defp var(name), do: Macro.var(name, __MODULE__) + + defp block(asts), do: {:__block__, [], Enum.reject(asts, &is_nil/1)} + + defp case_ast(subject, clauses), do: {:case, [], [subject, [do: clauses]]} + + defp clause(pattern, body), do: {:->, [], [[pattern], body]} + + defp guarded_clause(pattern, guard, body), + do: {:->, [], [[{:when, [], [pattern, guard]}], body]} +end diff --git a/lib/protobuf/encoder.ex b/lib/protobuf/encoder.ex index 23f215cd..fc5cb669 100644 --- a/lib/protobuf/encoder.ex +++ b/lib/protobuf/encoder.ex @@ -8,21 +8,120 @@ defmodule Protobuf.Encoder do @spec encode_to_iodata(struct()) :: iodata() def encode_to_iodata(%mod{} = struct) do - struct - |> transform_module(mod) - |> do_encode_to_iodata() + case mod.__encode_sized__(struct) do + {iodata, _size} -> iodata + :skip -> [] + end end @spec encode(struct()) :: binary() - def encode(%mod{} = struct) do + def encode(%_{} = struct) do struct - |> transform_module(mod) - |> encode_with_message_props(mod.__message_props__()) + |> encode_to_iodata() |> IO.iodata_to_binary() end + # Returns the encoded iodata together with its byte size, so that callers + # can write a length prefix without walking the iodata again. + @doc false + @spec encode_child_sized(term(), module()) :: {iodata(), non_neg_integer()} | :skip + def encode_child_sized(value, mod) do + case mod.transform_module() do + nil -> + sized(encode_rejected(mod, value)) + + transform_module -> + case transform_module.encode(value, mod) do + nil -> :skip + transformed -> sized(encode_from_type(mod, transformed)) + end + end + end + + # Only values the generated encoder rejected get here. One of them, a + # struct-tagged map with missing keys, still carries the module's struct tag, + # so encode_from_type/2 would hand it right back to the generated encoder, + # forever. The interpreter reads fields with Map.get/3 and encodes it fine. + defp encode_rejected(mod, %{__struct__: mod} = value) do + encode_with_message_props(value, mod.__message_props__()) + end + + defp encode_rejected(mod, value), do: encode_from_type(mod, value) + + # A value that isn't a message but looks like a proto3 default is only present + # when the child message transforms it into one, see Protobuf.DSL.Encoder. + @doc false + @spec encode_child_default_sized(term(), module()) :: {iodata(), non_neg_integer()} | :skip + def encode_child_default_sized(value, mod) do + if mod.transform_module() do + encode_child_sized(value, mod) + else + :skip + end + end + + # Indirection over __encode_entry_sized__/2: only entry modules with a + # transform module can return :skip, but whether the entry module has one is + # unknown where the call is compiled, and calling it directly makes Dialyzer + # flag the :skip clause as unreachable whenever it doesn't. + @doc false + @spec encode_entry_sized(module(), term(), term()) :: {iodata(), non_neg_integer()} | :skip + def encode_entry_sized(mod, key, value), do: mod.__encode_entry_sized__(key, value) + + # Fallback __encode_entry_sized__/2 for map entry modules that define a + # transform module; the generated entry encoders handle everything else. + @doc false + @spec encode_map_entry_sized({term(), term()}, module()) :: + {iodata(), non_neg_integer()} | :skip + def encode_map_entry_sized({_key, _value} = pair, mod) do + case transform_module(pair, mod) do + nil -> + :skip + + {key, value} -> + entry = struct(mod, %{key: key, value: value}) + sized(encode_with_message_props(entry, mod.__message_props__())) + end + end + + # The generated encoders don't wrap every field in a try/rescue like the + # interpreter does, so on failure the message is replayed through the + # interpreter, which raises the error naming the field that failed. If the + # interpreter is happy with the message, the generated encoder itself is at + # fault and the original error is re-raised. + @doc false + @spec reraise_generated_error(struct(), module(), Exception.t(), Exception.stacktrace()) :: + no_return() + def reraise_generated_error(struct, mod, error, stacktrace) do + _ = encode_with_message_props(struct, mod.__message_props__()) + reraise error, stacktrace + end + + @doc false + @spec reraise_generated_entry_error( + {term(), term()}, + module(), + Exception.t(), + Exception.stacktrace() + ) :: no_return() + def reraise_generated_entry_error({key, value}, mod, error, stacktrace) do + _ = encode_with_message_props(struct(mod, %{key: key, value: value}), mod.__message_props__()) + reraise error, stacktrace + end + + defp sized(iodata), do: {iodata, IO.iodata_length(iodata)} + + # Reached from the interpreted path, whose callers have already applied the + # transform module (if any) for this message. defp do_encode_to_iodata(%mod{} = struct) do - encode_with_message_props(struct, mod.__message_props__()) + if mod.transform_module() do + encode_with_message_props(struct, mod.__message_props__()) + else + case mod.__encode_sized__(struct) do + {iodata, _size} -> iodata + :skip -> [] + end + end end defp encode_with_message_props( @@ -110,7 +209,7 @@ defmodule Protobuf.Encoder do # so that oneof {:atom, val} can be encoded encoded = encode_from_type(type, val) byte_size = IO.iodata_length(encoded) - [fnum | Varint.encode(byte_size)] ++ encoded + [fnum, Varint.encode(byte_size) | encoded] end end) end @@ -121,10 +220,22 @@ defmodule Protobuf.Encoder do else encoded = Enum.map(val, &Wire.encode(type, &1)) byte_size = IO.iodata_length(encoded) - [fnum | Varint.encode(byte_size)] ++ encoded + [fnum, Varint.encode(byte_size) | encoded] end end + # Slow path of the generated encoders, for types not worth specializing for + # (which may not even be supported by Protobuf.Wire). + @doc false + @spec encode_wire(Wire.proto_type(), term()) :: iodata() + def encode_wire(type, value), do: Wire.encode(type, value) + + # Packed elements are appended to the field's payload binary, so the + # generated encoders need each one as a binary rather than as iodata. + @doc false + @spec encode_wire_binary(Wire.proto_type(), term()) :: binary() + def encode_wire_binary(type, value), do: type |> Wire.encode(value) |> IO.iodata_to_binary() + defp encode_from_type(mod, msg) do case msg do %{__struct__: ^mod} -> @@ -159,6 +270,12 @@ defmodule Protobuf.Encoder do end defp encode_unknown_fields(%_{__unknown_fields__: unknown_fields} = _message) do + encode_unknown_fields_iodata(unknown_fields) + end + + @doc false + @spec encode_unknown_fields_iodata([Protobuf.unknown_field()]) :: iodata() + def encode_unknown_fields_iodata(unknown_fields) do Enum.map(unknown_fields, fn {fnum, wire_type, value} -> [encode_fnum(fnum, wire_type), Wire.encode_from_wire_type(wire_type, value)] end) @@ -195,10 +312,12 @@ defmodule Protobuf.Encoder do # string b = 2 # } # Then this could return: %{a: "some value"} - defp oneof_actual_vals( - %MessageProps{field_tags: field_tags, field_props: field_props, oneof: oneof}, - struct - ) do + @doc false + @spec oneof_actual_vals(MessageProps.t(), struct()) :: %{optional(atom()) => term()} + def oneof_actual_vals( + %MessageProps{field_tags: field_tags, field_props: field_props, oneof: oneof}, + struct + ) do Enum.reduce(oneof, %{}, fn {field, index}, acc -> case Map.fetch(struct, field) do {:ok, {field_name, value}} when is_atom(field_name) -> @@ -236,6 +355,16 @@ defmodule Protobuf.Encoder do end defp encode_extensions(%mod{__pb_extensions__: pb_exts}) when is_map(pb_exts) do + encode_extensions_iodata(mod, pb_exts) + end + + defp encode_extensions(_) do + [] + end + + @doc false + @spec encode_extensions_iodata(module(), map()) :: iodata() + def encode_extensions_iodata(mod, pb_exts) do Enum.reduce(pb_exts, [], fn {{ext_mod, key}, val}, acc -> case Protobuf.Extension.get_extension_props(mod, ext_mod, key) do %{field_props: prop} -> @@ -249,8 +378,4 @@ defmodule Protobuf.Encoder do end end) end - - defp encode_extensions(_) do - [] - end end diff --git a/lib/protobuf/wire/varint.ex b/lib/protobuf/wire/varint.ex index 31203f36..37466cc8 100644 --- a/lib/protobuf/wire/varint.ex +++ b/lib/protobuf/wire/varint.ex @@ -35,14 +35,14 @@ defmodule Protobuf.Wire.Varint do # Refer to [efficiency guide](http://www1.erlang.org/doc/efficiency_guide/binaryhandling.html) # for more on efficient binary handling. # - # Encoding on the other hand is simpler. It takes an integer and returns an iolist with its + # Encoding on the other hand is simpler. It takes an integer and returns a binary with its # varint representation: # # iex> Protobuf.Wire.Varint.encode(35) - # [35] + # <<35>> # # iex> Protobuf.Wire.Varint.encode(1_234_567) - # [<<135>>, <<173>>, 75] + # <<135, 173, 75>> import Bitwise @@ -187,17 +187,53 @@ defmodule Protobuf.Wire.Varint do end end - @spec encode(integer) :: iolist + # One clause per encoded length, so that each varint is built as a single + # small binary in one bit-syntax instruction. + @spec encode(integer) :: binary + def encode(n) when n >= 1 <<< 64 or n < -(1 <<< 63) do + raise ArgumentError, "varint-encodable integers must fit in 64 bits, got: #{n}" + end + def encode(n) when n < 0 do <> = <> encode(n) end - def encode(n) when n <= 127 do - [n] - end + def encode(n) when n < 1 <<< 7, do: <> - def encode(n) do - [<<1::1, band(n, 127)::7>> | encode(bsr(n, 7))] - end + def encode(n) when n < 1 <<< 14, do: <<1::1, n::7, bsr(n, 7)>> + + def encode(n) when n < 1 <<< 21, do: <<1::1, n::7, 1::1, bsr(n, 7)::7, bsr(n, 14)>> + + def encode(n) when n < 1 <<< 28, + do: <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, bsr(n, 21)>> + + def encode(n) when n < 1 <<< 35, + do: <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, bsr(n, 28)>> + + def encode(n) when n < 1 <<< 42, + do: + <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, 1::1, + bsr(n, 28)::7, bsr(n, 35)>> + + def encode(n) when n < 1 <<< 49, + do: + <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, 1::1, + bsr(n, 28)::7, 1::1, bsr(n, 35)::7, bsr(n, 42)>> + + def encode(n) when n < 1 <<< 56, + do: + <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, 1::1, + bsr(n, 28)::7, 1::1, bsr(n, 35)::7, 1::1, bsr(n, 42)::7, bsr(n, 49)>> + + def encode(n) when n < 1 <<< 63, + do: + <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, 1::1, + bsr(n, 28)::7, 1::1, bsr(n, 35)::7, 1::1, bsr(n, 42)::7, 1::1, bsr(n, 49)::7, bsr(n, 56)>> + + def encode(n), + do: + <<1::1, n::7, 1::1, bsr(n, 7)::7, 1::1, bsr(n, 14)::7, 1::1, bsr(n, 21)::7, 1::1, + bsr(n, 28)::7, 1::1, bsr(n, 35)::7, 1::1, bsr(n, 42)::7, 1::1, bsr(n, 49)::7, 1::1, + bsr(n, 56)::7, bsr(n, 63)>> end diff --git a/test/protobuf/dsl/encoder_test.exs b/test/protobuf/dsl/encoder_test.exs new file mode 100644 index 00000000..31252821 --- /dev/null +++ b/test/protobuf/dsl/encoder_test.exs @@ -0,0 +1,129 @@ +defmodule Protobuf.DSL.EncoderTest do + use ExUnit.Case, async: true + + alias TestMsg.{ + EnumRepeatedUnpacked, + Ext, + FloatZeroDefault, + Foo, + Foo2, + Link, + MapFoo, + Maps, + Oneof, + OneofProto3, + Proto3Optional, + Scalars, + WithTransformModule + } + + # Encoding of every field kind is covered by Protobuf.EncoderTest. What is + # specific to the generated encoders is the {iodata, byte_size} pair they + # return: parents write their length prefix from that size instead of walking + # the iodata again, so a size that disagrees with the iodata would silently + # corrupt every enclosing message. + describe "__encode_sized__/1" do + test "returns the encoded message together with its byte size" do + for message <- sample_messages() do + %module{} = message + + assert {iodata, size} = module.__encode_sized__(message) + assert size == IO.iodata_length(iodata), "wrong size for #{inspect(message)}" + assert IO.iodata_to_binary(iodata) == Protobuf.encode(message) + end + end + + # A repeated field is absent only when it is empty: presence is never checked + # per element, so a zero enum value inside the list stays on the wire. + test "keeps the zero value of an element of an unpacked repeated enum" do + assert {iodata, size} = + EnumRepeatedUnpacked.__encode_sized__(%EnumRepeatedUnpacked{a: [:UNKNOWN, :A]}) + + assert IO.iodata_to_binary(iodata) == <<8, 0, 8, 1>> + assert size == IO.iodata_length(iodata) + end + + test "encodes a message with a transform module from its transformed value" do + assert {iodata, size} = WithTransformModule.__encode_sized__(42) + assert size == IO.iodata_length(iodata) + assert IO.iodata_to_binary(iodata) == <<8, 42>> + end + + # The generated struct clause can't match a struct-tagged map with missing + # keys, so it must fall back to the interpreter (which reads fields with + # Map.get/3) instead of bouncing between the two encoders forever. + test "encodes a struct with missing keys like the interpreter does" do + malformed = Map.delete(%Foo{a: 42}, :c) + + assert Protobuf.encode(malformed) == Protobuf.encode(%Foo{a: 42}) + end + + test "skips a proto2 float field holding its declared 0.0 default" do + assert Protobuf.encode(%FloatZeroDefault{a: 0.0}) == <<>> + assert Protobuf.encode(%FloatZeroDefault{a: 1.0}) == <<9, 0, 0, 0, 0, 0, 0, 240, 63>> + end + end + + describe "__encode_entry_sized__/2" do + test "encodes a map entry exactly like the entry message itself" do + for {key, value} <- [{"", 0}, {"key", 1}, {"key", -1}] do + assert {iodata, size} = MapFoo.__encode_entry_sized__(key, value) + assert size == IO.iodata_length(iodata) + assert IO.iodata_to_binary(iodata) == Protobuf.encode(%MapFoo{key: key, value: value}) + end + end + end + + defp sample_messages do + [ + %Foo{}, + %Foo{a: 0, c: "", k: false, n: 0.0, j: :UNKNOWN}, + %Foo{ + a: -1, + b: 1234, + c: "foo", + d: 1.5, + e: %Foo.Bar{a: 1, b: "bar"}, + g: [1, 2, 3], + h: [%Foo.Bar{}, %Foo.Bar{a: 2}], + i: [4, 5], + j: :A, + k: true, + l: %{"a" => 1, "b" => 0}, + o: [:A, :B], + p: "deprecated" + }, + %Foo2{a: 0}, + %Foo2{a: 1, b: 5, c: "", e: %Foo.Bar{}, g: [0], i: [1, 2], l: %{}}, + %Scalars{}, + %Scalars{ + string: "s", + bool: true, + float: -0.0, + double: 0.5, + int32: -1, + uint32: 1, + sint32: -1, + fixed32: 1, + sfixed32: -1, + int64: -1, + uint64: 1, + sint64: -1, + fixed64: 1, + sfixed64: -1, + bytes: <<0, 1>>, + repeated_string: ["a", ""], + repeated_bool: [true, false], + repeated_int32: [0, -1] + }, + %Maps{mapii: %{1 => 0}, mapbi: %{false => 1}, mapsi: %{"" => 0}}, + %Oneof{first: {:a, 0}, second: {:d, ""}}, + %Oneof{first: {:e, :UNKNOWN}, other: "other"}, + %OneofProto3{first: {:b, ""}, second: {:c, 0}}, + %Proto3Optional{a: 0, b: "", c: :UNKNOWN}, + %Link{value: 1, child: %Link{child: %Link{value: 2}}}, + %Foo{__unknown_fields__: [{3, 2, "unknown"}]}, + Ext.Foo1.put_extension(%Ext.Foo1{fa: 1}, Ext.PbExtension, :foo2, [1, 2]) + ] + end +end diff --git a/test/protobuf/wire/varint_test.exs b/test/protobuf/wire/varint_test.exs index e15d745b..0ffb9780 100644 --- a/test/protobuf/wire/varint_test.exs +++ b/test/protobuf/wire/varint_test.exs @@ -49,6 +49,16 @@ defmodule Protobuf.Wire.VarintTest do <<255, 255, 255, 255, 255, 255, 255, 255, 255, 1>> end + test "raises for integers that don't fit in 64 bits" do + assert_raise ArgumentError, ~r/must fit in 64 bits/, fn -> + Varint.encode(18_446_744_073_709_551_616) + end + + assert_raise ArgumentError, ~r/must fit in 64 bits/, fn -> + Varint.encode(-9_223_372_036_854_775_809) + end + end + defp encode(n) do n |> Varint.encode() diff --git a/test/support/test_msg.ex b/test/support/test_msg.ex index ca6a6bb7..4938fddb 100644 --- a/test/support/test_msg.ex +++ b/test/support/test_msg.ex @@ -96,6 +96,24 @@ defmodule TestMsg do field :non_matched, 101, type: :int32, optional: true end + # Repeated enums are packed by default, so an unpacked one is needed to check + # that its elements keep the zero value the field itself would drop. + defmodule EnumRepeatedUnpacked do + @moduledoc false + use Protobuf, syntax: :proto3 + + field :a, 1, repeated: true, type: EnumFoo, enum: true, packed: false + end + + # A 0.0 default must compile without the "pattern matching on 0.0" warning + # and still be skipped like any other proto2 declared default. + defmodule FloatZeroDefault do + @moduledoc false + use Protobuf, syntax: :proto2 + + field :a, 1, optional: true, type: :double, default: 0.0 + end + defmodule SignedInt32Repeated do @moduledoc false use Protobuf, syntax: :proto2