Replace `CopyingFieldHandler` with: * `FieldCopier()` for copying a single field. * `DynamicFieldCopier()` for copying a single field or multiple fields determined at runtime. * `AnyFieldCopier()` for copying any field, like `CopyingFieldHandler` before, but taking any `Context...`. PiperOrigin-RevId: 886062533
diff --git a/riegeli/messages/BUILD b/riegeli/messages/BUILD index b219445..735cddb 100644 --- a/riegeli/messages/BUILD +++ b/riegeli/messages/BUILD
@@ -223,6 +223,22 @@ ) cc_library( + name = "field_copier", + hdrs = ["field_copier.h"], + deps = [ + ":message_wire_format", + ":serialized_message_reader", + ":serialized_message_writer", + "//riegeli/base:cord_iterator_span", + "//riegeli/bytes:limiting_reader", + "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/base:nullability", + "@com_google_absl//absl/status", + "@com_google_absl//absl/strings:string_view", + ], +) + +cc_library( name = "serialized_message_writer", srcs = ["serialized_message_writer.cc"], hdrs = ["serialized_message_writer.h"], @@ -289,6 +305,7 @@ srcs = ["serialized_message_assembler.cc"], hdrs = ["serialized_message_assembler.h"], deps = [ + ":field_copier", ":field_handler_map", ":serialized_message_reader", ":serialized_message_writer",
diff --git a/riegeli/messages/field_copier.h b/riegeli/messages/field_copier.h new file mode 100644 index 0000000..796f8cc --- /dev/null +++ b/riegeli/messages/field_copier.h
@@ -0,0 +1,353 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef RIEGELI_MESSAGES_FIELD_COPIER_H_ +#define RIEGELI_MESSAGES_FIELD_COPIER_H_ + +#include <stdint.h> + +#include <optional> +#include <string> +#include <tuple> +#include <type_traits> +#include <utility> + +#include "absl/base/attributes.h" +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "riegeli/base/cord_iterator_span.h" +#include "riegeli/bytes/limiting_reader.h" +#include "riegeli/messages/message_wire_format.h" +#include "riegeli/messages/serialized_message_reader.h" +#include "riegeli/messages/serialized_message_writer.h" + +ABSL_POINTERS_DEFAULT_NONNULL + +namespace riegeli { + +// The type returned by `FieldCopier()`. +template <int field_number, WireTypeSet wire_types = AllWireTypes> +class FieldCopierType; + +// The type returned by `DynamicFieldCopier()`. +template <typename Accept, WireTypeSet wire_types = AllWireTypes> +class DynamicFieldCopierType; + +// The type of the `accept` function used by `AnyFieldCopier()`. +struct AcceptAnyField; + +// The type returned by `AnyFieldCopier()`. +using AnyFieldCopierType = DynamicFieldCopierType<AcceptAnyField>; + +// A field handler for `SerializedMessageReader` which copies the given field to +// a `SerializedMessageWriter`. +// +// As an optimization, `wire_types` constrains the set of wire types to handle. +// This yields smaller and faster code. +// +// `Context...` types must contain exactly one occurrence of +// `SerializedMessageWriter`. Use `ContextProjection()` to select the +// `SerializedMessageWriter` if this is not the case. +template <int field_number, WireTypeSet wire_types = AllWireTypes> +constexpr FieldCopierType<field_number, wire_types> FieldCopier() { + return FieldCopierType<field_number, wire_types>(); +} + +// A field handler for `SerializedMessageReader` which copies the given field to +// a `SerializedMessageWriter`, with the predicate over field numbers specified +// at runtime, possibly mapping them to other field numbers. +// +// `accept` is an invocable taking the field number as `int` and returning +// `std::optional<int>`. If it returns a value other than `std::nullopt`, +// the field is copied, and the returned value is used as the new field number. +// +// `Context...` types must contain exactly one occurrence of +// `SerializedMessageWriter`. Use `ContextProjection()` to select the +// `SerializedMessageWriter` if this is not the case. +template <WireTypeSet wire_types = AllWireTypes, typename Accept> +constexpr DynamicFieldCopierType<std::decay_t<Accept>, wire_types> +DynamicFieldCopier(Accept&& accept) { + return DynamicFieldCopierType<std::decay_t<Accept>, wire_types>( + std::forward<Accept>(accept)); +} + +// A field handler for `SerializedMessageReader` which copies any field to a +// `SerializedMessageWriter`. +// +// It is meant to be used as the last field handler, so that remaining fields +// not handled by previous field handlers will be copied unchanged. +// +// `Context...` types must contain exactly one occurrence of +// `SerializedMessageWriter`. Use `ContextProjection()` to select the +// `SerializedMessageWriter` if this is not the case. +constexpr AnyFieldCopierType AnyFieldCopier(); + +// Implementation details follow. + +template <int field_number, WireTypeSet wire_types> +class FieldCopierType { + public: + static constexpr int kFieldNumber = field_number; + + constexpr FieldCopierType() = default; + + FieldCopierType(const FieldCopierType& that) = default; + FieldCopierType& operator=(const FieldCopierType& that) = default; + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kVarint) == + WireTypeSet::kVarint, + int> = 0> + absl::Status HandleVarint(uint64_t repr, Context&... context) const { + return message_writer(context...).WriteUInt64(field_number, repr); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed32) == + WireTypeSet::kFixed32, + int> = 0> + absl::Status HandleFixed32(uint32_t repr, Context&... context) const { + return message_writer(context...).WriteFixed32(field_number, repr); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed64) == + WireTypeSet::kFixed64, + int> = 0> + absl::Status HandleFixed64(uint64_t repr, Context&... context) const { + return message_writer(context...).WriteFixed64(field_number, repr); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status HandleLengthDelimitedFromReader(ReaderSpan<> repr, + Context&... context) const { + return message_writer(context...) + .WriteString(field_number, std::move(repr)); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status HandleLengthDelimitedFromCord( + CordIteratorSpan repr, ABSL_ATTRIBUTE_UNUSED std::string& scratch, + Context&... context) const { + return message_writer(context...) + .WriteString(field_number, std::move(repr)); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status HandleLengthDelimitedFromString(absl::string_view repr, + Context&... context) const { + return message_writer(context...).WriteString(field_number, repr); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kStartGroup) == + WireTypeSet::kStartGroup, + int> = 0> + absl::Status HandleStartGroup(Context&... context) const { + return message_writer(context...).OpenGroup(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kEndGroup) == + WireTypeSet::kEndGroup, + int> = 0> + absl::Status HandleEndGroup(Context&... context) const { + return message_writer(context...).CloseGroup(field_number); + } + + private: + template <typename... Context> + static SerializedMessageWriter& message_writer(Context&... context) { + return std::get<SerializedMessageWriter&>( + std::tuple<Context&...>(context...)); + } +}; + +template <typename Accept, WireTypeSet wire_types> +class DynamicFieldCopierType { + public: + static constexpr int kFieldNumber = kDynamicFieldNumber; + + template <typename AcceptInitializer, + std::enable_if_t<std::is_convertible_v<AcceptInitializer&&, Accept>, + int> = 0> + explicit constexpr DynamicFieldCopierType(AcceptInitializer&& accept) + : accept_(std::forward<AcceptInitializer>(accept)) {} + + DynamicFieldCopierType() = default; + + DynamicFieldCopierType(const DynamicFieldCopierType& that) = default; + DynamicFieldCopierType& operator=(const DynamicFieldCopierType& that) = + default; + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kVarint) == + WireTypeSet::kVarint, + int> = 0> + std::optional<int> AcceptVarint(int field_number) const { + return accept_(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kVarint) == + WireTypeSet::kVarint, + int> = 0> + absl::Status DynamicHandleVarint(int field_number, uint64_t repr, + Context&... context) const { + return message_writer(context...).WriteUInt64(field_number, repr); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed32) == + WireTypeSet::kFixed32, + int> = 0> + std::optional<int> AcceptFixed32(int field_number) const { + return accept_(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed32) == + WireTypeSet::kFixed32, + int> = 0> + absl::Status DynamicHandleFixed32(int field_number, uint32_t repr, + Context&... context) const { + return message_writer(context...).WriteFixed32(field_number, repr); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed64) == + WireTypeSet::kFixed64, + int> = 0> + std::optional<int> AcceptFixed64(int field_number) const { + return accept_(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kFixed64) == + WireTypeSet::kFixed64, + int> = 0> + absl::Status DynamicHandleFixed64(int field_number, uint64_t repr, + Context&... context) const { + return message_writer(context...).WriteFixed64(field_number, repr); + } + + template <typename... Context> + std::optional<int> AcceptLengthDelimited(int field_number) const { + return accept_(field_number); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status DynamicHandleLengthDelimitedFromReader( + int field_number, ReaderSpan<> repr, Context&... context) const { + return message_writer(context...) + .WriteString(field_number, std::move(repr)); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status DynamicHandleLengthDelimitedFromCord( + int field_number, CordIteratorSpan repr, + ABSL_ATTRIBUTE_UNUSED std::string& scratch, Context&... context) const { + return message_writer(context...) + .WriteString(field_number, std::move(repr)); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kLengthDelimited) == + WireTypeSet::kLengthDelimited, + int> = 0> + absl::Status DynamicHandleLengthDelimitedFromString( + int field_number, absl::string_view repr, Context&... context) const { + return message_writer(context...).WriteString(field_number, repr); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kStartGroup) == + WireTypeSet::kStartGroup, + int> = 0> + std::optional<int> AcceptStartGroup(int field_number) const { + return accept_(field_number); + } + + template < + typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kStartGroup) == + WireTypeSet::kStartGroup, + int> = 0> + absl::Status DynamicHandleStartGroup(int field_number, + Context&... context) const { + return message_writer(context...).OpenGroup(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kEndGroup) == + WireTypeSet::kEndGroup, + int> = 0> + std::optional<int> AcceptEndGroup(int field_number) const { + return accept_(field_number); + } + + template <typename... Context, WireTypeSet dependent_wire_types = wire_types, + std::enable_if_t<(dependent_wire_types & WireTypeSet::kEndGroup) == + WireTypeSet::kEndGroup, + int> = 0> + absl::Status DynamicHandleEndGroup(int field_number, + Context&... context) const { + return message_writer(context...).CloseGroup(field_number); + } + + private: + template <typename... Context> + static SerializedMessageWriter& message_writer(Context&... context) { + return std::get<SerializedMessageWriter&>( + std::tuple<Context&...>(context...)); + } + + ABSL_ATTRIBUTE_NO_UNIQUE_ADDRESS Accept accept_; +}; + +struct AcceptAnyField { + std::optional<int> operator()(int field_number) const { return field_number; } +}; + +constexpr AnyFieldCopierType AnyFieldCopier() { + return AnyFieldCopierType(AcceptAnyField()); +} + +} // namespace riegeli + +#endif // RIEGELI_MESSAGES_FIELD_COPIER_H_
diff --git a/riegeli/messages/message_wire_format.h b/riegeli/messages/message_wire_format.h index 7fc1b38..2a7fc5e 100644 --- a/riegeli/messages/message_wire_format.h +++ b/riegeli/messages/message_wire_format.h
@@ -43,6 +43,47 @@ constexpr WireType GetTagWireType(uint32_t tag); constexpr int GetTagFieldNumber(uint32_t tag); +// Represents a set of wire types. + +enum class WireTypeSet : uint32_t { + kVarint = 1 << static_cast<uint32_t>(WireType::kVarint), + kFixed32 = 1 << static_cast<uint32_t>(WireType::kFixed32), + kFixed64 = 1 << static_cast<uint32_t>(WireType::kFixed64), + kLengthDelimited = 1 << static_cast<uint32_t>(WireType::kLengthDelimited), + kStartGroup = 1 << static_cast<uint32_t>(WireType::kStartGroup), + kEndGroup = 1 << static_cast<uint32_t>(WireType::kEndGroup), + kInvalid6 = 1 << static_cast<uint32_t>(WireType::kInvalid6), + kInvalid7 = 1 << static_cast<uint32_t>(WireType::kInvalid7), +}; + +constexpr WireTypeSet operator|(WireTypeSet a, WireTypeSet b) { + return WireTypeSet{static_cast<uint32_t>(a) | static_cast<uint32_t>(b)}; +} + +constexpr WireTypeSet operator^(WireTypeSet a, WireTypeSet b) { + return WireTypeSet{static_cast<uint32_t>(a) ^ static_cast<uint32_t>(b)}; +} + +constexpr WireTypeSet operator&(WireTypeSet a, WireTypeSet b) { + return WireTypeSet{static_cast<uint32_t>(a) & static_cast<uint32_t>(b)}; +} + +template <WireType... wire_types> +constexpr WireTypeSet WireTypeSetOf() { + return (WireTypeSet{0} | ... | + WireTypeSet{uint32_t{1} << static_cast<uint32_t>(wire_types)}); +} + +constexpr WireTypeSet NoWireTypes = WireTypeSetOf<>(); + +constexpr WireTypeSet AllWireTypes = + WireTypeSetOf<WireType::kVarint, WireType::kFixed32, WireType::kFixed64, + WireType::kLengthDelimited, WireType::kStartGroup, + WireType::kEndGroup, WireType::kInvalid6, + WireType::kInvalid7>(); + +constexpr WireTypeSet operator~(WireTypeSet a) { return AllWireTypes ^ a; } + // Implementation details follow. constexpr uint32_t MakeTag(int field_number, WireType wire_type) {
diff --git a/riegeli/messages/serialized_message_assembler.h b/riegeli/messages/serialized_message_assembler.h index 8f69d10..9e190e1 100644 --- a/riegeli/messages/serialized_message_assembler.h +++ b/riegeli/messages/serialized_message_assembler.h
@@ -43,6 +43,7 @@ #include "riegeli/bytes/reader.h" #include "riegeli/bytes/string_writer.h" #include "riegeli/bytes/writer.h" +#include "riegeli/messages/field_copier.h" #include "riegeli/messages/field_handler_map.h" #include "riegeli/messages/serialized_message_reader.h" #include "riegeli/messages/serialized_message_writer.h" @@ -202,12 +203,6 @@ // The destination message writer. SerializedMessageWriter>; - // Copies unhandled fields. - using CopyingHandler = - CopyingFieldHandler<const absl::Span<FieldValues>, - const absl::Span<const bool>, const absl::Span<bool>, - SerializedMessageWriter>; - // During registration, maintains information about a field or root. struct RegisteredFieldBuilder { ParentForAdd parent_for_add; @@ -561,7 +556,7 @@ SerializedMessageReader< const absl::Span<FieldValues>, const absl::Span<const bool>, const absl::Span<bool>, SerializedMessageWriter>( - std::cref(message.handlers), CopyingHandler()) + std::cref(message.handlers), AnyFieldCopier()) .ReadMessage( std::forward<Src>(src), fields_to_add, fields_to_remove, absl::MakeSpan(submessages_rewritten), message_writer);
diff --git a/riegeli/messages/serialized_message_reader.h b/riegeli/messages/serialized_message_reader.h index daf5983..ec5751b 100644 --- a/riegeli/messages/serialized_message_reader.h +++ b/riegeli/messages/serialized_message_reader.h
@@ -110,7 +110,7 @@ // in namespace `riegeli::field_handlers`. // // The primary dynamic field handlers are `DynamicFieldHandler`, -// `FieldHandlerMap`, and `CopyingFieldHandler`. +// `FieldHandlerMap`, `DynamicFieldCopier`, and `AnyFieldCopier`. // // All field handlers stored in a single `SerializedMessageReader` are usually // conceptually associated with a single message type.
diff --git a/riegeli/messages/serialized_message_reader_internal.h b/riegeli/messages/serialized_message_reader_internal.h index 869c03a..7b36556 100644 --- a/riegeli/messages/serialized_message_reader_internal.h +++ b/riegeli/messages/serialized_message_reader_internal.h
@@ -578,14 +578,7 @@ IsStaticFieldHandlerForFixed64<T, Context...>, IsStaticFieldHandlerForLengthDelimitedFromString<T, Context...>, IsStaticFieldHandlerForStartGroup<T, Context...>, - IsStaticFieldHandlerForEndGroup<T, Context...>>, - std::disjunction< - IsStaticFieldHandlerForLengthDelimitedFromString<T, Context...>, - std::negation<std::disjunction< - IsStaticFieldHandlerForLengthDelimitedFromReader<T, - Context...>, - IsStaticFieldHandlerForLengthDelimitedFromCord< - T, Context...>>>>> {}; + IsStaticFieldHandlerForEndGroup<T, Context...>>> {}; template <typename T, typename... Context> struct IsUnboundFieldHandler
diff --git a/riegeli/messages/serialized_message_writer.h b/riegeli/messages/serialized_message_writer.h index 9b78ae9..5a9ff89 100644 --- a/riegeli/messages/serialized_message_writer.h +++ b/riegeli/messages/serialized_message_writer.h
@@ -18,9 +18,6 @@ #include <stdint.h> #include <limits> -#include <optional> -#include <string> -#include <tuple> #include <type_traits> #include <utility> #include <vector> @@ -31,7 +28,6 @@ #include "absl/base/optimization.h" #include "absl/status/status.h" #include "absl/strings/cord.h" -#include "absl/strings/string_view.h" #include "google/protobuf/message_lite.h" #include "riegeli/base/any.h" #include "riegeli/base/arithmetic.h" @@ -350,100 +346,6 @@ // `writer_ == (submessages_.empty() ? dest_ : &submessages_.back())` }; -// A field handler for `SerializedMessageReader` which copies any field to a -// `SerializedMessageWriter`. -// -// It is meant to be used as the last field handler, so that remaining fields -// not handled by previous field handlers will be copied unchanged. -// -// `Context` types must contain exactly one occurrence of -// `SerializedMessageWriter`. -template <typename... Context> -class CopyingFieldHandler { - public: - // This is `kDynamicFieldNumber` from `serialized_message_reader.h`. - // Avoid adding a dependency just for that. - static constexpr int kFieldNumber = -1; - - CopyingFieldHandler() = default; - - CopyingFieldHandler(const CopyingFieldHandler&) = default; - CopyingFieldHandler& operator=(const CopyingFieldHandler&) = default; - - std::optional<int> AcceptVarint(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleVarint(int field_number, uint64_t repr, - Context&... context) const { - return message_writer(context...).WriteUInt64(field_number, repr); - } - - std::optional<int> AcceptFixed32(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleFixed32(int field_number, uint32_t repr, - Context&... context) const { - return message_writer(context...).WriteFixed32(field_number, repr); - } - - std::optional<int> AcceptFixed64(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleFixed64(int field_number, uint64_t repr, - Context&... context) const { - return message_writer(context...).WriteFixed64(field_number, repr); - } - - std::optional<int> AcceptLengthDelimited(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleLengthDelimitedFromReader( - int field_number, ReaderSpan<> repr, Context&... context) const { - return message_writer(context...) - .WriteString(field_number, std::move(repr)); - } - - absl::Status DynamicHandleLengthDelimitedFromCord( - int field_number, CordIteratorSpan repr, - ABSL_ATTRIBUTE_UNUSED std::string& scratch, Context&... context) const { - return message_writer(context...) - .WriteString(field_number, std::move(repr)); - } - - absl::Status DynamicHandleLengthDelimitedFromString( - int field_number, absl::string_view repr, Context&... context) const { - return message_writer(context...).WriteString(field_number, repr); - } - - std::optional<int> AcceptStartGroup(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleStartGroup(int field_number, - Context&... context) const { - return message_writer(context...).OpenGroup(field_number); - } - - std::optional<int> AcceptEndGroup(int field_number) const { - return field_number; - } - - absl::Status DynamicHandleEndGroup(int field_number, - Context&... context) const { - return message_writer(context...).CloseGroup(field_number); - } - - private: - static SerializedMessageWriter& message_writer(Context&... context) { - return std::get<SerializedMessageWriter&>( - std::tuple<Context&...>(context...)); - } -}; - // Implementation details follow. inline void SerializedMessageWriter::set_dest(Writer* absl_nullable dest) {