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: 974593707
diff --git a/domain_tests/map_filter_combinator_test.cc b/domain_tests/map_filter_combinator_test.cc index 353b123..e45050d 100644 --- a/domain_tests/map_filter_combinator_test.cc +++ b/domain_tests/map_filter_combinator_test.cc
@@ -191,7 +191,9 @@ auto domain = FlatMap([](int a) { return Just(~a); }, Arbitrary<int>()); 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<decltype(domain)::kInputCorpusValsOffset>(value.corpus_value)); } TEST(FlatMap, WorksWithDifferentCorpusType) { @@ -205,8 +207,9 @@ absl::BitGen bitgen; 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)); + EXPECT_EQ( + typename decltype(colors)::corpus_type{0}, + std::get<decltype(domain)::kInputCorpusValsOffset>(value.corpus_value)); EXPECT_EQ("Blue", value.user_value); } @@ -226,10 +229,20 @@ TEST(FlatMap, SerializationRoundTrip) { auto domain = FlatMap([](int len) { return AsciiString().WithSize(len); }, InRange(0, 10)); + using FlatMapDomain = decltype(domain); 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<FlatMapDomain::kOutputCorpusValIdx>(*parsed), + std::get<FlatMapDomain::kOutputCorpusValIdx>(value.corpus_value)); + EXPECT_EQ( + std::get<FlatMapDomain::kInputCorpusValsOffset>(*parsed), + std::get<FlatMapDomain::kInputCorpusValsOffset>(value.corpus_value)); } TEST(FlatMap, ValidationRejectsInvalidValue) { @@ -257,16 +270,19 @@ TEST(FlatMap, MutationAcceptsChangingDomains) { auto domain = FlatMap([](int len) { return AsciiString().WithSize(len); }, InRange(0, 10)); + using FlatMapDomain = decltype(domain); 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<FlatMapDomain::kInputCorpusValsOffset>(value.corpus_value) == + std::get<FlatMapDomain::kInputCorpusValsOffset>(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<FlatMapDomain::kInputCorpusValsOffset>(mutated)); } TEST(FlatMap, MutationAcceptsShrinkingOutputDomains) { @@ -484,7 +500,9 @@ 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<decltype(domain)::kInputCorpusValsOffset>(value.corpus_value)); } TEST(ReversibleFlatMap, AcceptsMultipleInnerDomains) { @@ -544,10 +562,21 @@ return std::optional(std::tuple<int>(s.size())); }, InRange(0, 10)); + using ReversibleFlatMapDomain = decltype(domain); 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<ReversibleFlatMapDomain::kOutputCorpusValIdx>(*parsed), + std::get<ReversibleFlatMapDomain::kOutputCorpusValIdx>( + value.corpus_value)); + EXPECT_EQ(std::get<ReversibleFlatMapDomain::kInputCorpusValsOffset>(*parsed), + std::get<ReversibleFlatMapDomain::kInputCorpusValsOffset>( + value.corpus_value)); } TEST(ReversibleFlatMap, ParseCorpusRejectsInvalidInputValues) {
diff --git a/fuzztest/internal/domains/aggregate_of_impl.h b/fuzztest/internal/domains/aggregate_of_impl.h index fa03b4a..1a6d9ac 100644 --- a/fuzztest/internal/domains/aggregate_of_impl.h +++ b/fuzztest/internal/domains/aggregate_of_impl.h
@@ -83,7 +83,7 @@ status = std::move(res).status(); return false; } - std::get<I>(results) = *std::move(res); + std::get<I>(results).emplace(*std::move(res)); return true; }; if (!(init_one(Is) && ...)) {
diff --git a/fuzztest/internal/domains/flat_map_impl.h b/fuzztest/internal/domains/flat_map_impl.h index 37f0cb9..0246f2d 100644 --- a/fuzztest/internal/domains/flat_map_impl.h +++ b/fuzztest/internal/domains/flat_map_impl.h
@@ -30,6 +30,7 @@ #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" @@ -61,16 +62,21 @@ Derived, // The user value is the user value of the output domain. value_type_t<FlatMapOutputDomain<FlatMapper, InputDomain...>>, - // 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<FlatMapper, InputDomain...>, corpus_type_t<FlatMapOutputDomain<FlatMapper, InputDomain...>>, corpus_type_t<InputDomain>...>> { public: using typename FlatMapImplBase::DomainBase::corpus_type; using typename FlatMapImplBase::DomainBase::value_type; + static constexpr size_t kOutputDomainIdx = 0; + static constexpr size_t kOutputCorpusValIdx = 1; + static constexpr size_t kInputCorpusValsOffset = 2; + FlatMapImplBase() = default; explicit FlatMapImplBase(FlatMapper flat_mapper, InputDomain... input_domains) : flat_mapper_(std::move(flat_mapper)), @@ -78,14 +84,19 @@ 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) { + // Use `std::make_tuple` instead of CTAD (`std::tuple{...}`) to avoid + // calling the copy constructor when there is a single input domain + // whose corpus type is already a `std::tuple`. return std::make_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 +110,77 @@ bool mutate_inputs = !only_shrink && absl::Bernoulli(prng, 0.1); if (mutate_inputs) { ApplyIndex<kNumInputValues>([&](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<I>(input_domains_) - .Mutate(std::get<I + 1>(val), prng, metadata, only_shrink), + .Mutate(std::get<I + kInputCorpusValsOffset>(val), prng, metadata, + only_shrink), ...); }); - std::get<0>(val) = GetOutputDomain(val).Init(prng); + // Generate a new output domain and output corpus value. + // We can't write `std::get<...>(val) = ...` because there are domains + // and corpus types that don't support assignment. So we manually destroy + // the old objects and construct new ones in place. + // We must compute both before reconstructing either to ensure `val` + // is not left in an invalid state if `GetOutputDomain` or `Init` throws. + auto output_domain = GetOutputDomain(val); + auto output_corpus_val = output_domain.Init(prng); + ReconstructInPlace(std::get<kOutputDomainIdx>(val), + std::move(output_domain)); + ReconstructInPlace(std::get<kOutputCorpusValIdx>(val), + std::move(output_corpus_val)); 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<kOutputDomainIdx>(val).Mutate(std::get<kOutputCorpusValIdx>(val), + prng, metadata, only_shrink); } value_type GetValue(const corpus_type& v) const { - return GetOutputDomain(v).GetValue(std::get<0>(v)); + return std::get<kOutputDomainIdx>(v).GetValue( + std::get<kOutputCorpusValIdx>(v)); } - auto GetPrinter() const { - return FlatMappedPrinter<FlatMapper, InputDomain...>{flat_mapper_, - input_domains_}; - } + auto GetPrinter() const { return Printer{}; } std::optional<corpus_type> 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<kOutputDomainIdx>(v).SerializeCorpus( + std::get<kOutputCorpusValIdx>(v))); + + // 2. Serialize the input corpus values. + ApplyIndex<kNumInputValues>([&](auto... I) { + (subs.push_back( + std::get<I>(input_domains_) + .SerializeCorpus(std::get<I + kInputCorpusValsOffset>(v))), + ...); + }); + return obj; } absl::Status ValidateCorpusValue(const corpus_type& corpus_value) const { @@ -154,8 +188,8 @@ 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<kOutputDomainIdx>(corpus_value) + .ValidateCorpusValue(std::get<kOutputCorpusValIdx>(corpus_value)); } protected: @@ -164,8 +198,8 @@ } 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 <typename Tuple> FlatMapOutputDomain<FlatMapper, InputDomain...> GetOutputDomain( @@ -181,8 +215,8 @@ }); } - // 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 <typename Tuple> absl::Status ValidateInputValues(const Tuple& tuple) const { @@ -208,6 +242,18 @@ } private: + struct Printer { + 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<kOutputDomainIdx>(corpus_value), + std::get<kOutputCorpusValIdx>(corpus_value), out, mode); + } + }; + FlatMapper flat_mapper_; std::tuple<InputDomain...> input_domains_; }; @@ -263,40 +309,46 @@ std::optional<corpus_type> 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<ReversibleFlatMapImpl::FlatMapImplBase::kNumInputValues>( [&](auto... I) -> std::optional<std::tuple<corpus_type_t<InputDomain>...>> { - auto inner_corpus_vals = - std::tuple{std::get<I>(this->input_domains()) - .FromValue(std::get<I>(*input_values_opt))...}; + // Use `std::make_tuple` instead of CTAD (`std::tuple{...}`) to + // avoid calling the copy constructor when there is a single input + // domain whose corpus type is already a `std::tuple`. + auto inner_corpus_vals = std::make_tuple( + std::get<I>(this->input_domains()) + .FromValue(std::get<I>(*input_user_vals))...); bool has_nullopt = (!std::get<I>(inner_corpus_vals).has_value() || ...); if (has_nullopt) return std::nullopt; - return std::tuple{*std::move(std::get<I>(inner_corpus_vals))...}; + return std::make_tuple( + *std::move(std::get<I>(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/domains/variant_of_impl.h b/fuzztest/internal/domains/variant_of_impl.h index 25b7267..0d0bee7 100644 --- a/fuzztest/internal/domains/variant_of_impl.h +++ b/fuzztest/internal/domains/variant_of_impl.h
@@ -64,7 +64,7 @@ // mutating case in order to explore more on a given type before we start // from scratch again. if (absl::Bernoulli(prng, 0.2)) { - val = Init(prng); + ReconstructInPlace(val, Init(prng)); } else { Switch<sizeof...(Inner)>(val.index(), [&](auto I) { std::get<I>(inner_).Mutate(std::get<I>(val), prng, metadata,
diff --git a/fuzztest/internal/type_support.h b/fuzztest/internal/type_support.h index 0273b7e..543b9c3 100644 --- a/fuzztest/internal/type_support.h +++ b/fuzztest/internal/type_support.h
@@ -20,6 +20,8 @@ #include <cstddef> #include <cstdint> #include <limits> +#include <memory> +#include <new> #include <string> #include <string_view> #include <tuple> @@ -52,6 +54,15 @@ namespace fuzztest::internal { +// Destroys `target` and constructs a new object of type `T` in place with +// `args`. This is useful for re-constructing objects in place when `T` is not +// assignable (e.g. tuples or variants containing non-assignable types). +template <typename T, typename... Args> +void ReconstructInPlace(T& target, Args&&... args) { + std::destroy_at(&target); + ::new (static_cast<void*>(&target)) T(std::forward<Args>(args)...); +} + // Return a best effort printer for type `T`. // This is useful for cases where the domain can't figure out how to print the // value. @@ -525,27 +536,6 @@ } }; -template <typename FlatMapper, typename... Inner> -struct FlatMappedPrinter { - const FlatMapper& mapper; - const std::tuple<Inner...>& inner; - - template <typename CorpusT> - void PrintCorpusValue(const CorpusT& corpus_value, - domain_implementor::RawSink out, - domain_implementor::PrintMode mode) const { - auto output_domain = ApplyIndex<sizeof...(Inner)>([&](auto... I) { - return mapper( - // the first field of `corpus_value` is the output value, so skip it - std::get<I>(inner).GetValue(std::get<I + 1>(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 8669236..5b123b1 100644 --- a/fuzztest/internal/type_support_test.cc +++ b/fuzztest/internal/type_support_test.cc
@@ -482,36 +482,49 @@ }; auto input_domain = InRange(1, 3); auto flat_map_domain = FlatMap(optional_sized_strings, input_domain); + using FlatMapDomain = decltype(flat_map_domain); - corpus_type_t<decltype(flat_map_domain)> abc_corpus_val = { + corpus_type_t<FlatMapDomain> abc_corpus_val = { + // Output domain + optional_sized_strings(3), // String of size GenericDomainCorpusType(std::in_place_type<std::string>, "ABC"), // Size 3}; + // 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()); + input_domain + .ValidateCorpusValue( + std::get<FlatMapDomain::kInputCorpusValsOffset>(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)) + std::get<FlatMapDomain::kOutputDomainIdx>(abc_corpus_val) + .ValidateCorpusValue( + std::get<FlatMapDomain::kOutputCorpusValIdx>(abc_corpus_val)) .ok()); EXPECT_THAT(TestPrintValue(abc_corpus_val, flat_map_domain), ElementsAre("(\"ABC\")", "\"ABC\"")); - corpus_type_t<decltype(flat_map_domain)> nullopt_corpus_val = { - // Corpus value of nullopt - std::monostate{}, - // Size (here irrelevant) - 2}; + corpus_type_t<FlatMapDomain> nullopt_corpus_val = {// Output domain + optional_sized_strings(2), + // Corpus value of nullopt + std::monostate{}, + // Size (here irrelevant) + 2}; // 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)) - .ok()); + input_domain + .ValidateCorpusValue(std::get<FlatMapDomain::kInputCorpusValsOffset>( + nullopt_corpus_val)) + .ok()); + ASSERT_TRUE( + std::get<FlatMapDomain::kOutputDomainIdx>(nullopt_corpus_val) + .ValidateCorpusValue( + std::get<FlatMapDomain::kOutputCorpusValIdx>(nullopt_corpus_val)) + .ok()); EXPECT_THAT(TestPrintValue(nullopt_corpus_val, flat_map_domain), Each("std::nullopt")); }