Refactor FlatbuffersTableUntypedDomainImpl::MutateSelectedField to use a visitor and clean up CountNumberOfMutableFieldsVisitor. Specifically: * Introduce `MutateSelectedFieldVisitor` to encapsulate mutation logic for specific field indices, passing nested table recursion through the visitor pattern. * Clean up `CountNumberOfMutableFieldsVisitor` and rename fields (`total_weight` -> `field_count`, `val` -> `corpus_value`) to better reflect their semantic meaning. * Limit recursive mutations and field counting to nested tables (`FlatbuffersTableTag`) via `std::is_same_v` constexpr checks. * Add TODO markers referencing b/418993255 for extending support to structs, unions, and containers of these types. PiperOrigin-RevId: 941534253
diff --git a/fuzztest/internal/domains/flatbuffers_domain_impl.cc b/fuzztest/internal/domains/flatbuffers_domain_impl.cc index 9031344..90980fe 100644 --- a/fuzztest/internal/domains/flatbuffers_domain_impl.cc +++ b/fuzztest/internal/domains/flatbuffers_domain_impl.cc
@@ -131,28 +131,13 @@ } for (const auto* field : *table_object_->fields()) { - if (!IsSupportedField(field)) { - if (only_shrink && !val.contains(field->id())) continue; - } + if (!IsSupportedField(field)) continue; + if (only_shrink && !val.contains(field->id())) continue; - ++field_counter; - if (field_counter == selected_field_index) { - VisitFlatbufferField( - schema_, field, - MutateVisitor{*this, prng, metadata, only_shrink, val}); - return field_counter; - } - - if (field->type()->base_type() == reflection::BaseType::Obj) { - auto sub_object = schema_->objects()->Get(field->type()->index()); - if (!sub_object->is_struct()) { - field_counter += - GetCachedDomain<FlatbuffersTableTag>(field).MutateSelectedField( - val[field->id()], prng, metadata, only_shrink, - selected_field_index - field_counter); - } - // TODO: Add support for structs. - } + VisitFlatbufferField( + schema_, field, + MutateSelectedFieldVisitor{*this, field_counter, val, prng, metadata, + only_shrink, selected_field_index}); if (field_counter >= selected_field_index) { return field_counter;
diff --git a/fuzztest/internal/domains/flatbuffers_domain_impl.h b/fuzztest/internal/domains/flatbuffers_domain_impl.h index 328b2e7..2199970 100644 --- a/fuzztest/internal/domains/flatbuffers_domain_impl.h +++ b/fuzztest/internal/domains/flatbuffers_domain_impl.h
@@ -598,42 +598,53 @@ struct CountNumberOfMutableFieldsVisitor { const FlatbuffersTableUntypedDomainImpl& self; - uint64_t& total_weight; - corpus_type& val; - bool only_shrink = false; - - template <typename T> - void Visit(const reflection::Field* absl_nonnull field) const { - if (!self.IsSupportedField(field)) return; - if (only_shrink && !val.contains(field->id())) return; - - // Add the weight of the field itself. - total_weight += 1; - - auto& domain = self.GetCachedDomain<T>(field); - if (auto it = val.find(field->id()); it != val.end()) { - // Add the weight of the field corpus. - total_weight += domain.CountNumberOfFields(it->second); - } - } - }; - - struct MutateVisitor { - FlatbuffersTableUntypedDomainImpl& self; - absl::BitGenRef prng; - const domain_implementor::MutationMetadata& metadata; - bool only_shrink; + uint64_t& field_count; corpus_type& corpus_value; + const bool only_shrink = false; template <typename T> void Visit(const reflection::Field* absl_nonnull field) { - auto& domain = self.GetCachedDomain<T>(field); + if (!self.IsSupportedField(field)) return; auto it = corpus_value.find(field->id()); - if (it == corpus_value.end()) { - if (only_shrink) return; - it = corpus_value.try_emplace(field->id(), domain.Init(prng)).first; + if (only_shrink && it == corpus_value.end()) return; + + field_count++; + + if (it == corpus_value.end()) return; + auto& domain = self.GetCachedDomain<T>(field); + field_count += domain.CountNumberOfFields(it->second); + } + }; + + struct MutateSelectedFieldVisitor { + FlatbuffersTableUntypedDomainImpl& self; + uint64_t& field_counter; + corpus_type& corpus_value; + absl::BitGenRef prng; + const domain_implementor::MutationMetadata& metadata; + const bool only_shrink; + const uint64_t selected_field_index; + + template <typename T> + void Visit(const reflection::Field* absl_nonnull field) { + if (!self.IsSupportedField(field)) return; + auto it = corpus_value.find(field->id()); + if (only_shrink && it == corpus_value.end()) return; + + field_counter++; + auto& domain = self.GetCachedDomain<T>(field); + if (field_counter == selected_field_index) { + if (it == corpus_value.end()) { + it = corpus_value.try_emplace(field->id(), domain.Init(prng)).first; + } + domain.Mutate(it->second, prng, metadata, only_shrink); + return; } - domain.Mutate(it->second, prng, metadata, only_shrink); + + if (it == corpus_value.end()) return; + field_counter += + domain.MutateSelectedField(it->second, prng, metadata, only_shrink, + selected_field_index - field_counter); } }; @@ -764,6 +775,14 @@ return inner_->CountNumberOfFields(val.untyped_corpus); } + uint64_t MutateSelectedField( + corpus_type& val, absl::BitGenRef prng, + const domain_implementor::MutationMetadata& metadata, bool only_shrink, + uint64_t selected_field_index) { + return inner_->MutateSelectedField(val.untyped_corpus, prng, metadata, + only_shrink, selected_field_index); + } + // Mutates the given corpus value. void Mutate(corpus_type& val, absl::BitGenRef prng, const domain_implementor::MutationMetadata& metadata,