diff --git a/rust/cpp_kernel/extension.cc b/rust/cpp_kernel/extension.cc index b9f6a4bb4ceb2..e929c008f0a99 100644 --- a/rust/cpp_kernel/extension.cc +++ b/rust/cpp_kernel/extension.cc @@ -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" diff --git a/src/google/protobuf/extension_set.cc b/src/google/protobuf/extension_set.cc index ad6c308b4f0e4..abdfe77b93bc7 100644 --- a/src/google/protobuf/extension_set.cc +++ b/src/google/protobuf/extension_set.cc @@ -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" @@ -446,12 +448,14 @@ 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 +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); @@ -459,15 +463,25 @@ const MessageLite& ExtensionSet::GetMessage( } } +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 +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; @@ -475,7 +489,7 @@ MessageLite* ExtensionSet::MutableMessage(Arena* arena, int number, 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 { @@ -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, @@ -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. @@ -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. diff --git a/src/google/protobuf/extension_set.h b/src/google/protobuf/extension_set.h index 402db21e33916..559561bfbd002 100644 --- a/src/google/protobuf/extension_set.h +++ b/src/google/protobuf/extension_set.h @@ -28,8 +28,6 @@ #include #include -#include "absl/log/absl_log.h" - #include "google/protobuf/stubs/common.h" #include "absl/base/casts.h" #include "absl/base/prefetch.h" @@ -393,7 +391,13 @@ class PROTOBUF_EXPORT ExtensionSet { } } - PROTOBUF_FUTURE_ADD_EARLY_NODISCARD const MessageLite& GetMessage( + template + 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, @@ -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 + 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 @@ -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, @@ -1652,11 +1664,14 @@ class MessageTypeTraits { typedef MessageTypeTraits Singular; static constexpr bool kLifetimeBound = true; + static constexpr const internal::ClassData* class_data() { + return internal::MessageTraits::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( - 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 */, @@ -1666,8 +1681,8 @@ class MessageTypeTraits { } static inline MutableType Mutable(Arena* arena, int number, FieldType field_type, ExtensionSet* set) { - return static_cast(set->MutableMessage( - arena, number, field_type, Type::default_instance(), nullptr)); + return static_cast(set->MutableMessageByClassData( + arena, number, field_type, class_data(), nullptr)); } static inline void SetAllocated(Arena* arena, int number, FieldType field_type, MutableType message, @@ -1684,14 +1699,13 @@ class MessageTypeTraits { [[nodiscard]] static inline MutableType Release(Arena* arena, int number, FieldType /* field_type */, ExtensionSet* set) { - return static_cast( - set->ReleaseMessage(arena, number, Type::default_instance())); + return static_cast(set->ReleaseMessage(arena, number, class_data())); } static inline MutableType UnsafeArenaRelease(Arena* arena, int number, FieldType /* field_type */, ExtensionSet* set) { - return static_cast(set->UnsafeArenaReleaseMessage( - arena, number, Type::default_instance())); + return static_cast( + set->UnsafeArenaReleaseMessage(arena, number, class_data())); } }; diff --git a/src/google/protobuf/extension_set_heavy.cc b/src/google/protobuf/extension_set_heavy.cc index b4936cdcd378b..1cf73dcd04b36 100644 --- a/src/google/protobuf/extension_set_heavy.cc +++ b/src/google/protobuf/extension_set_heavy.cc @@ -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" @@ -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" diff --git a/src/google/protobuf/extension_set_inl.h b/src/google/protobuf/extension_set_inl.h index bcb71a53c52cb..5afaea717a1a9 100644 --- a/src/google/protobuf/extension_set_inl.h +++ b/src/google/protobuf/extension_set_inl.h @@ -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); } @@ -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); } } @@ -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 diff --git a/src/google/protobuf/generated_message_tctable_lite.cc b/src/google/protobuf/generated_message_tctable_lite.cc index 55e7086c96c58..420ef9b703ac7 100644 --- a/src/google/protobuf/generated_message_tctable_lite.cc +++ b/src/google/protobuf/generated_message_tctable_lite.cc @@ -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" @@ -686,28 +687,12 @@ PROTOBUF_ALWAYS_INLINE const char* TcParser::SingularParseMessageAuxImpl( SyncHasbits(msg, hasbits, table); auto& field = RefAt(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); 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); }; diff --git a/src/google/protobuf/message_lite.h b/src/google/protobuf/message_lite.h index e085d95e611ac..a918a7f63c8a8 100644 --- a/src/google/protobuf/message_lite.h +++ b/src/google/protobuf/message_lite.h @@ -206,6 +206,7 @@ class InternalMetadataOffset; template struct InternalMetadataOffsetHelper; class LazyField; +class ReflectionVisit; class RepeatedPtrFieldBase; class TcParser; struct TcParseTableBase; @@ -965,6 +966,7 @@ class PROTOBUF_EXPORT MessageLite { template friend struct internal::InternalMetadataOffsetHelper; friend class internal::LazyField; + friend internal::ReflectionVisit; friend internal::RepeatedPtrFieldBase; friend class internal::SwapFieldHelper; friend class internal::TcParser; diff --git a/src/google/protobuf/reflection_visit_field_info.h b/src/google/protobuf/reflection_visit_field_info.h index 3362fc86eebf2..c0da4463c73c0 100644 --- a/src/google/protobuf/reflection_visit_field_info.h +++ b/src/google/protobuf/reflection_visit_field_info.h @@ -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" diff --git a/src/google/protobuf/reflection_visit_fields.h b/src/google/protobuf/reflection_visit_fields.h index 556d4d28df033..6eb2f5294a374 100644 --- a/src/google/protobuf/reflection_visit_fields.h +++ b/src/google/protobuf/reflection_visit_fields.h @@ -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" diff --git a/src/google/protobuf/static_message_factory.h b/src/google/protobuf/static_message_factory.h index ebc9e6277ac1d..09c2b6573bfdc 100644 --- a/src/google/protobuf/static_message_factory.h +++ b/src/google/protobuf/static_message_factory.h @@ -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. @@ -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 class ByTemplate { public: