Internal change. PiperOrigin-RevId: 962102266
diff --git a/src/google/protobuf/compiler/java/full/enum_field.cc b/src/google/protobuf/compiler/java/full/enum_field.cc index fd7f982..6e3dec0 100644 --- a/src/google/protobuf/compiler/java/full/enum_field.cc +++ b/src/google/protobuf/compiler/java/full/enum_field.cc
@@ -501,14 +501,23 @@ WriteFieldEnumValueAccessorDocComment(printer, descriptor_, SETTER, context_->options(), /* builder */ true); - printer->Print(variables_, - "$deprecation$public Builder " - "${$set$capitalized_name$Value$}$(int value) {\n" - " $set_oneof_case_message$;\n" - " $oneof_name$_ = value;\n" - " onChanged();\n" - " return this;\n" - "}\n"); + printer->Print( + variables_, + "$deprecation$public Builder " + "${$set$capitalized_name$Value$}$(int value) {\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" + " $oneof_name$_ = value;\n" + " onChanged();\n" + " return this;\n" + "}\n"); printer->Annotate("{", "}", descriptor_, Semantic::kSet); } } @@ -539,7 +548,15 @@ "$deprecation$public Builder " "${$set$capitalized_name$$}$($type$ value) {\n" " $null_check$\n" - " $set_oneof_case_message$;\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" " $oneof_name$_ = value.getNumber();\n" " onChanged();\n" " return this;\n" @@ -602,19 +619,36 @@ if (SupportUnknownEnumValue(descriptor_)) { printer->Print(variables_, "int rawValue = input.readEnum();\n" - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "$oneof_name$_ = rawValue;\n"); } else { - printer->Print(variables_, - "int rawValue = input.readEnum();\n" - "$type$ value =\n" - " $type$.forNumber(rawValue);\n" - "if (value == null) {\n" - " mergeUnknownVarintField($number$, rawValue);\n" - "} else {\n" - " $set_oneof_case_message$;\n" - " $oneof_name$_ = rawValue;\n" - "}\n"); + printer->Print( + variables_, + "int rawValue = input.readEnum();\n" + "$type$ value =\n" + " $type$.forNumber(rawValue);\n" + "if (value == null) {\n" + " mergeUnknownVarintField($number$, rawValue);\n" + "} else {\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" + " $oneof_name$_ = rawValue;\n" + "}\n"); } }
diff --git a/src/google/protobuf/compiler/java/full/field_generator.cc b/src/google/protobuf/compiler/java/full/field_generator.cc index a72b7be..bf830ee 100644 --- a/src/google/protobuf/compiler/java/full/field_generator.cc +++ b/src/google/protobuf/compiler/java/full/field_generator.cc
@@ -1,6 +1,7 @@ #include "google/protobuf/compiler/java/full/field_generator.h" #include "google/protobuf/compiler/java/context.h" +#include "google/protobuf/compiler/java/helpers.h" #include "google/protobuf/compiler/java/name_resolver.h" namespace google { @@ -15,6 +16,14 @@ context_(context), name_resolver_(context->GetNameResolver()) {} +bool ImmutableFieldGenerator::HasHasbit() const { + return ::google::protobuf::compiler::java::HasHasbit(descriptor_); +} + +bool ImmutableFieldGenerator::IsRealOneof() const { + return ::google::protobuf::compiler::java::IsRealOneof(descriptor_); +} + } // namespace java } // namespace compiler } // namespace protobuf
diff --git a/src/google/protobuf/compiler/java/full/field_generator.h b/src/google/protobuf/compiler/java/full/field_generator.h index fddc8d5..193acb8 100644 --- a/src/google/protobuf/compiler/java/full/field_generator.h +++ b/src/google/protobuf/compiler/java/full/field_generator.h
@@ -23,6 +23,9 @@ ImmutableFieldGenerator& operator=(const ImmutableFieldGenerator&) = delete; ~ImmutableFieldGenerator() override = default; + bool HasHasbit() const; + bool IsRealOneof() const; + int GetBitIndex() const { return bit_index_; } constexpr int GetNumBits() const { return 1; } virtual void GenerateInterfaceMembers(io::Printer* printer) const = 0;
diff --git a/src/google/protobuf/compiler/java/full/make_field_gens.cc b/src/google/protobuf/compiler/java/full/make_field_gens.cc index e88deb6..0c11451 100644 --- a/src/google/protobuf/compiler/java/full/make_field_gens.cc +++ b/src/google/protobuf/compiler/java/full/make_field_gens.cc
@@ -86,7 +86,13 @@ } } } - +bool HasExplicitPresence(const FieldDescriptor* field) { + return HasHasbit(field); +} +bool HasNoPresence(const FieldDescriptor* field) { return IsRealOneof(field); } +bool HasHintBitFields(const FieldDescriptor* field) { + return !HasExplicitPresence(field) && !HasNoPresence(field); +} } // namespace FieldGeneratorMap<ImmutableFieldGenerator> MakeImmutableFieldGenerators( @@ -95,11 +101,35 @@ // bit fields. int bit_index = 0; FieldGeneratorMap<ImmutableFieldGenerator> ret(descriptor); + + // First pass: fields with real presence bits. for (int i = 0; i < descriptor->field_count(); i++) { const FieldDescriptor* field = descriptor->field(i); - auto generator = MakeImmutableGenerator(field, bit_index, context); - bit_index += generator->GetNumBits(); - ret.Add(field, std::move(generator)); + if (HasExplicitPresence(field)) { + auto generator = MakeImmutableGenerator(field, bit_index, context); + bit_index += generator->GetNumBits(); + ret.Add(field, std::move(generator)); + } + } + + // Second pass: fields with hint presence bits. + for (int i = 0; i < descriptor->field_count(); i++) { + const FieldDescriptor* field = descriptor->field(i); + if (HasHintBitFields(field)) { + auto generator = MakeImmutableGenerator(field, bit_index, context); + bit_index += generator->GetNumBits(); + ret.Add(field, std::move(generator)); + } + } + + // Third pass: fields with no presence tracking. + for (int i = 0; i < descriptor->field_count(); i++) { + const FieldDescriptor* field = descriptor->field(i); + if (HasNoPresence(field)) { + auto generator = MakeImmutableGenerator(field, bit_index, context); + bit_index += generator->GetNumBits(); + ret.Add(field, std::move(generator)); + } } return ret; }
diff --git a/src/google/protobuf/compiler/java/full/message.cc b/src/google/protobuf/compiler/java/full/message.cc index 82ce443..c493c4a 100644 --- a/src/google/protobuf/compiler/java/full/message.cc +++ b/src/google/protobuf/compiler/java/full/message.cc
@@ -116,7 +116,8 @@ const OneofDescriptor* oneof = descriptor_->field(i)->containing_oneof(); auto& generator = oneof_generators_[oneof->index()]; if (generator == nullptr) { - generator = std::make_unique<OneofGenerator>(oneof, context_); + generator = std::make_unique<OneofGenerator>(oneof, context_, + field_generators_); } } }
diff --git a/src/google/protobuf/compiler/java/full/message_builder.cc b/src/google/protobuf/compiler/java/full/message_builder.cc index 6781472..68252c7 100644 --- a/src/google/protobuf/compiler/java/full/message_builder.cc +++ b/src/google/protobuf/compiler/java/full/message_builder.cc
@@ -18,10 +18,7 @@ #include <vector> #include "absl/container/btree_map.h" -#include "absl/container/btree_set.h" -#include "absl/container/flat_hash_map.h" #include "absl/log/absl_check.h" -#include "absl/strings/ascii.h" #include "absl/strings/str_cat.h" #include "absl/strings/str_replace.h" #include "absl/strings/string_view.h" @@ -30,7 +27,7 @@ #include "google/protobuf/compiler/java/context.h" #include "google/protobuf/compiler/java/doc_comment.h" #include "google/protobuf/compiler/java/field_common.h" -#include "google/protobuf/compiler/java/generator_factory.h" +#include "google/protobuf/compiler/java/generator_common.h" #include "google/protobuf/compiler/java/helpers.h" #include "google/protobuf/compiler/java/full/enum.h" #include "google/protobuf/compiler/java/full/extension.h" @@ -128,7 +125,7 @@ // oneof for (const auto& kv : oneof_generators_) { - kv.second->GenerateCommonBuilderMethods(printer); + kv.second->GenerateCommonBuilderMethods(printer, field_generators_); } // Integers for bit fields. @@ -571,10 +568,6 @@ } } - if (!oneof_generators_.empty()) { - printer->Print("buildPartialOneofs(result);\n"); - } - printer->Outdent(); printer->Print( " onBuilt();\n" @@ -584,69 +577,44 @@ "classname", name_resolver_->GetImmutableClassName(descriptor_)); // Build all fields in shards organized by bitfield membership. - int start_field = 0; for (int i = 0; i < totalInts; i++) { - start_field = GenerateBuildPartialShard(printer, i, start_field); - } - - // Build Oneofs - if (!oneof_generators_.empty()) { - printer->Print("private void buildPartialOneofs($classname$ result) {\n", - "classname", - name_resolver_->GetImmutableClassName(descriptor_)); - printer->Indent(); - for (const auto& kv : oneof_generators_) { - kv.second->GenerateBuildingCode(printer, field_generators_); - } - printer->Outdent(); - printer->Print("}\n\n"); + GenerateBuildPartialShard(printer, i); } } -int MessageBuilderGenerator::GenerateBuildPartialShard(io::Printer* printer, - int shard, - int first_field) { +void MessageBuilderGenerator::GenerateBuildPartialShard(io::Printer* printer, + int shard) { printer->Print( "private void buildPartial_autosplit_$shard$($classname$ result) {\n" - " int from_$bit_field_name$ = $bit_field_name$;\n", + " int from_$bit_field_name$ = $bit_field_name$;\n" + " int to_$bit_field_name$ = 0;\n", "classname", name_resolver_->GetImmutableClassName(descriptor_), "shard", absl::StrCat(shard), "bit_field_name", GetBitFieldName(shard)); printer->Indent(); - absl::btree_set<int> declared_to_bitfields; - int bit = 0; - int next = first_field; - for (; bit < 32 && next < descriptor_->field_count(); ++next, ++bit) { + int i = shard * 32; + int shard_end = std::min(i + 32, static_cast<int>(field_generators_.size())); + for (; i < shard_end; ++i) { const ImmutableFieldGenerator& field = - field_generators_.get(descriptor_->field(next)); - - // Skip oneof fields that are handled separately - if (IsRealOneof(descriptor_->field(next))) { - continue; - } - - // Track message bits if necessary - int to_bitfield = field.GetBitIndex() / 32; - if (declared_to_bitfields.count(to_bitfield) == 0) { - printer->Print("int to_$bit_field_name$ = 0;\n", "bit_field_name", - GetBitFieldName(to_bitfield)); - declared_to_bitfields.insert(to_bitfield); - } + field_generators_.getInInsertOrder(i); // Copy the field from the builder to the message field.GenerateBuildingCode(printer); } - // Copy the bit field results to the generated message - for (int to_bitfield : declared_to_bitfields) { - printer->Print("result.$bit_field_name$ |= to_$bit_field_name$;\n", - "bit_field_name", GetBitFieldName(to_bitfield)); + for (const auto& kv : oneof_generators_) { + if (kv.second->HasScalarFields(shard)) { + kv.second->GenerateBuildingCode(printer, shard); + } } printer->Outdent(); - printer->Print("}\n\n"); - return next; + // Copy the bit field results to the generated message + printer->Print( + " result.$bit_field_name$ |= to_$bit_field_name$;\n" + "}\n\n", + "bit_field_name", GetBitFieldName(shard)); } // ===================================================================
diff --git a/src/google/protobuf/compiler/java/full/message_builder.h b/src/google/protobuf/compiler/java/full/message_builder.h index e793fe3..b338f90 100644 --- a/src/google/protobuf/compiler/java/full/message_builder.h +++ b/src/google/protobuf/compiler/java/full/message_builder.h
@@ -14,11 +14,11 @@ #include <memory> #include <string> -#include <vector> #include "absl/container/btree_map.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" +#include "google/protobuf/compiler/java/generator_common.h" #include "google/protobuf/compiler/java/full/field_generator.h" #include "google/protobuf/compiler/java/full/oneof_generator.h" #include "google/protobuf/descriptor.h" @@ -67,8 +67,7 @@ io::Printer* printer, absl::Span<const std::string> merging_code_blocks, absl::string_view method_suffix); void GenerateBuildPartial(io::Printer* printer); - int GenerateBuildPartialShard(io::Printer* printer, int shard, - int first_field); + void GenerateBuildPartialShard(io::Printer* printer, int shard); void GenerateDescriptorMethods(io::Printer* printer); void GenerateBuilderParsingMethods(io::Printer* printer); void GenerateBuilderFieldParsingCases(io::Printer* printer);
diff --git a/src/google/protobuf/compiler/java/full/message_field.cc b/src/google/protobuf/compiler/java/full/message_field.cc index 26cb67a..1b62132 100644 --- a/src/google/protobuf/compiler/java/full/message_field.cc +++ b/src/google/protobuf/compiler/java/full/message_field.cc
@@ -494,11 +494,9 @@ "if ($get_has_field_bit_from_local$) {\n" " result.$name$_ = $name$Builder_ == null\n" " ? $name$_\n" - " : $name$Builder_.build();\n"); - if (GetNumBits() > 0) { - printer->Print(variables_, " $set_has_field_bit_to_local$;\n"); - } - printer->Print("}\n"); + " : $name$Builder_.build();\n" + " $set_has_field_bit_to_local$;\n" + "}\n"); } void ImmutableMessageFieldGenerator::GenerateBuilderParsingCode( @@ -679,7 +677,15 @@ "$name$Builder_.setMessage(value);\n", - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "return this;\n", Semantic::kSet); } @@ -720,7 +726,15 @@ " $name$Builder_.setMessage(value);\n" "}\n", - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "return this;\n", Semantic::kSet); } @@ -800,7 +814,15 @@ " isClean());\n" " $oneof_name$_ = null;\n" " }\n" - " $set_oneof_case_message$;\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" " $on_changed$\n" " return $name$Builder_;\n" "}\n"); @@ -842,9 +864,12 @@ void ImmutableMessageOneofFieldGenerator::GenerateBuildingCode( io::Printer* printer) const { printer->Print(variables_, - "if ($has_oneof_case_message$ &&\n" - " $name$Builder_ != null) {\n" - " result.$oneof_name$_ = $name$Builder_.build();\n" + "if ($get_has_field_bit_from_local$) {\n" + " result.$oneof_name$_ = $name$Builder_ == null\n" + " ? $oneof_name$_\n" + " : $name$Builder_.build();\n" + " result.$oneof_name$Case_ = $number$;\n" + " $set_has_field_bit_to_local$;\n" "}\n"); } @@ -862,14 +887,30 @@ " " "internalGet$capitalized_name$FieldBuilder().getBuilder(),\n" " extensionRegistry);\n" - "$set_oneof_case_message$;\n"); + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n"); } else { printer->Print(variables_, "input.readMessage(\n" " " "internalGet$capitalized_name$FieldBuilder().getBuilder(),\n" " extensionRegistry);\n" - "$set_oneof_case_message$;\n"); + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n"); } }
diff --git a/src/google/protobuf/compiler/java/full/message_field.h b/src/google/protobuf/compiler/java/full/message_field.h index c14b32e..1db7d64 100644 --- a/src/google/protobuf/compiler/java/full/message_field.h +++ b/src/google/protobuf/compiler/java/full/message_field.h
@@ -111,8 +111,9 @@ void GenerateMembers(io::Printer* printer) const override; void GenerateBuilderMembers(io::Printer* printer) const override; void GenerateBuilderClearCode(io::Printer* printer) const override; - void GenerateBuildingCode(io::Printer* printer) const override; + void GenerateMergingCode(io::Printer* printer) const override; + void GenerateBuildingCode(io::Printer* printer) const override; void GenerateBuilderParsingCode(io::Printer* printer) const override; void GenerateSerializationCode(io::Printer* printer) const override; void GenerateSerializedSizeCode(io::Printer* printer) const override;
diff --git a/src/google/protobuf/compiler/java/full/oneof_generator.cc b/src/google/protobuf/compiler/java/full/oneof_generator.cc index 55c6f3c..ccc3a57 100644 --- a/src/google/protobuf/compiler/java/full/oneof_generator.cc +++ b/src/google/protobuf/compiler/java/full/oneof_generator.cc
@@ -53,21 +53,36 @@ } // namespace -OneofGenerator::OneofGenerator(const OneofDescriptor* descriptor, - Context* context) +OneofGenerator::OneofGenerator( + const OneofDescriptor* descriptor, Context* context, + const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) : descriptor_(descriptor) { SetOneofVariables(descriptor, context, &variables_); + for (int i = 0; i < descriptor_->field_count(); i++) { + const FieldDescriptor* field = descriptor_->field(i); + const ImmutableFieldGenerator& generator = field_generators.get(field); + int bit_index = generator.GetBitIndex(); + if (bit_index >= 0) { + clear_masks_[bit_index / 32] |= (1u << (bit_index % 32)); + if (field->message_type() == nullptr) { + scalar_masks_[bit_index / 32] |= (1u << (bit_index % 32)); + } + } + } } OneofGenerator::~OneofGenerator() = default; -void OneofGenerator::GenerateCommonBuilderMethods(io::Printer* printer) const { +void OneofGenerator::GenerateCommonBuilderMethods( + io::Printer* printer, + const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const { // oneofCase_ and oneof_ printer->Print(variables_, "private int $oneof_name$Case_ = 0;\n" "private java.lang.Object $oneof_name$_;\n"); GenerateBuilderGetOneofCase(printer); GenerateBuilderClearOneof(printer); + GenerateBuilderClearOneofHasBits(printer); } void OneofGenerator::GenerateBuilderGetOneofCase(io::Printer* printer) const { @@ -83,7 +98,13 @@ void OneofGenerator::GenerateBuilderClearOneof(io::Printer* printer) const { printer->Print(variables_, "\n" - "public Builder ${$clear$oneof_capitalized_name$$}$() {\n" + "public Builder ${$clear$oneof_capitalized_name$$}$() {\n"); + for (const auto& kv : clear_masks_) { + printer->Print(" $bit_field_name$ = ($bit_field_name$ & ~$mask$);\n", + "bit_field_name", GetBitFieldName(kv.first), "mask", + absl::StrCat("0x", absl::Hex(kv.second, absl::kZeroPad8))); + } + printer->Print(variables_, " $oneof_name$Case_ = 0;\n" " $oneof_name$_ = null;\n" " onChanged();\n" @@ -98,6 +119,11 @@ printer->Print(variables_, "$oneof_name$Case_ = 0;\n" "$oneof_name$_ = null;\n"); + for (const auto& kv : clear_masks_) { + printer->Print("$bit_field_name$ = ($bit_field_name$ & ~$mask$);\n", + "bit_field_name", GetBitFieldName(kv.first), "mask", + absl::StrCat("0x", absl::Hex(kv.second, absl::kZeroPad8))); + } } void OneofGenerator::GenerateMergingCode( @@ -125,19 +151,37 @@ "}\n"); } -void OneofGenerator::GenerateBuildingCode( - io::Printer* printer, - const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const { - printer->Print(variables_, - "result.$oneof_name$Case_ = $oneof_name$Case_;\n" - "result.$oneof_name$_ = this.$oneof_name$_;\n"); - for (int i = 0; i < descriptor_->field_count(); ++i) { - if (descriptor_->field(i)->message_type() != nullptr) { - const ImmutableFieldGenerator& field = - field_generators.get(descriptor_->field(i)); - field.GenerateBuildingCode(printer); - } +void OneofGenerator::GenerateBuilderClearOneofHasBits( + io::Printer* printer) const { + printer->Print( + variables_, + "\n" + "private void ${$clear$oneof_capitalized_name$HasBits$}$() {\n"); + for (const auto& kv : clear_masks_) { + printer->Print(" $bit_field_name$ = ($bit_field_name$ & ~$mask$);\n", + "bit_field_name", GetBitFieldName(kv.first), "mask", + absl::StrCat("0x", absl::Hex(kv.second, absl::kZeroPad8))); } + printer->Print(variables_, "}\n\n"); +} + +bool OneofGenerator::HasScalarFields(int shard) const { + return scalar_masks_.find(shard) != scalar_masks_.end(); +} + +void OneofGenerator::GenerateBuildingCode(io::Printer* printer, + int shard) const { + auto it = scalar_masks_.find(shard); + if (it == scalar_masks_.end()) return; + auto vars = variables_; + vars["bit_field_name"] = GetBitFieldName(shard); + vars["mask"] = absl::StrCat("0x", absl::Hex(it->second, absl::kZeroPad8)); + printer->Print(vars, + "to_$bit_field_name$ |= (from_$bit_field_name$ & $mask$);\n" + "if ((from_$bit_field_name$ & $mask$) != 0) {\n" + " result.$oneof_name$Case_ = $oneof_name$Case_;\n" + " result.$oneof_name$_ = $oneof_name$_;\n" + "}\n"); } void OneofGenerator::GenerateInterfaceMembers(io::Printer* printer) const {
diff --git a/src/google/protobuf/compiler/java/full/oneof_generator.h b/src/google/protobuf/compiler/java/full/oneof_generator.h index 27511cb..445f12b 100644 --- a/src/google/protobuf/compiler/java/full/oneof_generator.h +++ b/src/google/protobuf/compiler/java/full/oneof_generator.h
@@ -12,6 +12,7 @@ #ifndef GOOGLE_PROTOBUF_COMPILER_JAVA_IMMUTABLE_ONEOF_GENERATOR_H__ #define GOOGLE_PROTOBUF_COMPILER_JAVA_IMMUTABLE_ONEOF_GENERATOR_H__ +#include <map> #include <string> #include "absl/container/flat_hash_map.h" @@ -39,7 +40,9 @@ class OneofGenerator { public: - OneofGenerator(const OneofDescriptor* descriptor, Context* context); + OneofGenerator( + const OneofDescriptor* descriptor, Context* context, + const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators); OneofGenerator(const OneofGenerator&) = delete; OneofGenerator& operator=(const OneofGenerator&) = delete; ~OneofGenerator(); @@ -53,22 +56,29 @@ io::Printer* printer, const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const; - void GenerateCommonBuilderMethods(io::Printer* printer) const; + void GenerateCommonBuilderMethods( + io::Printer* printer, + const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const; void GenerateBuilderClearMethod(io::Printer* printer) const; void GenerateMergingCode( io::Printer* printer, const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const; - void GenerateBuildingCode( - io::Printer* printer, - const FieldGeneratorMap<ImmutableFieldGenerator>& field_generators) const; + bool HasScalarFields(int shard) const; + void GenerateBuildingCode(io::Printer* printer, int shard) const; private: void GenerateBuilderGetOneofCase(io::Printer* printer) const; void GenerateBuilderClearOneof(io::Printer* printer) const; + void GenerateBuilderClearOneofHasBits(io::Printer* printer) const; const OneofDescriptor* descriptor_; absl::flat_hash_map<absl::string_view, std::string> variables_; + // A map from bit_field index (e.g. 0 for hasBits0) to the mask + // representing the bits to clear for this oneof's fields. + std::map<int, uint32_t> clear_masks_; + // A map from bit_field index to the mask of scalar fields for this oneof. + std::map<int, uint32_t> scalar_masks_; }; } // namespace java
diff --git a/src/google/protobuf/compiler/java/full/primitive_field.cc b/src/google/protobuf/compiler/java/full/primitive_field.cc index 7a34996..a07ca99 100644 --- a/src/google/protobuf/compiler/java/full/primitive_field.cc +++ b/src/google/protobuf/compiler/java/full/primitive_field.cc
@@ -578,7 +578,15 @@ "$deprecation$public Builder " "${$set$capitalized_name$$}$($type$ value) {\n" " $null_check$\n" - " $set_oneof_case_message$;\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" " $oneof_name$_ = value;\n" " $on_changed$\n" " return this;\n" @@ -633,7 +641,15 @@ io::Printer* printer) const { printer->Print(variables_, "$oneof_name$_ = input.read$capitalized_type$();\n" - "$set_oneof_case_message$;\n"); + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n"); } void ImmutablePrimitiveOneofFieldGenerator::GenerateSerializationCode(
diff --git a/src/google/protobuf/compiler/java/full/string_field.cc b/src/google/protobuf/compiler/java/full/string_field.cc index 3309ad9..fbe542b 100644 --- a/src/google/protobuf/compiler/java/full/string_field.cc +++ b/src/google/protobuf/compiler/java/full/string_field.cc
@@ -654,7 +654,15 @@ "$deprecation$public Builder ${$set$capitalized_name$$}$(\n" " java.lang.String value) {\n" " $null_check$\n" - " $set_oneof_case_message$;\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" " $oneof_name$_ = value;\n" " $on_changed$\n" " return this;\n" @@ -695,7 +703,15 @@ printer->Print(variables_, " checkByteStringIsUtf8(value);\n"); } printer->Print(variables_, - " $set_oneof_case_message$;\n" + " switch ($oneof_name$Case_) {\n" + " default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + " case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + " case $number$:\n" + " break;\n" + " }\n" " $oneof_name$_ = value;\n" " $on_changed$\n" " return this;\n" @@ -722,7 +738,15 @@ // Allow a slight breach of abstraction here in order to avoid forcing // all string fields to Strings when copying fields from a Message. printer->Print(variables_, - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "$oneof_name$_ = other.$oneof_name$_;\n" "$on_changed$\n"); } @@ -736,14 +760,30 @@ io::Printer* printer) const { if (CheckUtf8(descriptor_)) { printer->Print(variables_, - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "$oneof_name$_ = " "input.readStringRequireUtf8();\n" ); } else { printer->Print(variables_, "com.google.protobuf.ByteString bs = input.readBytes();\n" - "$set_oneof_case_message$;\n" + "switch ($oneof_name$Case_) {\n" + "default:\n" + " clear$oneof_capitalized_name$HasBits(); // fallthrough\n" + "case 0:\n" + " $set_oneof_case_message$;\n" + " $set_has_field_bit$ // fallthrough\n" + "case $number$:\n" + " break;\n" + "}\n" "$oneof_name$_ = bs;\n"); } }
diff --git a/src/google/protobuf/compiler/java/generator_common.h b/src/google/protobuf/compiler/java/generator_common.h index 7ec12c1..0ff66d9 100644 --- a/src/google/protobuf/compiler/java/generator_common.h +++ b/src/google/protobuf/compiler/java/generator_common.h
@@ -30,11 +30,12 @@ public: explicit FieldGeneratorMap(const Descriptor* descriptor) : descriptor_(descriptor) { - field_generators_.reserve(static_cast<size_t>(descriptor->field_count())); + insert_order_.reserve(static_cast<size_t>(descriptor->field_count())); + index_order_.resize(static_cast<size_t>(descriptor->field_count())); } ~FieldGeneratorMap() { - for (const auto* g : field_generators_) { + for (const auto* g : insert_order_) { delete g; } } @@ -45,21 +46,28 @@ FieldGeneratorMap(const FieldGeneratorMap&) = delete; FieldGeneratorMap& operator=(const FieldGeneratorMap&) = delete; + size_t size() const { return insert_order_.size(); } + void Add(const FieldDescriptor* field, std::unique_ptr<FieldGeneratorType> field_generator) { ABSL_CHECK_EQ(field->containing_type(), descriptor_); - field_generators_.push_back(field_generator.release()); + insert_order_.push_back(field_generator.release()); + index_order_[static_cast<size_t>(field->index())] = insert_order_.back(); } const FieldGeneratorType& get(const FieldDescriptor* field) const { ABSL_CHECK_EQ(field->containing_type(), descriptor_); - return *field_generators_[static_cast<size_t>(field->index())]; + return *index_order_[static_cast<size_t>(field->index())]; + } + + const FieldGeneratorType& getInInsertOrder(int index) const { + return *insert_order_[static_cast<size_t>(index)]; } std::vector<const FieldGenerator*> field_generators() const { std::vector<const FieldGenerator*> field_generators; - field_generators.reserve(field_generators_.size()); - for (const auto* g : field_generators_) { + field_generators.reserve(index_order_.size()); + for (const auto* g : index_order_) { field_generators.push_back(g); } return field_generators; @@ -67,7 +75,8 @@ private: const Descriptor* descriptor_; - std::vector<const FieldGeneratorType*> field_generators_; + std::vector<const FieldGeneratorType*> insert_order_; + std::vector<const FieldGeneratorType*> index_order_; }; inline void ReportUnexpectedPackedFieldsCall() {
diff --git a/src/google/protobuf/compiler/java/generator_unittest.cc b/src/google/protobuf/compiler/java/generator_unittest.cc index c3c2096..a61128a 100644 --- a/src/google/protobuf/compiler/java/generator_unittest.cc +++ b/src/google/protobuf/compiler/java/generator_unittest.cc
@@ -529,6 +529,44 @@ "foo.FooProto.fileOpt);"))); } +} + +TEST_F(JavaGeneratorTest, OneofHasBitsBuildPartial) { + CreateTempFile("oneof_has_bits.proto", + R"schema( + syntax = "proto2"; + package com.google.protos; + option java_multiple_files = true; + message SubMsg { + optional int32 x = 1; + } + message TestOneofHasBits { + optional int32 regular_int = 1; + oneof my_oneof { + int32 oneof_int = 2; + string oneof_str = 3; + SubMsg oneof_msg = 4; + } + })schema"); + + RunProtoc( + "protocol_compiler --proto_path=$tmpdir --java_out=$tmpdir " + "oneof_has_bits.proto"); + + ExpectNoErrors(); + EXPECT_TRUE( + FileGenerated(PACKAGE_PREFIX "com/google/protos/TestOneofHasBits.java")); + // Verify message typed field in oneof is treated like regular message typed + // fields (checked via from_bitField0_ & 0x00000008). + EXPECT_TRUE(FileContainsSubstring(PACKAGE_PREFIX + "com/google/protos/TestOneofHasBits.java", + "((from_bitField0_ & 0x00000008) != 0)")); + // Verify non-message oneof fields are grouped together by bitField + // membership (0x00000002 | 0x00000004 = 0x00000006). + EXPECT_TRUE(FileContainsSubstring( + PACKAGE_PREFIX "com/google/protos/TestOneofHasBits.java", + "to_bitField0_ |= (from_bitField0_ & 0x00000006);")); +} } // namespace } // namespace java