diff --git a/common/values/bytes_value_input_stream.h b/common/values/bytes_value_input_stream.h index 35050d2de..fc2972a05 100644 --- a/common/values/bytes_value_input_stream.h +++ b/common/values/bytes_value_input_stream.h @@ -21,15 +21,12 @@ #include #include #include -#include +#include +#include -#include "absl/base/attributes.h" -#include "absl/base/nullability.h" #include "absl/log/absl_check.h" #include "absl/strings/cord.h" #include "absl/strings/string_view.h" -#include "absl/types/variant.h" -#include "absl/utility/utility.h" #include "common/internal/byte_string.h" #include "common/values/bytes_value.h" #include "google/protobuf/io/zero_copy_stream.h" @@ -39,80 +36,44 @@ namespace cel { class BytesValueInputStream final : public google::protobuf::io::ZeroCopyInputStream { public: - explicit BytesValueInputStream( - const BytesValue* absl_nonnull value ABSL_ATTRIBUTE_LIFETIME_BOUND) { - Construct(value); - } - - ~BytesValueInputStream() override { AsVariant().~variant(); } + explicit BytesValueInputStream(const BytesValue& value) { Construct(value); } bool Next(const void** data, int* size) override { - return absl::visit( - [&data, &size](auto& alternative) -> bool { - return alternative.stream.Next(data, size); - }, - AsVariant()); + return stream_->Next(data, size); } - void BackUp(int count) override { - absl::visit( - [&count](auto& alternative) -> void { - alternative.stream.BackUp(count); - }, - AsVariant()); - } + void BackUp(int count) override { stream_->BackUp(count); } - bool Skip(int count) override { - return absl::visit( - [&count](auto& alternative) -> bool { - return alternative.stream.Skip(count); - }, - AsVariant()); - } + bool Skip(int count) override { return stream_->Skip(count); } - int64_t ByteCount() const override { - return absl::visit( - [](const auto& alternative) -> int64_t { - return alternative.stream.ByteCount(); - }, - AsVariant()); - } + int64_t ByteCount() const override { return stream_->ByteCount(); } bool ReadCord(absl::Cord* cord, int count) override { - return absl::visit( - [&cord, &count](auto& alternative) -> bool { - return alternative.stream.ReadCord(cord, count); - }, - AsVariant()); + return stream_->ReadCord(cord, count); } private: - struct ArrayStream { - ArrayStream(const char* data, int size) : stream(data, size) {} - - google::protobuf::io::ArrayInputStream stream; - }; + using ArrayStream = google::protobuf::io::ArrayInputStream; struct CordStream { - explicit CordStream(const absl::Cord& cord) - : cord(cord), stream(&this->cord) {} + explicit CordStream(absl::Cord cord) + : cord(std::move(cord)), stream(&this->cord) {} absl::Cord cord; google::protobuf::io::CordInputStream stream; }; - using Variant = absl::variant; + using Variant = std::variant; - void Construct(const BytesValue* absl_nonnull value) { - ABSL_DCHECK(value != nullptr); - - switch (value->value_.GetKind()) { + void Construct(const BytesValue& value) { + switch (value.value_.GetKind()) { case common_internal::ByteStringKind::kSmall: - Construct(value->value_.GetSmall()); + small_ = value.value_.rep_.small; + Construct(absl::string_view(small_.data, small_.size)); break; case common_internal::ByteStringKind::kMedium: - Construct(value->value_.GetMedium()); + Construct(value.value_.GetMedium()); break; case common_internal::ByteStringKind::kLarge: - Construct(value->value_.GetLarge()); + Construct(value.value_.GetLarge()); break; } } @@ -120,27 +81,17 @@ class BytesValueInputStream final : public google::protobuf::io::ZeroCopyInputSt void Construct(absl::string_view value) { ABSL_DCHECK_LE(value.size(), static_cast(std::numeric_limits::max())); - ::new (static_cast(&impl_[0])) - Variant(absl::in_place_type, value.data(), - static_cast(value.size())); - } - - void Construct(const absl::Cord& value) { - ::new (static_cast(&impl_[0])) - Variant(absl::in_place_type, value); - } - - void Destruct() { AsVariant().~variant(); } - - Variant& AsVariant() ABSL_ATTRIBUTE_LIFETIME_BOUND { - return *std::launder(reinterpret_cast(&impl_[0])); + stream_ = &variant_.emplace(value.data(), + static_cast(value.size())); } - const Variant& AsVariant() const ABSL_ATTRIBUTE_LIFETIME_BOUND { - return *std::launder(reinterpret_cast(&impl_[0])); + void Construct(absl::Cord value) { + stream_ = &variant_.emplace(std::move(value)).stream; } - alignas(Variant) char impl_[sizeof(Variant)]; + google::protobuf::io::ZeroCopyInputStream* stream_; + common_internal::SmallByteStringRep small_; + Variant variant_; }; } // namespace cel diff --git a/common/values/bytes_value_output_stream.h b/common/values/bytes_value_output_stream.h index b23d1bd3d..5ec3b7ea7 100644 --- a/common/values/bytes_value_output_stream.h +++ b/common/values/bytes_value_output_stream.h @@ -19,18 +19,16 @@ #define THIRD_PARTY_CEL_CPP_COMMON_VALUES_BYTES_VALUE_OUTPUT_STREAM_H_ #include -#include #include #include +#include -#include "absl/base/attributes.h" #include "absl/base/nullability.h" +#include "absl/base/optimization.h" #include "absl/functional/overload.h" #include "absl/log/absl_check.h" #include "absl/strings/cord.h" #include "absl/strings/string_view.h" -#include "absl/types/variant.h" -#include "absl/utility/utility.h" #include "common/internal/byte_string.h" #include "common/values/bytes_value.h" #include "google/protobuf/arena.h" @@ -41,99 +39,55 @@ namespace cel { class BytesValueOutputStream final : public google::protobuf::io::ZeroCopyOutputStream { public: + BytesValueOutputStream() { Construct(); } + explicit BytesValueOutputStream(const BytesValue& value) { Construct(value); } bool Next(void** data, int* size) override { - return absl::visit(absl::Overload( - [&data, &size](String& string) -> bool { - return string.stream.Next(data, size); - }, - [&data, &size](Cord& cord) -> bool { - return cord.stream.Next(data, size); - }), - AsVariant()); + return stream_->Next(data, size); } - void BackUp(int count) override { - absl::visit( - absl::Overload( - [&count](String& string) -> void { string.stream.BackUp(count); }, - [&count](Cord& cord) -> void { cord.stream.BackUp(count); }), - AsVariant()); - } + void BackUp(int count) override { stream_->BackUp(count); } - int64_t ByteCount() const override { - return absl::visit(absl::Overload( - [](const String& string) -> int64_t { - return string.stream.ByteCount(); - }, - [](const Cord& cord) -> int64_t { - return cord.stream.ByteCount(); - }), - AsVariant()); - } + int64_t ByteCount() const override { return stream_->ByteCount(); } bool WriteAliasedRaw(const void* data, int size) override { - return absl::visit(absl::Overload( - [&data, &size](String& string) -> bool { - return string.stream.WriteAliasedRaw(data, size); - }, - [&data, &size](Cord& cord) -> bool { - return cord.stream.WriteAliasedRaw(data, size); - }), - AsVariant()); + return stream_->WriteAliasedRaw(data, size); } - bool AllowsAliasing() const override { - return absl::visit(absl::Overload( - [](const String& string) -> bool { - return string.stream.AllowsAliasing(); - }, - [](const Cord& cord) -> bool { - return cord.stream.AllowsAliasing(); - }), - AsVariant()); - } + bool AllowsAliasing() const override { return stream_->AllowsAliasing(); } bool WriteCord(const absl::Cord& out) override { - return absl::visit( - absl::Overload( - [&out](String& string) -> bool { - return string.stream.WriteCord(out); - }, - [&out](Cord& cord) -> bool { return cord.stream.WriteCord(out); }), - AsVariant()); + return stream_->WriteCord(out); } + [[nodiscard]] BytesValue Consume(google::protobuf::Arena* absl_nonnull arena) && { ABSL_DCHECK(arena != nullptr); - return absl::visit( - absl::Overload( - [arena](String& string) -> BytesValue { - return BytesValue::From(std::move(string.target), arena); - }, - [arena](Cord& cord) -> BytesValue { - return BytesValue::From(cord.stream.Consume(), arena); - }), - AsVariant()); + return std::visit( + absl::Overload([](std::monostate) -> BytesValue { ABSL_UNREACHABLE(); }, + [arena](StringStream& stream) -> BytesValue { + return BytesValue::From(std::move(stream.target), + arena); + }, + [arena](CordStream& stream) -> BytesValue { + return BytesValue::From(stream.Consume(), arena); + }), + variant_); } private: - struct String final { - explicit String(absl::string_view target) + struct StringStream final { + explicit StringStream(absl::string_view target) : target(target), stream(&this->target) {} std::string target; google::protobuf::io::StringOutputStream stream; }; + using CordStream = google::protobuf::io::CordOutputStream; + using Variant = std::variant; - struct Cord final { - explicit Cord(const absl::Cord& cord) : stream(cord) {} - - google::protobuf::io::CordOutputStream stream; - }; - - using Variant = absl::variant; + void Construct() { stream_ = &variant_.emplace(); } void Construct(const BytesValue& value) { switch (value.value_.GetKind()) { @@ -150,26 +104,15 @@ class BytesValueOutputStream final : public google::protobuf::io::ZeroCopyOutput } void Construct(absl::string_view value) { - ::new (static_cast(&impl_[0])) - Variant(absl::in_place_type, value); - } - - void Construct(const absl::Cord& value) { - ::new (static_cast(&impl_[0])) - Variant(absl::in_place_type, value); - } - - void Destruct() { AsVariant().~variant(); } - - Variant& AsVariant() ABSL_ATTRIBUTE_LIFETIME_BOUND { - return *std::launder(reinterpret_cast(&impl_[0])); + stream_ = &variant_.emplace(value).stream; } - const Variant& AsVariant() const ABSL_ATTRIBUTE_LIFETIME_BOUND { - return *std::launder(reinterpret_cast(&impl_[0])); + void Construct(absl::Cord value) { + stream_ = &variant_.emplace(std::move(value)); } - alignas(Variant) char impl_[sizeof(Variant)]; + google::protobuf::io::ZeroCopyOutputStream* stream_; + Variant variant_; }; } // namespace cel diff --git a/common/values/bytes_value_test.cc b/common/values/bytes_value_test.cc index e4e7ad665..d2d62b5af 100644 --- a/common/values/bytes_value_test.cc +++ b/common/values/bytes_value_test.cc @@ -168,9 +168,25 @@ TEST_F(BytesValueTest, Comparison) { EXPECT_FALSE(BytesValue::WrapUnsafe("foo") < BytesValue::WrapUnsafe("bar")); } -TEST_F(BytesValueTest, StringInputStream) { +TEST_F(BytesValueTest, SmallStringInputStream) { + BytesValue value = BytesValue::From("foo", arena()); + BytesValueInputStream stream(value); + const void* data; + int size; + absl::Cord cord; + ASSERT_TRUE(stream.Next(&data, &size)); + EXPECT_THAT(data, NotNull()); + EXPECT_EQ(size, 3); + EXPECT_EQ(stream.ByteCount(), 3); + stream.BackUp(size); + ASSERT_TRUE(stream.Skip(3)); + EXPECT_FALSE(stream.ReadCord(&cord, 3)); + EXPECT_FALSE(stream.Next(&data, &size)); +} + +TEST_F(BytesValueTest, MediumStringInputStream) { BytesValue value = BytesValue::WrapUnsafe("foo"); - BytesValueInputStream stream(&value); + BytesValueInputStream stream(value); const void* data; int size; absl::Cord cord; @@ -185,8 +201,9 @@ TEST_F(BytesValueTest, StringInputStream) { } TEST_F(BytesValueTest, CordInputStream) { - BytesValue value = BytesValue::From(absl::Cord("foo"), arena()); - BytesValueInputStream stream(&value); + absl::Cord value_cord("foo"); + BytesValue value = BytesValue::WrapUnsafe(&value_cord); + BytesValueInputStream stream(value); const void* data; int size; absl::Cord cord;