Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 25 additions & 74 deletions common/values/bytes_value_input_stream.h
Original file line number Diff line number Diff line change
Expand Up @@ -21,15 +21,12 @@
#include <cstddef>
#include <cstdint>
#include <limits>
#include <new>
#include <utility>
#include <variant>

#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"
Expand All @@ -39,108 +36,62 @@ 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<ArrayStream, CordStream>;
using Variant = std::variant<std::monostate, ArrayStream, CordStream>;

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;
}
}

void Construct(absl::string_view value) {
ABSL_DCHECK_LE(value.size(),
static_cast<size_t>(std::numeric_limits<int>::max()));
::new (static_cast<void*>(&impl_[0]))
Variant(absl::in_place_type<ArrayStream>, value.data(),
static_cast<int>(value.size()));
}

void Construct(const absl::Cord& value) {
::new (static_cast<void*>(&impl_[0]))
Variant(absl::in_place_type<CordStream>, value);
}

void Destruct() { AsVariant().~variant(); }

Variant& AsVariant() ABSL_ATTRIBUTE_LIFETIME_BOUND {
return *std::launder(reinterpret_cast<Variant*>(&impl_[0]));
stream_ = &variant_.emplace<ArrayStream>(value.data(),
static_cast<int>(value.size()));
}

const Variant& AsVariant() const ABSL_ATTRIBUTE_LIFETIME_BOUND {
return *std::launder(reinterpret_cast<const Variant*>(&impl_[0]));
void Construct(absl::Cord value) {
stream_ = &variant_.emplace<CordStream>(std::move(value)).stream;
}

alignas(Variant) char impl_[sizeof(Variant)];
google::protobuf::io::ZeroCopyInputStream* stream_;
common_internal::SmallByteStringRep small_;
Variant variant_;
};

} // namespace cel
Expand Down
119 changes: 31 additions & 88 deletions common/values/bytes_value_output_stream.h
Original file line number Diff line number Diff line change
Expand Up @@ -19,18 +19,16 @@
#define THIRD_PARTY_CEL_CPP_COMMON_VALUES_BYTES_VALUE_OUTPUT_STREAM_H_

#include <cstdint>
#include <new>
#include <string>
#include <utility>
#include <variant>

#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"
Expand All @@ -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<std::monostate, StringStream, CordStream>;

struct Cord final {
explicit Cord(const absl::Cord& cord) : stream(cord) {}

google::protobuf::io::CordOutputStream stream;
};

using Variant = absl::variant<String, Cord>;
void Construct() { stream_ = &variant_.emplace<CordStream>(); }

void Construct(const BytesValue& value) {
switch (value.value_.GetKind()) {
Expand All @@ -150,26 +104,15 @@ class BytesValueOutputStream final : public google::protobuf::io::ZeroCopyOutput
}

void Construct(absl::string_view value) {
::new (static_cast<void*>(&impl_[0]))
Variant(absl::in_place_type<String>, value);
}

void Construct(const absl::Cord& value) {
::new (static_cast<void*>(&impl_[0]))
Variant(absl::in_place_type<Cord>, value);
}

void Destruct() { AsVariant().~variant(); }

Variant& AsVariant() ABSL_ATTRIBUTE_LIFETIME_BOUND {
return *std::launder(reinterpret_cast<Variant*>(&impl_[0]));
stream_ = &variant_.emplace<StringStream>(value).stream;
}

const Variant& AsVariant() const ABSL_ATTRIBUTE_LIFETIME_BOUND {
return *std::launder(reinterpret_cast<const Variant*>(&impl_[0]));
void Construct(absl::Cord value) {
stream_ = &variant_.emplace<CordStream>(std::move(value));
}

alignas(Variant) char impl_[sizeof(Variant)];
google::protobuf::io::ZeroCopyOutputStream* stream_;
Variant variant_;
};

} // namespace cel
Expand Down
25 changes: 21 additions & 4 deletions common/values/bytes_value_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down
Loading