From 266ec3554c83f22e857ba6ba5496038858262898 Mon Sep 17 00:00:00 2001 From: Filip Niksic Date: Wed, 26 Aug 2026 14:32:41 -0700 Subject: [PATCH] Store the output domain as part of the corpus value in FlatMap. This greatly improves the efficiency of FlatMap. Before the CL, the majority of time when using this domain was spent in calls to `GetOutputDomain`. PiperOrigin-RevId: 971505250 --- domain_tests/map_filter_combinator_test.cc | 26 +++- fuzztest/internal/domains/flat_map_impl.h | 144 +++++++++++++-------- fuzztest/internal/type_support.h | 21 --- fuzztest/internal/type_support_test.cc | 20 +-- 4 files changed, 119 insertions(+), 92 deletions(-) diff --git a/domain_tests/map_filter_combinator_test.cc b/domain_tests/map_filter_combinator_test.cc index 353b1230f..641140f26 100644 --- a/domain_tests/map_filter_combinator_test.cc +++ b/domain_tests/map_filter_combinator_test.cc @@ -191,7 +191,7 @@ TEST(FlatMap, WorksWithSameCorpusType) { auto domain = FlatMap([](int a) { return Just(~a); }, Arbitrary()); absl::BitGen bitgen; Value value(domain, bitgen); - EXPECT_EQ(value.user_value, ~std::get<1>(value.corpus_value)); + EXPECT_EQ(value.user_value, ~std::get<2>(value.corpus_value)); } TEST(FlatMap, WorksWithDifferentCorpusType) { @@ -206,7 +206,7 @@ TEST(FlatMap, WorksWithDifferentCorpusType) { Value value(domain, bitgen); // `0` is the index in the ElementOf EXPECT_EQ(typename decltype(colors)::corpus_type{0}, - std::get<1>(value.corpus_value)); + std::get<2>(value.corpus_value)); EXPECT_EQ("Blue", value.user_value); } @@ -229,7 +229,13 @@ TEST(FlatMap, SerializationRoundTrip) { absl::BitGen bitgen; Value value(domain, bitgen); auto serialized = domain.SerializeCorpus(value.corpus_value); - EXPECT_EQ(domain.ParseCorpus(serialized), value.corpus_value); + auto parsed = domain.ParseCorpus(serialized); + ASSERT_TRUE(parsed.has_value()); + // Corpus value is a tuple: + // (output_domain, output_corpus_val, input_corpus_val...) + // We ignore the output domain itself since it doesn't have equality defined. + EXPECT_EQ(std::get<1>(*parsed), std::get<1>(value.corpus_value)); + EXPECT_EQ(std::get<2>(*parsed), std::get<2>(value.corpus_value)); } TEST(FlatMap, ValidationRejectsInvalidValue) { @@ -260,13 +266,13 @@ TEST(FlatMap, MutationAcceptsChangingDomains) { absl::BitGen bitgen; Value value(domain, bitgen); auto mutated = value.corpus_value; - while (std::get<1>(value.corpus_value) == std::get<1>(mutated)) { + while (std::get<2>(value.corpus_value) == std::get<2>(mutated)) { // We demand that our output domain has size `len` above. This will check // fail in ContainerOfImpl if we try to generate a string of the wrong // length. domain.Mutate(mutated, bitgen, {}, false); } - EXPECT_EQ(domain.GetValue(mutated).size(), std::get<1>(mutated)); + EXPECT_EQ(domain.GetValue(mutated).size(), std::get<2>(mutated)); } TEST(FlatMap, MutationAcceptsShrinkingOutputDomains) { @@ -484,7 +490,7 @@ TEST(ReversibleFlatMap, WorksWithSameCorpusType) { absl::BitGen bitgen; Value value(domain, bitgen); // Corpus value is a tuple: (output_corpus, input_corpus...) - EXPECT_EQ(value.user_value, ~std::get<1>(value.corpus_value)); + EXPECT_EQ(value.user_value, ~std::get<2>(value.corpus_value)); } TEST(ReversibleFlatMap, AcceptsMultipleInnerDomains) { @@ -547,7 +553,13 @@ TEST(ReversibleFlatMap, SerializationRoundTrip) { absl::BitGen bitgen; Value value(domain, bitgen); auto serialized = domain.SerializeCorpus(value.corpus_value); - EXPECT_EQ(domain.ParseCorpus(serialized), value.corpus_value); + auto parsed = domain.ParseCorpus(serialized); + ASSERT_TRUE(parsed.has_value()); + // Corpus value is a tuple: + // (output_domain, output_corpus_val, input_corpus_val...) + // We ignore the output domain itself since it doesn't have equality defined. + EXPECT_EQ(std::get<1>(*parsed), std::get<1>(value.corpus_value)); + EXPECT_EQ(std::get<2>(*parsed), std::get<2>(value.corpus_value)); } TEST(ReversibleFlatMap, ParseCorpusRejectsInvalidInputValues) { diff --git a/fuzztest/internal/domains/flat_map_impl.h b/fuzztest/internal/domains/flat_map_impl.h index 37f0cb98e..ecf52ad4a 100644 --- a/fuzztest/internal/domains/flat_map_impl.h +++ b/fuzztest/internal/domains/flat_map_impl.h @@ -16,6 +16,8 @@ #define FUZZTEST_FUZZTEST_INTERNAL_DOMAINS_FLAT_MAP_IMPL_H_ #include +#include +#include #include #include #include @@ -30,9 +32,9 @@ #include "./fuzztest/internal/domains/serialization_helpers.h" #include "./fuzztest/internal/logging.h" #include "./fuzztest/internal/meta.h" +#include "./fuzztest/internal/printer.h" #include "./fuzztest/internal/serialization.h" #include "./fuzztest/internal/status.h" -#include "./fuzztest/internal/type_support.h" namespace fuzztest::internal { @@ -61,10 +63,11 @@ class FlatMapImplBase Derived, // The user value is the user value of the output domain. value_type_t>, - // The corpus value is a tuple where the first element is the corpus - // value of the output domain, and the rest is the corpus value of the - // input domains. + // The corpus value is a tuple where the first element is the output + // domain itself, the second element is the corpus value of the output + // domain, and the rest are the corpus values of the input domains. std::tuple< + FlatMapOutputDomain, corpus_type_t>, corpus_type_t...>> { public: @@ -78,14 +81,16 @@ class FlatMapImplBase corpus_type Init(absl::BitGenRef prng) { if (auto seed = this->MaybeGetRandomSeed(prng)) return *seed; - auto input_corpus = std::apply( + auto input_corpus_vals = std::apply( [&](auto&... input_domains) { - return std::make_tuple(input_domains.Init(prng)...); + return std::tuple{input_domains.Init(prng)...}; }, input_domains_); - auto output_domain = GetOutputDomain(input_corpus); - return std::tuple_cat(std::make_tuple(output_domain.Init(prng)), - input_corpus); + auto output_domain = GetOutputDomain(input_corpus_vals); + auto output_corpus_val = output_domain.Init(prng); + return std::tuple_cat( + std::tuple{std::move(output_domain), std::move(output_corpus_val)}, + std::move(input_corpus_vals)); } void Mutate(corpus_type& val, absl::BitGenRef prng, @@ -99,54 +104,68 @@ class FlatMapImplBase bool mutate_inputs = !only_shrink && absl::Bernoulli(prng, 0.1); if (mutate_inputs) { ApplyIndex([&](auto... I) { - // The first field of `val` is the output corpus value, so skip it. + // The first two fields of `val` are the output domain and the output + // corpus value, so skip them. (std::get(input_domains_) - .Mutate(std::get(val), prng, metadata, only_shrink), + .Mutate(std::get(val), prng, metadata, only_shrink), ...); }); - std::get<0>(val) = GetOutputDomain(val).Init(prng); + // Generate a new output domain and store it as `std::get<0>(val)`. + // We can't write `std::get<0>(val) = GetOutputDomain(val)` because + // there are domains that don't support copy-assignment. So we manually + // destroy the old domain and construct a new one in place. + std::destroy_at(&std::get<0>(val)); + ::new (static_cast(&std::get<0>(val))) + FlatMapOutputDomain(GetOutputDomain(val)); + std::get<1>(val) = std::get<0>(val).Init(prng); return; } - // For simplicity, we create a new output domain each call to `Mutate`. This - // means that stateful domains don't work, but this is currently a matter of - // convenience, not correctness. For example, `Filter` won't automatically - // find when something is too restrictive. - // TODO(b/246423623): Support stateful domains. - GetOutputDomain(val).Mutate(std::get<0>(val), prng, metadata, only_shrink); + std::get<0>(val).Mutate(std::get<1>(val), prng, metadata, only_shrink); } value_type GetValue(const corpus_type& v) const { - return GetOutputDomain(v).GetValue(std::get<0>(v)); + return std::get<0>(v).GetValue(std::get<1>(v)); } - auto GetPrinter() const { - return FlatMappedPrinter{flat_mapper_, - input_domains_}; - } + auto GetPrinter() const { return Printer{input_domains_}; } std::optional ParseCorpus(const IRObject& obj) const { - auto input_corpus = ParseWithDomainTuple(input_domains_, obj, /*skip=*/1); - if (!input_corpus.has_value()) { + auto input_corpus_vals = + ParseWithDomainTuple(input_domains_, obj, /*skip=*/1); + if (!input_corpus_vals.has_value()) { return std::nullopt; } - absl::Status input_values_validity = ValidateInputValues(*input_corpus); + absl::Status input_values_validity = + ValidateInputValues(*input_corpus_vals); if (!input_values_validity.ok()) { absl::FPrintF(GetStderr(), "[!] %s", input_values_validity.message()); return std::nullopt; } - auto output_domain = GetOutputDomain(*input_corpus); + auto output_domain = GetOutputDomain(*input_corpus_vals); // We know obj.Subs()[0] exists because ParseWithDomainTuple succeeded. - auto output_corpus = output_domain.ParseCorpus((*obj.Subs())[0]); - if (!output_corpus.has_value()) { + auto output_corpus_val = output_domain.ParseCorpus((*obj.Subs())[0]); + if (!output_corpus_val.has_value()) { return std::nullopt; } - return std::tuple_cat(std::make_tuple(*output_corpus), *input_corpus); + return std::tuple_cat( + std::tuple{std::move(output_domain), *std::move(output_corpus_val)}, + *std::move(input_corpus_vals)); } IRObject SerializeCorpus(const corpus_type& v) const { - auto domain = - std::tuple_cat(std::make_tuple(GetOutputDomain(v)), input_domains_); - return SerializeWithDomainTuple(domain, v); + IRObject obj; + auto& subs = obj.MutableSubs(); + + // 1. Serialize the output corpus value. + subs.push_back(std::get<0>(v).SerializeCorpus(std::get<1>(v))); + + // 2. Serialize the input corpus values. + ApplyIndex([&](auto... I) { + (subs.push_back( + std::get(input_domains_).SerializeCorpus(std::get(v))), + ...); + }); + return obj; } absl::Status ValidateCorpusValue(const corpus_type& corpus_value) const { @@ -154,8 +173,8 @@ class FlatMapImplBase absl::Status input_values_validity = ValidateInputValues(corpus_value); if (!input_values_validity.ok()) return input_values_validity; // Check the output value. - return GetOutputDomain(corpus_value) - .ValidateCorpusValue(std::get<0>(corpus_value)); + return std::get<0>(corpus_value) + .ValidateCorpusValue(std::get<1>(corpus_value)); } protected: @@ -164,8 +183,8 @@ class FlatMapImplBase } static constexpr size_t kNumInputValues = sizeof...(InputDomain); - // Returns the output domain for a `tuple` with or without the output value - // as the leading element, and with the input values as the last + // Returns the output domain for a `tuple` with or without the output domain + // and value as the leading elements, and with the input values as the last // `kNumInputValues` elements. template FlatMapOutputDomain GetOutputDomain( @@ -181,8 +200,8 @@ class FlatMapImplBase }); } - // Validates the input values for a `tuple` with or without the output value - // as the leading element, and with the input values as the last + // Validates the input values for a `tuple` with or without the output domain + // and value as the leading elements, and with the input values as the last // `kNumInputValues` elements. template absl::Status ValidateInputValues(const Tuple& tuple) const { @@ -208,6 +227,19 @@ class FlatMapImplBase } private: + struct Printer { + const std::tuple& input_domains; + + void PrintCorpusValue(const corpus_type& corpus_value, + domain_implementor::RawSink out, + domain_implementor::PrintMode mode) const { + // There is no useful way to print the input values, so we just print the + // output value by delegating to the output domain. + domain_implementor::PrintValue(std::get<0>(corpus_value), + std::get<1>(corpus_value), out, mode); + } + }; + FlatMapper flat_mapper_; std::tuple input_domains_; }; @@ -263,40 +295,42 @@ class ReversibleFlatMapImpl std::optional FromValue(const value_type& v) const { // 1. Recover the input values using the user-provided inverse mapper. - auto input_values_opt = std::invoke(inv_mapper_, v); - if (!input_values_opt.has_value()) return std::nullopt; + auto input_user_vals = std::invoke(inv_mapper_, v); + if (!input_user_vals.has_value()) return std::nullopt; - // 2. Map input values into input corpus values. - auto input_corpus_opt = + // 2. Map input user values into input corpus values. + auto input_corpus_vals = ApplyIndex( [&](auto... I) -> std::optional...>> { auto inner_corpus_vals = std::tuple{std::get(this->input_domains()) - .FromValue(std::get(*input_values_opt))...}; + .FromValue(std::get(*input_user_vals))...}; bool has_nullopt = (!std::get(inner_corpus_vals).has_value() || ...); if (has_nullopt) return std::nullopt; return std::tuple{*std::move(std::get(inner_corpus_vals))...}; }); - if (!input_corpus_opt.has_value()) return std::nullopt; + if (!input_corpus_vals.has_value()) return std::nullopt; - if (!this->ValidateInputValues(*input_corpus_opt).ok()) return std::nullopt; + if (!this->ValidateInputValues(*input_corpus_vals).ok()) + return std::nullopt; // 3. Re-instantiate the dynamically generated output domain. - auto output_domain = this->GetOutputDomain(*input_corpus_opt); + auto output_domain = this->GetOutputDomain(*input_corpus_vals); - // 4. Map the output value into the output corpus value. - auto output_corpus_opt = output_domain.FromValue(v); - if (!output_corpus_opt.has_value()) return std::nullopt; + // 4. Map the output user value into the output corpus value. + auto output_corpus_val = output_domain.FromValue(v); + if (!output_corpus_val.has_value()) return std::nullopt; - if (!output_domain.ValidateCorpusValue(*output_corpus_opt).ok()) { + if (!output_domain.ValidateCorpusValue(*output_corpus_val).ok()) { return std::nullopt; } - // 5. Assemble the final corpus tuple (output corpus followed by input - // corpus). - return std::tuple_cat(std::make_tuple(*std::move(output_corpus_opt)), - *std::move(input_corpus_opt)); + + // 5. Assemble the final corpus tuple. + return std::tuple_cat( + std::tuple{std::move(output_domain), *std::move(output_corpus_val)}, + *std::move(input_corpus_vals)); } private: diff --git a/fuzztest/internal/type_support.h b/fuzztest/internal/type_support.h index 0273b7e13..5ab6de48b 100644 --- a/fuzztest/internal/type_support.h +++ b/fuzztest/internal/type_support.h @@ -525,27 +525,6 @@ struct MappedPrinter { } }; -template -struct FlatMappedPrinter { - const FlatMapper& mapper; - const std::tuple& inner; - - template - void PrintCorpusValue(const CorpusT& corpus_value, - domain_implementor::RawSink out, - domain_implementor::PrintMode mode) const { - auto output_domain = ApplyIndex([&](auto... I) { - return mapper( - // the first field of `corpus_value` is the output value, so skip it - std::get(inner).GetValue(std::get(corpus_value))...); - }); - - // Delegate to the output domain's printer. - domain_implementor::PrintValue(output_domain, std::get<0>(corpus_value), - out, mode); - } -}; - struct DurationPrinter { void PrintUserValue(const absl::Duration duration, domain_implementor::RawSink out, diff --git a/fuzztest/internal/type_support_test.cc b/fuzztest/internal/type_support_test.cc index 86692368c..454bf78ef 100644 --- a/fuzztest/internal/type_support_test.cc +++ b/fuzztest/internal/type_support_test.cc @@ -484,6 +484,8 @@ TEST(FlatMapTest, DelegatesToOutputDomainPrinter) { auto flat_map_domain = FlatMap(optional_sized_strings, input_domain); corpus_type_t abc_corpus_val = { + // Output domain + optional_sized_strings(3), // String of size GenericDomainCorpusType(std::in_place_type, "ABC"), // Size @@ -491,15 +493,16 @@ TEST(FlatMapTest, DelegatesToOutputDomainPrinter) { // Sanity checks that the components of `abc_corpus_val` are in the respective // domains. ASSERT_TRUE( - input_domain.ValidateCorpusValue(std::get<1>(abc_corpus_val)).ok()); - ASSERT_TRUE( - optional_sized_strings(input_domain.GetValue(std::get<1>(abc_corpus_val))) - .ValidateCorpusValue(std::get<0>(abc_corpus_val)) - .ok()); + input_domain.ValidateCorpusValue(std::get<2>(abc_corpus_val)).ok()); + ASSERT_TRUE(std::get<0>(abc_corpus_val) + .ValidateCorpusValue(std::get<1>(abc_corpus_val)) + .ok()); EXPECT_THAT(TestPrintValue(abc_corpus_val, flat_map_domain), ElementsAre("(\"ABC\")", "\"ABC\"")); corpus_type_t nullopt_corpus_val = { + // Output domain + optional_sized_strings(2), // Corpus value of nullopt std::monostate{}, // Size (here irrelevant) @@ -507,10 +510,9 @@ TEST(FlatMapTest, DelegatesToOutputDomainPrinter) { // Sanity checks that the components of `nullopt_corpus_val` are in the // respective domains. ASSERT_TRUE( - input_domain.ValidateCorpusValue(std::get<1>(nullopt_corpus_val)).ok()); - ASSERT_TRUE(optional_sized_strings( - input_domain.GetValue(std::get<1>(nullopt_corpus_val))) - .ValidateCorpusValue(std::get<0>(nullopt_corpus_val)) + input_domain.ValidateCorpusValue(std::get<2>(nullopt_corpus_val)).ok()); + ASSERT_TRUE(std::get<0>(nullopt_corpus_val) + .ValidateCorpusValue(std::get<1>(nullopt_corpus_val)) .ok()); EXPECT_THAT(TestPrintValue(nullopt_corpus_val, flat_map_domain), Each("std::nullopt"));