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
8 changes: 4 additions & 4 deletions rust/cpp_kernel/extension.cc
Original file line number Diff line number Diff line change
Expand Up @@ -86,15 +86,15 @@ google::protobuf::rust::PtrAndLen proto2_rust_Message_get_extension_string(
const google::protobuf::MessageLite* proto2_rust_Message_get_extension_message(
const google::protobuf::MessageLite* m, int32_t number,
const google::protobuf::MessageLite* default_instance) {
return &GetExtensionSet(m)->GetMessage(m->GetArena(), number,
*default_instance);
return &GetExtensionSet(m)->GetMessageByPrototype(m->GetArena(), number,
*default_instance);
}

google::protobuf::MessageLite* proto2_rust_Message_mutable_extension_message(
google::protobuf::MessageLite* m, int32_t number, int32_t type,
const google::protobuf::MessageLite* default_instance) {
return GetExtensionSet(m)->MutableMessage(m->GetArena(), number, type,
*default_instance, nullptr);
return GetExtensionSet(m)->MutableMessageByPrototype(
m->GetArena(), number, type, *default_instance, nullptr);
}

} // extern "C"
48 changes: 38 additions & 10 deletions src/google/protobuf/extension_set.cc
Original file line number Diff line number Diff line change
Expand Up @@ -35,10 +35,12 @@
#include "google/protobuf/internal_visibility.h"
#include "google/protobuf/io/coded_stream.h"
#include "google/protobuf/message_lite.h"
#include "google/protobuf/message_traits.h"
#include "google/protobuf/metadata_lite.h"
#include "google/protobuf/parse_context.h"
#include "google/protobuf/port.h"
#include "google/protobuf/repeated_field.h"
#include "google/protobuf/static_message_factory.h"
#include "google/protobuf/wire_format_lite.h"


Expand Down Expand Up @@ -446,36 +448,48 @@ std::string* ExtensionSet::AddString(Arena* arena, int number, FieldType type,
// -------------------------------------------------------------------
// Messages

const MessageLite& ExtensionSet::GetMessage(
Arena* arena, int number, const MessageLite& default_value) const {
template <typename Strategy>
const MessageLite& ExtensionSet::GetMessageGeneric(Strategy strategy,
Arena* arena,
int number) const {
const Extension* extension = FindOrNull(number);
if (extension == nullptr) {
// Not present. Return the default value.
return default_value;
return strategy.Default();
} else {
ABSL_DCHECK_TYPE(*extension, OPTIONAL_FIELD, MESSAGE);
ABSL_DCHECK(!extension->is_lazy);
return *extension->ptr.message_value;
}
}

const MessageLite& ExtensionSet::GetMessageByClassData(
Arena* arena, int number, const ClassData* class_data) const {
return GetMessageGeneric(ByClassData(class_data), arena, number);
}

const MessageLite& ExtensionSet::GetMessageByPrototype(
Arena* arena, int number, const MessageLite& default_value) const {
return GetMessageGeneric(ByPrototype(&default_value), arena, number);
}

// Defined in extension_set_heavy.cc.
// const MessageLite& ExtensionSet::GetMessage(Arena* arena, int number,
// const Descriptor* message_type,
// MessageFactory* factory) const

MessageLite* ExtensionSet::MutableMessage(Arena* arena, int number,
FieldType type,
const MessageLite& prototype,
const FieldDescriptor* descriptor) {
template <typename Strategy>
MessageLite* ExtensionSet::MutableMessageGeneric(
Strategy strategy, Arena* arena, int number, FieldType type,
const FieldDescriptor* descriptor) {
Extension* extension;
if (MaybeNewExtension(arena, number, descriptor, &extension)) {
extension->type = type;
ABSL_DCHECK_EQ(cpp_type(extension->type), WireFormatLite::CPPTYPE_MESSAGE);
extension->is_repeated = false;
extension->is_pointer = true;
extension->is_lazy = false;
extension->ptr.message_value = prototype.New(arena);
extension->ptr.message_value = strategy.New(arena);
extension->is_cleared = false;
return extension->ptr.message_value;
} else {
Expand All @@ -486,6 +500,20 @@ MessageLite* ExtensionSet::MutableMessage(Arena* arena, int number,
}
}

MessageLite* ExtensionSet::MutableMessageByPrototype(
Arena* arena, int number, FieldType type, const MessageLite& prototype,
const FieldDescriptor* descriptor) {
return MutableMessageGeneric(ByPrototype(&prototype), arena, number, type,
descriptor);
}

MessageLite* ExtensionSet::MutableMessageByClassData(
Arena* arena, int number, FieldType type, const ClassData* class_data,
const FieldDescriptor* descriptor) {
return MutableMessageGeneric(ByClassData(class_data), arena, number, type,
descriptor);
}

// Defined in extension_set_heavy.cc.
// MessageLite* ExtensionSet::MutableMessage(Arena* arena, int number,
// FieldType type,
Expand Down Expand Up @@ -570,7 +598,7 @@ void ExtensionSet::UnsafeArenaSetAllocatedMessage(
}

MessageLite* ExtensionSet::ReleaseMessage(Arena* arena, int number,
const MessageLite& prototype) {
const ClassData* class_data) {
Extension* extension = FindOrNull(number);
if (extension == nullptr) {
// Not present. Return nullptr.
Expand All @@ -596,7 +624,7 @@ MessageLite* ExtensionSet::ReleaseMessage(Arena* arena, int number,
}

MessageLite* ExtensionSet::UnsafeArenaReleaseMessage(
Arena* arena, int number, const MessageLite& prototype) {
Arena* arena, int number, const ClassData* class_data) {
Extension* extension = FindOrNull(number);
if (extension == nullptr) {
// Not present. Return nullptr.
Expand Down
44 changes: 29 additions & 15 deletions src/google/protobuf/extension_set.h
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,6 @@
#include <variant>
#include <vector>

#include "absl/log/absl_log.h"

#include "google/protobuf/stubs/common.h"
#include "absl/base/casts.h"
#include "absl/base/prefetch.h"
Expand Down Expand Up @@ -393,7 +391,13 @@ class PROTOBUF_EXPORT ExtensionSet {
}
}

PROTOBUF_FUTURE_ADD_EARLY_NODISCARD const MessageLite& GetMessage(
template <typename Strategy>
const MessageLite& GetMessageGeneric(Strategy strategy, Arena* arena,
int number) const;

[[nodiscard]] const MessageLite& GetMessageByClassData(
Arena* arena, int number, const ClassData* class_data) const;
[[nodiscard]] const MessageLite& GetMessageByPrototype(
Arena* arena, int number, const MessageLite& default_value) const;
PROTOBUF_FUTURE_ADD_EARLY_NODISCARD const MessageLite& GetMessage(
Arena* arena, int number, const Descriptor* message_type,
Expand All @@ -404,8 +408,16 @@ class PROTOBUF_EXPORT ExtensionSet {
// type.
#define desc const FieldDescriptor* descriptor // avoid line wrapping
std::string* MutableString(Arena* arena, int number, FieldType type, desc);
MessageLite* MutableMessage(Arena* arena, int number, FieldType type,
const MessageLite& prototype, desc);

template <typename Strategy>
MessageLite* MutableMessageGeneric(Strategy strategy, Arena* arena,
int number, FieldType type, desc);
MessageLite* MutableMessageByPrototype(Arena* arena, int number,
FieldType type,
const MessageLite& prototype, desc);
MessageLite* MutableMessageByClassData(Arena* arena, int number,
FieldType type,
const ClassData* class_data, desc);
MessageLite* MutableMessage(Arena* arena, const FieldDescriptor* descriptor,
MessageFactory* factory);
// Adds the given message to the ExtensionSet, taking ownership of the
Expand All @@ -418,9 +430,9 @@ class PROTOBUF_EXPORT ExtensionSet {
const FieldDescriptor* descriptor,
MessageLite* message);
[[nodiscard]] MessageLite* ReleaseMessage(Arena* arena, int number,
const MessageLite& prototype);
const ClassData* class_data);
MessageLite* UnsafeArenaReleaseMessage(Arena* arena, int number,
const MessageLite& prototype);
const ClassData* class_data);

[[nodiscard]] MessageLite* ReleaseMessage(Arena* arena,
const FieldDescriptor* descriptor,
Expand Down Expand Up @@ -1652,11 +1664,14 @@ class MessageTypeTraits {
typedef MessageTypeTraits<Type> Singular;
static constexpr bool kLifetimeBound = true;

static constexpr const internal::ClassData* class_data() {
return internal::MessageTraits<Type>::class_data();
}
PROTOBUF_FUTURE_ADD_EARLY_NODISCARD static inline ConstType Get(
Arena* arena, int number, const ExtensionSet& set,
ConstType default_value) {
ConstType /* default_value */) {
return static_cast<const Type&>(
set.GetMessage(arena, number, default_value));
set.GetMessageByClassData(arena, number, class_data()));
}
PROTOBUF_FUTURE_ADD_EARLY_NODISCARD static inline std::nullptr_t GetPtr(
int /* number */, const ExtensionSet& /* set */,
Expand All @@ -1666,8 +1681,8 @@ class MessageTypeTraits {
}
static inline MutableType Mutable(Arena* arena, int number,
FieldType field_type, ExtensionSet* set) {
return static_cast<Type*>(set->MutableMessage(
arena, number, field_type, Type::default_instance(), nullptr));
return static_cast<Type*>(set->MutableMessageByClassData(
arena, number, field_type, class_data(), nullptr));
}
static inline void SetAllocated(Arena* arena, int number,
FieldType field_type, MutableType message,
Expand All @@ -1684,14 +1699,13 @@ class MessageTypeTraits {
[[nodiscard]] static inline MutableType Release(Arena* arena, int number,
FieldType /* field_type */,
ExtensionSet* set) {
return static_cast<Type*>(
set->ReleaseMessage(arena, number, Type::default_instance()));
return static_cast<Type*>(set->ReleaseMessage(arena, number, class_data()));
}
static inline MutableType UnsafeArenaRelease(Arena* arena, int number,
FieldType /* field_type */,
ExtensionSet* set) {
return static_cast<Type*>(set->UnsafeArenaReleaseMessage(
arena, number, Type::default_instance()));
return static_cast<Type*>(
set->UnsafeArenaReleaseMessage(arena, number, class_data()));
}
};

Expand Down
2 changes: 2 additions & 0 deletions src/google/protobuf/extension_set_heavy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
#include "absl/strings/string_view.h"
#include "absl/types/span.h"
#include "google/protobuf/arena.h"
#include "google/protobuf/class_data.h"
#include "google/protobuf/descriptor.h"
#include "google/protobuf/descriptor.pb.h"
#include "google/protobuf/extension_set.h"
Expand All @@ -39,6 +40,7 @@
#include "google/protobuf/io/coded_stream.h"
#include "google/protobuf/message.h"
#include "google/protobuf/message_lite.h"
#include "google/protobuf/message_traits.h"
#include "google/protobuf/parse_context.h"
#include "google/protobuf/port.h"
#include "google/protobuf/repeated_field.h"
Expand Down
19 changes: 10 additions & 9 deletions src/google/protobuf/extension_set_inl.h
Original file line number Diff line number Diff line change
Expand Up @@ -168,9 +168,9 @@ const char* ExtensionSet::ParseFieldWithExtensionInfo(
info.is_repeated
? AddMessage(arena, number, WireFormatLite::TYPE_GROUP,
info.message_info.GetClassData(), info.descriptor)
: MutableMessage(arena, number, WireFormatLite::TYPE_GROUP,
*info.message_info.GetPrototype(),
info.descriptor);
: MutableMessageByClassData(
arena, number, WireFormatLite::TYPE_GROUP,
info.message_info.GetClassData(), info.descriptor);
uint32_t tag = (number << 3) + WireFormatLite::WIRETYPE_START_GROUP;
return ctx->ParseGroup(value, ptr, tag);
}
Expand All @@ -180,9 +180,9 @@ const char* ExtensionSet::ParseFieldWithExtensionInfo(
info.is_repeated
? AddMessage(arena, number, WireFormatLite::TYPE_MESSAGE,
info.message_info.GetClassData(), info.descriptor)
: MutableMessage(arena, number, WireFormatLite::TYPE_MESSAGE,
*info.message_info.GetPrototype(),
info.descriptor);
: MutableMessageByClassData(
arena, number, WireFormatLite::TYPE_MESSAGE,
info.message_info.GetClassData(), info.descriptor);
return ctx->ParseMessage(value, ptr);
}
}
Expand Down Expand Up @@ -225,9 +225,10 @@ const char* ExtensionSet::ParseMessageSetItemTmpl(
? AddMessage(arena, type_id, WireFormatLite::TYPE_MESSAGE,
extension.message_info.GetClassData(),
extension.descriptor)
: MutableMessage(arena, type_id, WireFormatLite::TYPE_MESSAGE,
*extension.message_info.GetPrototype(),
extension.descriptor);
: MutableMessageByClassData(
arena, type_id, WireFormatLite::TYPE_MESSAGE,
extension.message_info.GetClassData(),
extension.descriptor);

const char* p;
// We can't use regular parse from string as we have to track
Expand Down
25 changes: 5 additions & 20 deletions src/google/protobuf/generated_message_tctable_lite.cc
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
#include "google/protobuf/io/zero_copy_stream_impl_lite.h"
#include "google/protobuf/map.h"
#include "google/protobuf/message_lite.h"
#include "google/protobuf/message_traits.h"
#include "google/protobuf/micro_string.h"
#include "google/protobuf/parse_context.h"
#include "google/protobuf/port.h"
Expand Down Expand Up @@ -686,28 +687,12 @@ PROTOBUF_ALWAYS_INLINE const char* TcParser::SingularParseMessageAuxImpl(
SyncHasbits(msg, hasbits, table);
auto& field = RefAt<MessageLite*>(msg, data.offset());
const auto aux = *table->field_aux(data.aux_idx());
#ifndef PROTOBUF_MESSAGE_GLOBALS
const auto* inner_table =
aux_is_table ? aux.table_ptr() : aux.message_default()->GetTcParseTable();

auto [inner_table, class_data] =
GetTableAndClassDataFromAux<aux_is_table>(aux);
if (field == nullptr) {
field = NewMessage(inner_table->class_data, msg->GetArena());
}
#else
const TcParseTableBase* inner_table;
if constexpr (aux_is_table) {
inner_table = MessageGlobalsBase::ToParseTableBase(aux.message_globals());
if (field == nullptr) {
field =
NewMessage(MessageGlobalsBase::GetClassData(aux.message_globals()),
msg->GetArena());
}
} else {
inner_table = aux.message_default()->GetTcParseTable();
if (field == nullptr) {
field = NewMessage(inner_table->class_data, msg->GetArena());
}
field = NewMessage(class_data, msg->GetArena());
}
#endif // PROTOBUF_MESSAGE_GLOBALS
const auto inner_loop = [&](const char* ptr) {
return ParseLoop(field, ptr, ctx, inner_table);
};
Expand Down
2 changes: 2 additions & 0 deletions src/google/protobuf/message_lite.h
Original file line number Diff line number Diff line change
Expand Up @@ -206,6 +206,7 @@ class InternalMetadataOffset;
template <typename T, size_t kFieldOffset>
struct InternalMetadataOffsetHelper;
class LazyField;
class ReflectionVisit;
class RepeatedPtrFieldBase;
class TcParser;
struct TcParseTableBase;
Expand Down Expand Up @@ -965,6 +966,7 @@ class PROTOBUF_EXPORT MessageLite {
template <typename T, size_t kFieldOffset>
friend struct internal::InternalMetadataOffsetHelper;
friend class internal::LazyField;
friend internal::ReflectionVisit;
friend internal::RepeatedPtrFieldBase;
friend class internal::SwapFieldHelper;
friend class internal::TcParser;
Expand Down
1 change: 1 addition & 0 deletions src/google/protobuf/reflection_visit_field_info.h
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
#include "absl/strings/cord.h"
#include "absl/strings/string_view.h"
#include "google/protobuf/arenastring.h"
#include "google/protobuf/class_data.h"
#include "google/protobuf/descriptor.h"
#include "google/protobuf/descriptor.pb.h"
#include "google/protobuf/extension_set.h"
Expand Down
1 change: 1 addition & 0 deletions src/google/protobuf/reflection_visit_fields.h
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
#include "google/protobuf/generated_message_reflection.h"
#include "google/protobuf/has_bits.h"
#include "google/protobuf/message.h"
#include "google/protobuf/message_traits.h"
#include "google/protobuf/port.h"
#include "google/protobuf/reflection.h"
#include "google/protobuf/reflection_visit_field_info.h"
Expand Down
16 changes: 16 additions & 0 deletions src/google/protobuf/static_message_factory.h
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

#include "absl/log/absl_check.h"
#include "google/protobuf/arena.h"
#include "google/protobuf/class_data.h"
#include "google/protobuf/message_lite.h"

// Must be included last.
Expand All @@ -25,6 +26,21 @@ class ByPrototype {
const MessageLite* prototype_;
};

class ByClassData {
public:
explicit ByClassData(const internal::ClassData* class_data)
: class_data_(class_data) {}

MessageLite* New(Arena* arena) const { return class_data_->New(arena); }

const MessageLite& Default() const {
return *class_data_->default_instance();
}

private:
const internal::ClassData* class_data_;
};

template <typename MessageType>
class ByTemplate {
public:
Expand Down
Loading