Initialize all required fields customized by users during protobuf Init.

PiperOrigin-RevId: 715371239
diff --git a/e2e_tests/functional_test.cc b/e2e_tests/functional_test.cc
index a5ccbbf..d510252 100644
--- a/e2e_tests/functional_test.cc
+++ b/e2e_tests/functional_test.cc
@@ -469,6 +469,12 @@
   EXPECT_THAT(std_err, HasSubstr("recursive fields"));
 }
 
+TEST_F(UnitTestModeTest, InitializesRequiredFieldsEvenIfTheyAreRecursive) {
+  auto [status, std_out, std_err] =
+      Run("MySuite.FailIfRequiredRecursiveFieldsAreUnset");
+  EXPECT_THAT(status, Eq(ExitCode(0)));
+}
+
 TEST_F(UnitTestModeTest,
        AvoidsFailureIfSetByDefaultPolicyIsOverwrittenOnRecursiveStructures) {
   auto [status, std_out, std_err] =
diff --git a/e2e_tests/testdata/fuzz_tests_for_functional_testing.cc b/e2e_tests/testdata/fuzz_tests_for_functional_testing.cc
index fc88e68..5b661e3 100644
--- a/e2e_tests/testdata/fuzz_tests_for_functional_testing.cc
+++ b/e2e_tests/testdata/fuzz_tests_for_functional_testing.cc
@@ -444,6 +444,10 @@
   return absl::StrContains(field->name(), "parent");
 }
 
+bool IsParent1(const FieldDescriptor* field) {
+  return absl::StrContains(field->name(), "parent1");
+}
+
 void FailsIfCantInitializeProto(const TestProtobufWithRecursion& proto) {}
 FUZZ_TEST(MySuite, FailsIfCantInitializeProto)
     .WithDomains(Arbitrary<TestProtobufWithRecursion>()
@@ -451,6 +455,14 @@
                      .WithFieldsUnset(IsChildId)
                      .WithFieldUnset("id"));
 
+void FailIfRequiredRecursiveFieldsAreUnset(
+    const TestProtobufWithRecursion& proto) {
+  if (proto.has_child() && !proto.child().has_parent1()) std::abort();
+}
+FUZZ_TEST(MySuite, FailIfRequiredRecursiveFieldsAreUnset)
+    .WithDomains(
+        Arbitrary<TestProtobufWithRecursion>().WithFieldsAlwaysSet(IsParent1));
+
 void InitializesRecursiveProtoIfInfiniteRecursivePolicyIsOverwritten(
     const TestProtobufWithRecursion& proto) {}
 FUZZ_TEST(MySuite,
diff --git a/fuzztest/internal/domains/protobuf_domain_impl.h b/fuzztest/internal/domains/protobuf_domain_impl.h
index acd1ef1..42da939 100644
--- a/fuzztest/internal/domains/protobuf_domain_impl.h
+++ b/fuzztest/internal/domains/protobuf_domain_impl.h
@@ -457,7 +457,7 @@
   corpus_type Init(absl::BitGenRef prng) {
     if (auto seed = this->MaybeGetRandomSeed(prng)) return *seed;
     FUZZTEST_INTERNAL_CHECK(
-        !IsCustomizedRecursivelyOnly() || !IsNonTerminatingRecursive(),
+        !IsCustomizedRecursivelyOnly() || !IsInfinitelyRecursive(),
         "Cannot set recursive fields by default.");
     const auto* descriptor = prototype_.Get()->GetDescriptor();
     corpus_type val;
@@ -467,13 +467,12 @@
     for (const FieldDescriptor* field : GetProtobufFields(descriptor)) {
       if (auto* oneof = field->containing_oneof()) {
         if (!oneof_to_field.contains(oneof->index())) {
-          oneof_to_field[oneof->index()] = SelectAFieldIndexInOneof(
-              oneof, prng,
-              /*non_recursive_only=*/IsCustomizedRecursivelyOnly());
+          oneof_to_field[oneof->index()] =
+              SelectAFieldIndexInOneof(oneof, prng);
         }
         if (oneof_to_field[oneof->index()] != field->index()) continue;
-      } else if (!IsRequired(field) && IsCustomizedRecursivelyOnly() &&
-                 IsFieldRecursive(field)) {
+      } else if (!MustBeSet(field) && IsCustomizedRecursivelyOnly() &&
+                 IsFieldFinitelyRecursive(field)) {
         // We avoid initializing non-required recursive fields by default (if
         // they are not explicitly customized). Otherwise, the initialization
         // may never terminate. If a proto has only non-required recursive
@@ -944,12 +943,15 @@
 
   template <typename OneofDescriptor>
   int SelectAFieldIndexInOneof(const OneofDescriptor* oneof,
-                               absl::BitGenRef prng, bool non_recursive_only) {
+                               absl::BitGenRef prng) {
     std::vector<int> fields;
     for (int i = 0; i < oneof->field_count(); ++i) {
       OptionalPolicy policy = GetOneofFieldPolicy(oneof->field(i));
       if (policy == OptionalPolicy::kAlwaysNull) continue;
-      if (non_recursive_only && IsFieldRecursive(oneof->field(i))) continue;
+      if (IsCustomizedRecursivelyOnly() &&
+          IsFieldFinitelyRecursive(oneof->field(i))) {
+        continue;
+      }
       fields.push_back(i);
     }
     if (fields.empty()) {  // This can happen if all fields are unset.
@@ -1506,8 +1508,11 @@
       return ModifyDomainForRequiredFieldRule(std::move(domain));
     } else {
       return ModifyDomainForOptionalFieldRule(
-          std::move(domain), use_policy ? policy_.GetOptionalPolicy(field)
-                                        : OptionalPolicy::kWithNull);
+          std::move(domain),
+          use_policy
+              ? (field->containing_oneof() ? GetOneofFieldPolicy(field)
+                                           : policy_.GetOptionalPolicy(field))
+              : OptionalPolicy::kWithNull);
     }
   }
 
@@ -1682,13 +1687,21 @@
     return GetDomainForField<T, is_repeated>(field, /*use_policy=*/false);
   }
 
-  bool IsNonTerminatingRecursive() {
+  // Analysis type for protobuf recursions.
+  enum class RecursionType {
+    // The proto contains a proto of type P, that must contain another P.
+    kInfinitelyRecursive,
+    // The proto contains a proto of type P, that can contain another P.
+    kFinitelyRecursive,
+  };
+
+  bool IsInfinitelyRecursive() {
     absl::flat_hash_set<decltype(prototype_.Get()->GetDescriptor())> parents;
     return IsProtoRecursive(prototype_.Get()->GetDescriptor(), parents,
-                            /*consider_non_terminating_recursions=*/true);
+                            RecursionType::kInfinitelyRecursive);
   }
 
-  bool IsFieldRecursive(const FieldDescriptor* field) {
+  bool IsFieldFinitelyRecursive(const FieldDescriptor* field) {
     if (!field->message_type()) return false;
     static absl::NoDestructor<absl::flat_hash_map<const FieldDescriptor*, bool>>
         cache;
@@ -1698,9 +1711,8 @@
       if (it != cache->end()) return it->second;
     }
     absl::flat_hash_set<decltype(field->message_type())> parents;
-    bool result =
-        IsProtoRecursive(field->message_type(), parents,
-                         /*consider_non_terminating_recursions=*/false);
+    bool result = IsProtoRecursive(field->message_type(), parents,
+                                   RecursionType::kFinitelyRecursive);
     if (IsCustomizedRecursivelyOnly()) cache->insert(it, {field, result});
     return result;
   }
@@ -1715,24 +1727,22 @@
 
   bool IsOneofRecursive(const OneofDescriptor* oneof,
                         absl::flat_hash_set<const Descriptor*>& parents,
-                        bool consider_non_terminating_recursions) const {
+                        RecursionType recursion_type) const {
     bool is_oneof_recursive = false;
     for (int i = 0; i < oneof->field_count(); ++i) {
       const auto* field = oneof->field(i);
       const auto field_policy = policy_.GetOptionalPolicy(field);
       if (field_policy == OptionalPolicy::kAlwaysNull) continue;
       const auto* child = field->message_type();
-      if (consider_non_terminating_recursions) {
-        is_oneof_recursive =
-            field_policy != OptionalPolicy::kWithNull && child &&
-            IsProtoRecursive(child, parents,
-                             consider_non_terminating_recursions);
+      if (recursion_type == RecursionType::kInfinitelyRecursive) {
+        is_oneof_recursive = field_policy != OptionalPolicy::kWithNull &&
+                             child &&
+                             IsProtoRecursive(child, parents, recursion_type);
         if (!is_oneof_recursive) {
           return false;
         }
       } else {
-        if (child && IsProtoRecursive(child, parents,
-                                      consider_non_terminating_recursions)) {
+        if (child && IsProtoRecursive(child, parents, recursion_type)) {
           return true;
         }
       }
@@ -1740,16 +1750,47 @@
     return is_oneof_recursive;
   }
 
+  bool MustBeSet(const FieldDescriptor* field) const {
+    if (IsRequired(field)) {
+      return true;
+    } else if (field->containing_oneof()) {
+      return GetOneofFieldPolicy(field) == OptionalPolicy::kWithoutNull;
+    } else if (field->is_optional()) {
+      return policy_.GetOptionalPolicy(field) == OptionalPolicy::kWithoutNull;
+    } else if (field->is_repeated()) {
+      return policy_.GetMinRepeatedFieldSize(field).has_value() &&
+             *policy_.GetMinRepeatedFieldSize(field) > 0;
+    }
+    FUZZTEST_INTERNAL_CHECK(false,
+                            "Field is not optional, repeated, or required");
+    return false;
+  }
+
+  bool MustBeUnset(const FieldDescriptor* field) const {
+    if (IsRequired(field)) {
+      return false;
+    } else if (field->containing_oneof()) {
+      return GetOneofFieldPolicy(field) == OptionalPolicy::kAlwaysNull;
+    } else if (field->is_optional()) {
+      return policy_.GetOptionalPolicy(field) == OptionalPolicy::kAlwaysNull;
+    } else if (field->is_repeated()) {
+      return policy_.GetMaxRepeatedFieldSize(field).has_value() &&
+             *policy_.GetMaxRepeatedFieldSize(field) == 0;
+    }
+    FUZZTEST_INTERNAL_CHECK(false,
+                            "Field is not optional, repeated, or required");
+    return false;
+  }
+
   template <typename Descriptor>
   bool IsProtoRecursive(const Descriptor* descriptor,
                         absl::flat_hash_set<const Descriptor*>& parents,
-                        bool consider_non_terminating_recursions) const {
+                        RecursionType recursion_type) const {
     if (parents.contains(descriptor)) return true;
     parents.insert(descriptor);
     for (int i = 0; i < descriptor->oneof_decl_count(); ++i) {
       const auto* oneof = descriptor->oneof_decl(i);
-      if (IsOneofRecursive(oneof, parents,
-                           consider_non_terminating_recursions)) {
+      if (IsOneofRecursive(oneof, parents, recursion_type)) {
         parents.erase(descriptor);
         return true;
       }
@@ -1763,27 +1804,12 @@
         // its default domain. Otherwise, this field can always be set safely.
         continue;
       }
-      if (consider_non_terminating_recursions) {
-        const bool should_be_set =
-            IsRequired(field) ||
-            (field->is_optional() && policy_.GetOptionalPolicy(field) ==
-                                         OptionalPolicy::kWithoutNull) ||
-            (field->is_repeated() &&
-             policy_.GetMinRepeatedFieldSize(field).has_value() &&
-             *policy_.GetMinRepeatedFieldSize(field) > 0);
-        if (!should_be_set) continue;
+      if (recursion_type == RecursionType::kInfinitelyRecursive) {
+        if (!MustBeSet(field)) continue;
       } else {
-        const bool can_be_set =
-            IsRequired(field) ||
-            (field->is_optional() &&
-             policy_.GetOptionalPolicy(field) != OptionalPolicy::kAlwaysNull) ||
-            (field->is_repeated() &&
-             (!policy_.GetMaxRepeatedFieldSize(field).has_value() ||
-              *policy_.GetMaxRepeatedFieldSize(field) > 0));
-        if (!can_be_set) continue;
+        if (MustBeUnset(field)) continue;
       }
-      if (IsProtoRecursive(child, parents,
-                           consider_non_terminating_recursions)) {
+      if (IsProtoRecursive(child, parents, recursion_type)) {
         parents.erase(descriptor);
         return true;
       }
@@ -1793,10 +1819,6 @@
   }
 
   bool IsRequired(const FieldDescriptor* field) const {
-    if (field->containing_oneof() &&
-        GetOneofFieldPolicy(field) == OptionalPolicy::kWithoutNull) {
-      return true;
-    }
     return field->is_required() || IsMapValueMessage(field);
   }