Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 19 additions & 7 deletions domain_tests/map_filter_combinator_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -191,7 +191,7 @@ TEST(FlatMap, WorksWithSameCorpusType) {
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<2>(value.corpus_value));
}

TEST(FlatMap, WorksWithDifferentCorpusType) {
Expand All @@ -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);
}

Expand All @@ -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) {
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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) {
Expand Down
144 changes: 89 additions & 55 deletions fuzztest/internal/domains/flat_map_impl.h
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
#define FUZZTEST_FUZZTEST_INTERNAL_DOMAINS_FLAT_MAP_IMPL_H_

#include <cstddef>
#include <memory>
#include <new>
#include <optional>
#include <tuple>
#include <type_traits>
Expand All @@ -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 {

Expand Down Expand Up @@ -61,10 +63,11 @@ class FlatMapImplBase
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:
Expand All @@ -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,
Expand All @@ -99,63 +104,77 @@ class FlatMapImplBase
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 + 2>(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<void*>(&std::get<0>(val)))
FlatMapOutputDomain<FlatMapper, InputDomain...>(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<FlatMapper, InputDomain...>{flat_mapper_,
input_domains_};
}
auto GetPrinter() const { return Printer{input_domains_}; }

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<0>(v).SerializeCorpus(std::get<1>(v)));

// 2. Serialize the input corpus values.
ApplyIndex<kNumInputValues>([&](auto... I) {
(subs.push_back(
std::get<I>(input_domains_).SerializeCorpus(std::get<I + 2>(v))),
...);
});
return obj;
}

absl::Status ValidateCorpusValue(const corpus_type& corpus_value) const {
// Check input values first.
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:
Expand All @@ -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 <typename Tuple>
FlatMapOutputDomain<FlatMapper, InputDomain...> GetOutputDomain(
Expand All @@ -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 <typename Tuple>
absl::Status ValidateInputValues(const Tuple& tuple) const {
Expand All @@ -208,6 +227,19 @@ class FlatMapImplBase
}

private:
struct Printer {
const std::tuple<InputDomain...>& 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<InputDomain...> input_domains_;
};
Expand Down Expand Up @@ -263,40 +295,42 @@ class ReversibleFlatMapImpl

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))...};
.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))...};
});
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:
Expand Down
21 changes: 0 additions & 21 deletions fuzztest/internal/type_support.h
Original file line number Diff line number Diff line change
Expand Up @@ -525,27 +525,6 @@ struct MappedPrinter {
}
};

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,
Expand Down
Loading
Loading