Add support for optimizing reading packed repeated fields. This is not used automatically for now. PiperOrigin-RevId: 896261060
diff --git a/riegeli/messages/field_handlers.h b/riegeli/messages/field_handlers.h index 930d878..786469d 100644 --- a/riegeli/messages/field_handlers.h +++ b/riegeli/messages/field_handlers.h
@@ -64,6 +64,8 @@ class OnOptionalFixedType; template <typename Value, int field_number, typename Action> class OnRepeatedFixedType; +template <typename BaseFieldHandler, typename PackedFieldHandler> +class OnPackedType; template <int field_number, typename Action> class OnLengthDelimitedType; template <int field_number, typename Action> @@ -387,6 +389,25 @@ std::forward<Action>(action)); } +// Field handler with a dedicated implementation for a packed repeated field. +// +// Uses `BaseFieldHandler` for scalar wire types, and `PackedFieldHandler` for +// length-delimited wire type. +// +// Regular `OnRepeated...()` field handlers already support packed repeated +// fields, but they call the base action repeatedly for each element, which +// can be less efficient than a dedicated implementation. +template <typename BaseFieldHandler, typename PackedFieldHandler> +constexpr OnPackedType<std::decay_t<BaseFieldHandler>, + std::decay_t<PackedFieldHandler>> +OnPacked(BaseFieldHandler&& base_field_handler, + PackedFieldHandler&& packed_action) { + return OnPackedType<std::decay_t<BaseFieldHandler>, + std::decay_t<PackedFieldHandler>>( + std::forward<BaseFieldHandler>(base_field_handler), + std::forward<PackedFieldHandler>(packed_action)); +} + // Field handler of a singular or an element of a repeated `string`, `bytes`, // or submessage field. // @@ -673,6 +694,144 @@ int> = 0> absl::Status HandleLengthDelimitedFromString(absl::string_view repr, Context&... context) const; + + private: + template <typename... Context> + absl::Status HandleLengthDelimitedFromStringInternal( + absl::string_view repr, Context&... context) const { + const char* cursor = repr.data(); + const char* const limit = repr.data() + repr.size(); + while (cursor < limit) { + const Value element = ReadLittleEndian<Value>(cursor); + if (absl::Status status = this->action()(element, context...); + ABSL_PREDICT_FALSE(!status.ok())) { + return status; + } + cursor += sizeof(Value); + } + return absl::OkStatus(); + } +}; + +template <typename BaseFieldHandler, typename PackedFieldHandler> +class OnPackedType { + private: + template <typename... Context> + struct ExpectedFieldHandlers + : std::conjunction< + std::disjunction< + serialized_message_reader_internal:: + IsStaticFieldHandlerForVarint<BaseFieldHandler, Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForFixed32<BaseFieldHandler, + Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForFixed64<BaseFieldHandler, + Context...>>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForLengthDelimited<PackedFieldHandler, + Context...>> {}; + + public: + static constexpr int kFieldNumber = BaseFieldHandler::kFieldNumber; + + template < + typename BaseFieldHandlerInitializer, + typename PackedFieldHandlerInitializer, + std::enable_if_t<std::conjunction_v< + std::is_convertible<BaseFieldHandlerInitializer&&, + BaseFieldHandler>, + std::is_convertible<PackedFieldHandlerInitializer&&, + PackedFieldHandler>>, + int> = 0> + explicit constexpr OnPackedType( + BaseFieldHandlerInitializer&& base_field_handler, + PackedFieldHandlerInitializer&& packed_action) + : base_field_handler_( + std::forward<BaseFieldHandlerInitializer>(base_field_handler)), + packed_field_handler_( + std::forward<PackedFieldHandlerInitializer>(packed_action)) {} + + template < + typename... Context, + std::enable_if_t< + std::conjunction_v< + ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal::IsStaticFieldHandlerForVarint< + BaseFieldHandler, Context...>>, + int> = 0> + absl::Status HandleVarint(uint64_t repr, Context&... context) const { + return base_field_handler_.HandleVarint(repr, context...); + } + + template <typename... Context, + std::enable_if_t< + std::conjunction_v<ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForFixed32< + BaseFieldHandler, Context...>>, + int> = 0> + absl::Status HandleFixed32(uint32_t repr, Context&... context) const { + return base_field_handler_.HandleFixed32(repr, context...); + } + + template <typename... Context, + std::enable_if_t< + std::conjunction_v<ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForFixed64< + BaseFieldHandler, Context...>>, + int> = 0> + absl::Status HandleFixed64(uint64_t repr, Context&... context) const { + return base_field_handler_.HandleFixed64(repr, context...); + } + + template < + typename... Context, + std::enable_if_t<std::conjunction_v< + ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForLengthDelimitedFromReader< + PackedFieldHandler, Context...>>, + int> = 0> + absl::Status HandleLengthDelimitedFromReader(ReaderSpan<> repr, + Context&... context) const { + return packed_field_handler_.HandleLengthDelimitedFromReader( + std::move(repr), context...); + } + + template < + typename... Context, + std::enable_if_t< + std::conjunction_v<ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForLengthDelimitedFromCord< + PackedFieldHandler, Context...>>, + int> = 0> + absl::Status HandleLengthDelimitedFromCord(CordIteratorSpan repr, + std::string& scratch, + Context&... context) const { + return packed_field_handler_.HandleLengthDelimitedFromCord( + std::move(repr), scratch, context...); + } + + template < + typename... Context, + std::enable_if_t<std::conjunction_v< + ExpectedFieldHandlers<Context...>, + serialized_message_reader_internal:: + IsStaticFieldHandlerForLengthDelimitedFromString< + PackedFieldHandler, Context...>>, + int> = 0> + absl::Status HandleLengthDelimitedFromString(absl::string_view repr, + Context&... context) const { + return packed_field_handler_.HandleLengthDelimitedFromString(repr, + context...); + } + + private: + ABSL_ATTRIBUTE_NO_UNIQUE_ADDRESS BaseFieldHandler base_field_handler_; + ABSL_ATTRIBUTE_NO_UNIQUE_ADDRESS PackedFieldHandler packed_field_handler_; }; template <int field_number, typename Action> @@ -855,6 +1014,20 @@ absl::Status OnRepeatedVarintType<Value, kind, field_number, Action>:: HandleLengthDelimitedFromReader(ReaderSpan<> repr, Context&... context) const { + if (repr.reader().Pull(1, IntCast<size_t>(repr.length())) && + repr.reader().available() >= IntCast<size_t>(repr.length())) { + const absl::string_view value(repr.reader().cursor(), + IntCast<size_t>(repr.length())); + repr.reader().move_cursor(IntCast<size_t>(repr.length())); + absl::Status status = HandleLengthDelimitedFromString(value, context...); + // Comparison against `absl::CancelledError()` is a fast path of + // `absl::IsCancelled()`. + if (ABSL_PREDICT_FALSE(!status.ok() && status != absl::CancelledError())) { + status = field_handlers_internal::AnnotateByReader(std::move(status), + repr.reader()); + } + return status; + } ScopedLimiter scoped_limiter(repr); uint64_t element; while (ReadVarint64(repr.reader(), element)) { @@ -890,6 +1063,15 @@ HandleLengthDelimitedFromCord(CordIteratorSpan repr, ABSL_ATTRIBUTE_UNUSED std::string& scratch, Context&... context) const { + if (const absl::string_view chunk = + absl::Cord::ChunkRemaining(repr.iterator()); + chunk.size() >= IntCast<size_t>(repr.length())) { + const absl::string_view value = + chunk.substr(0, IntCast<size_t>(repr.length())); + absl::Cord::AdvanceAndRead(&repr.iterator(), + IntCast<size_t>(repr.length())); + return HandleLengthDelimitedFromString(value, context...); + } const size_t limit = CordIteratorSpan::Remaining(repr.iterator()) - repr.length(); uint64_t element; @@ -955,6 +1137,21 @@ return field_handlers_internal::ReadPackedFixedError<sizeof(Value)>( repr.reader()); } + if (repr.reader().Pull(1, IntCast<size_t>(repr.length())) && + repr.reader().available() >= IntCast<size_t>(repr.length())) { + const absl::string_view value(repr.reader().cursor(), + IntCast<size_t>(repr.length())); + repr.reader().move_cursor(IntCast<size_t>(repr.length())); + absl::Status status = + HandleLengthDelimitedFromStringInternal(value, context...); + // Comparison against `absl::CancelledError()` is a fast path of + // `absl::IsCancelled()`. + if (ABSL_PREDICT_FALSE(!status.ok() && status != absl::CancelledError())) { + status = field_handlers_internal::AnnotateByReader(std::move(status), + repr.reader()); + } + return status; + } Position length = repr.length(); while (length > 0) { Value element; @@ -988,6 +1185,15 @@ if (ABSL_PREDICT_FALSE(repr.length() % sizeof(Value) > 0)) { return field_handlers_internal::ReadPackedFixedError<sizeof(Value)>(); } + if (const absl::string_view chunk = + absl::Cord::ChunkRemaining(repr.iterator()); + chunk.size() >= IntCast<size_t>(repr.length())) { + const absl::string_view value = + chunk.substr(0, IntCast<size_t>(repr.length())); + absl::Cord::AdvanceAndRead(&repr.iterator(), + IntCast<size_t>(repr.length())); + return HandleLengthDelimitedFromStringInternal(value, context...); + } Position length = repr.length(); while (length > 0) { char buffer[sizeof(Value)]; @@ -1012,17 +1218,7 @@ if (ABSL_PREDICT_FALSE(repr.size() % sizeof(Value) > 0)) { return field_handlers_internal::ReadPackedFixedError<sizeof(Value)>(); } - const char* cursor = repr.data(); - const char* const limit = repr.data() + repr.size(); - while (cursor < limit) { - const Value element = ReadLittleEndian<Value>(cursor); - if (absl::Status status = this->action()(element, context...); - ABSL_PREDICT_FALSE(!status.ok())) { - return status; - } - cursor += sizeof(Value); - } - return absl::OkStatus(); + return HandleLengthDelimitedFromStringInternal(repr, context...); } } // namespace field_handlers
diff --git a/riegeli/varint/BUILD b/riegeli/varint/BUILD index 2708279..8c5aedd 100644 --- a/riegeli/varint/BUILD +++ b/riegeli/varint/BUILD
@@ -18,7 +18,9 @@ "//riegeli/base:arithmetic", "//riegeli/base:assert", "//riegeli/bytes:reader", + "@com_google_absl//absl/base:config", "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/numeric:bits", "@com_google_absl//absl/strings:cord", "@com_google_absl//absl/strings:string_view", ],
diff --git a/riegeli/varint/varint_reading.cc b/riegeli/varint/varint_reading.cc index f30e515..4765cef 100644 --- a/riegeli/varint/varint_reading.cc +++ b/riegeli/varint/varint_reading.cc
@@ -20,14 +20,16 @@ #include <cstring> #include "absl/base/attributes.h" +#include "absl/base/config.h" #include "absl/base/optimization.h" +#include "absl/numeric/bits.h" #include "absl/strings/cord.h" #include "absl/strings/string_view.h" #include "riegeli/base/arithmetic.h" #include "riegeli/base/assert.h" #include "riegeli/bytes/reader.h" -namespace riegeli::varint_internal { +namespace riegeli { namespace { @@ -415,8 +417,16 @@ return index + 1; } +inline uint64_t ReadNativeEndian(const char* src) { + uint64_t dest; + std::memcpy(&dest, src, sizeof(dest)); + return dest; +} + } // namespace +namespace varint_internal { + template <typename T, bool canonical, size_t initial_index> bool ReadVarintFromReaderBuffer(Reader& src, const char* cursor, T acc, T& dest) { @@ -753,4 +763,67 @@ template size_t SkipVarintFromArray<uint64_t, true, 2>(const char* src, size_t available); -} // namespace riegeli::varint_internal +} // namespace varint_internal + +size_t CountVarints(absl::string_view value) { + // The number of varints is the number of bytes with the highest bit clear. + // This is easier to compute as the total number of bytes, minus the number + // of bytes with the highest bit set. + size_t num_varints = value.size(); + if (value.size() < sizeof(uint64_t)) { + // Count byte by byte. + for (const char byte : value) { + num_varints -= static_cast<uint8_t>(byte) >> 7; + } + return num_varints; + } + + // Count in whole blocks, except for the last one. + const char* const limit = value.data() + value.size() - sizeof(uint64_t); + const char* cursor = value.data(); + while (cursor < limit) { + const uint64_t block = ReadNativeEndian(cursor); + num_varints -= + IntCast<size_t>(absl::popcount(block & uint64_t{0x8080808080808080})); + cursor += 8; + } + + // Count in the last, possibly incomplete block. + const uint64_t block = ReadNativeEndian(limit); + uint64_t mask = uint64_t{0x8080808080808080}; +#if ABSL_IS_LITTLE_ENDIAN + mask <<= PtrDistance(limit, cursor) * 8; +#elif ABSL_IS_BIG_ENDIAN + mask >>= PtrDistance(limit, cursor) * 8; +#else +#error Unknown endianness +#endif + num_varints -= IntCast<size_t>(absl::popcount(block & mask)); + + return num_varints; +} + +bool VerifyBools(absl::string_view value) { + uint64_t bit_or = 0; + if (value.size() < sizeof(uint64_t)) { + // Verify byte by byte. + for (const char byte : value) { + bit_or |= static_cast<uint8_t>(byte); + } + return bit_or <= 1; + } + + // Verify whole blocks, except for the last one. + const char* const limit = value.data() + value.size() - sizeof(uint64_t); + const char* cursor = value.data(); + while (cursor < limit) { + bit_or |= ReadNativeEndian(cursor); + cursor += 8; + } + // Verify the last, possibly incomplete block. + bit_or |= ReadNativeEndian(limit); + + return (bit_or & ~uint64_t{0x0101010101010101}) == 0; +} + +} // namespace riegeli
diff --git a/riegeli/varint/varint_reading.h b/riegeli/varint/varint_reading.h index 2467401..790286c 100644 --- a/riegeli/varint/varint_reading.h +++ b/riegeli/varint/varint_reading.h
@@ -352,6 +352,17 @@ constexpr int32_t DecodeVarintSigned32(uint32_t repr); constexpr int64_t DecodeVarintSigned64(uint64_t repr); +// Counts the number of varints in `value`. +// +// If varints are valid only up to some point, then returns at least the number +// of valid varints. Returns `value.size()` only if each varint is valid and +// takes one byte. +size_t CountVarints(absl::string_view value); + +// Checks if each byte of `value` is a valid representation for a `bool`, +// i.e. 0 or 1. Optimized for the result being `true`. +bool VerifyBools(absl::string_view value); + // Implementation details follow. namespace varint_internal {