Add support for writing packed repeated fields to a preallocated array rather than to a `Writer`. This avoids a bounds check for each element. Cosmetics: write bools with specialized code instead of as varints. This avoids relying on the compiler to statically compute lengths of these varints, which is only possible due to their limited range of [0..1]. PiperOrigin-RevId: 906770507
diff --git a/riegeli/messages/serialized_message_backward_writer.h b/riegeli/messages/serialized_message_backward_writer.h index 3da4e4e..4a17c2d 100644 --- a/riegeli/messages/serialized_message_backward_writer.h +++ b/riegeli/messages/serialized_message_backward_writer.h
@@ -332,7 +332,13 @@ inline absl::Status SerializedMessageBackwardWriter::WriteBool(int field_number, bool value) { - return WriteUInt32(field_number, value ? 1 : 0); + const uint32_t tag = MakeTag(field_number, WireType::kVarint); + const size_t length = LengthVarint32(tag) + 1; + if (ABSL_PREDICT_FALSE(!writer().Push(length))) return writer().status(); + writer().move_cursor(length); + char* const ptr = WriteVarint32(tag, writer().cursor()); + *ptr = value ? '\1' : '\0'; + return absl::OkStatus(); } inline absl::Status SerializedMessageBackwardWriter::WriteFixed32( @@ -424,7 +430,10 @@ inline absl::Status SerializedMessageBackwardWriter::WritePackedBool( bool value) { - return WritePackedUInt32(value ? 1 : 0); + if (ABSL_PREDICT_FALSE(!writer().Write(value ? '\1' : '\0'))) { + return writer().status(); + } + return absl::OkStatus(); } inline absl::Status SerializedMessageBackwardWriter::WritePackedFixed32(
diff --git a/riegeli/messages/serialized_message_writer.h b/riegeli/messages/serialized_message_writer.h index fea8f1a..9a7c03e 100644 --- a/riegeli/messages/serialized_message_writer.h +++ b/riegeli/messages/serialized_message_writer.h
@@ -333,6 +333,29 @@ // This is useful for `WriteLengthUnchecked()`. static Position LengthOfOpenPlusCloseGroup(int field_number); + // Writes an element of a packed repeated field to an array. + // + // The field must have been opened with `WriteLengthUnchecked()` and + // `writer().Push()`. + static char* WritePackedInt32(int32_t value, char* dest); + static char* WritePackedInt64(int64_t value, char* dest); + static char* WritePackedUInt32(uint32_t value, char* dest); + static char* WritePackedUInt64(uint64_t value, char* dest); + static char* WritePackedSInt32(int32_t value, char* dest); + static char* WritePackedSInt64(int64_t value, char* dest); + static char* WritePackedBool(bool value, char* dest); + static char* WritePackedFixed32(uint32_t value, char* dest); + static char* WritePackedFixed64(uint64_t value, char* dest); + static char* WritePackedSFixed32(int32_t value, char* dest); + static char* WritePackedSFixed64(int64_t value, char* dest); + static char* WritePackedFloat(float value, char* dest); + static char* WritePackedDouble(double value, char* dest); + template <typename EnumType, + std::enable_if_t<std::disjunction_v<std::is_enum<EnumType>, + std::is_integral<EnumType>>, + int> = 0> + static char* WritePackedEnum(EnumType value, char* dest); + private: ABSL_ATTRIBUTE_COLD static absl::Status LengthOverflowError(Position length); ABSL_ATTRIBUTE_COLD static absl::Status WriteStringFailed(Reader& src, @@ -431,7 +454,19 @@ inline absl::Status SerializedMessageWriter::WriteBool(int field_number, bool value) { - return WriteUInt32(field_number, value ? 1 : 0); + const uint32_t tag = MakeTag(field_number, WireType::kVarint); + if (ABSL_PREDICT_FALSE(!writer().Push( + (RIEGELI_IS_CONSTANT(tag) || + (RIEGELI_IS_CONSTANT(tag < 0x80) && tag < 0x80) + ? LengthVarint32(tag) + : kMaxLengthVarint32) + + 1))) { + return writer().status(); + } + char* ptr = WriteVarint32(tag, writer().cursor()); + *ptr++ = value ? '\1' : '\0'; + writer().set_cursor(ptr); + return absl::OkStatus(); } inline absl::Status SerializedMessageWriter::WriteFixed32(int field_number, @@ -530,7 +565,10 @@ } inline absl::Status SerializedMessageWriter::WritePackedBool(bool value) { - return WritePackedUInt32(value ? 1 : 0); + if (ABSL_PREDICT_FALSE(!writer().Write(value ? '\1' : '\0'))) { + return writer().status(); + } + return absl::OkStatus(); } inline absl::Status SerializedMessageWriter::WritePackedFixed32( @@ -861,6 +899,82 @@ return 2 * LengthVarint32(MakeTag(field_number, WireType::kStartGroup)); } +inline char* SerializedMessageWriter::WritePackedInt32(int32_t value, + char* dest) { + return WritePackedUInt64(static_cast<uint64_t>(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedInt64(int64_t value, + char* dest) { + return WritePackedUInt64(static_cast<uint64_t>(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedUInt32(uint32_t value, + char* dest) { + return WriteVarint32(value, dest); +} + +inline char* SerializedMessageWriter::WritePackedUInt64(uint64_t value, + char* dest) { + return WriteVarint64(value, dest); +} + +inline char* SerializedMessageWriter::WritePackedSInt32(int32_t value, + char* dest) { + return WriteVarint32(EncodeVarintSigned32(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedSInt64(int64_t value, + char* dest) { + return WriteVarint64(EncodeVarintSigned64(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedBool(bool value, char* dest) { + *dest = value ? '\1' : '\0'; + return dest + 1; +} + +inline char* SerializedMessageWriter::WritePackedFixed32(uint32_t value, + char* dest) { + WriteLittleEndian<uint32_t>(value, dest); + return dest + sizeof(uint32_t); +} + +inline char* SerializedMessageWriter::WritePackedFixed64(uint64_t value, + char* dest) { + WriteLittleEndian<uint64_t>(value, dest); + return dest + sizeof(uint64_t); +} + +inline char* SerializedMessageWriter::WritePackedSFixed32(int32_t value, + char* dest) { + return WritePackedFixed32(static_cast<uint32_t>(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedSFixed64(int64_t value, + char* dest) { + return WritePackedFixed64(static_cast<uint64_t>(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedFloat(float value, + char* dest) { + return WritePackedFixed32(absl::bit_cast<uint32_t>(value), dest); +} + +inline char* SerializedMessageWriter::WritePackedDouble(double value, + char* dest) { + return WritePackedFixed64(absl::bit_cast<uint64_t>(value), dest); +} + +template <typename EnumType, + std::enable_if_t<std::disjunction_v<std::is_enum<EnumType>, + std::is_integral<EnumType>>, + int>> +inline char* SerializedMessageWriter::WritePackedEnum(EnumType value, + char* dest) { + return WritePackedUInt64(static_cast<uint64_t>(value), dest); +} + } // namespace riegeli #endif // RIEGELI_MESSAGES_SERIALIZED_MESSAGE_WRITER_H_