Fix a validation bug in `FlatMapImpl` and clean up the code.

The bug was that `ParseCorpus` wasn't checking validity of the input values
before getting the output domain. For invalid input values, this was leading to
undefined behavior.

PiperOrigin-RevId: 603402138
diff --git a/domain_tests/map_filter_combinator_test.cc b/domain_tests/map_filter_combinator_test.cc
index d8c16b7..cc80539 100644
--- a/domain_tests/map_filter_combinator_test.cc
+++ b/domain_tests/map_filter_combinator_test.cc
@@ -313,6 +313,18 @@
   EXPECT_THAT(MutateUntilFoundN(domain, /*n=*/2), UnorderedElementsAre(0, 1));
 }
 
+TEST(FlatMap, ParseCorpusRejectsInvalidInputValues) {
+  absl::BitGen bitgen;
+
+  auto domain_a = FlatMap([](int a) { return Just(a); }, InRange(0, 9));
+  auto domain_b = FlatMap([](int a) { return Just(a); }, InRange(10, 19));
+
+  Value value(domain_a, bitgen);
+  auto serialized = domain_a.SerializeCorpus(value.corpus_value);
+
+  EXPECT_EQ(domain_b.ParseCorpus(serialized), std::nullopt);
+}
+
 TEST(Filter, CanFilterInitCalls) {
   Domain<int> domain = Filter([](int i) { return i % 2 == 0; }, InRange(1, 10));
   absl::BitGen bitgen;
diff --git a/fuzztest/internal/domains/flat_map_impl.h b/fuzztest/internal/domains/flat_map_impl.h
index b6712e3..abeaa9e 100644
--- a/fuzztest/internal/domains/flat_map_impl.h
+++ b/fuzztest/internal/domains/flat_map_impl.h
@@ -15,15 +15,19 @@
 #ifndef FUZZTEST_FUZZTEST_INTERNAL_DOMAINS_FLAT_MAP_IMPL_H_
 #define FUZZTEST_FUZZTEST_INTERNAL_DOMAINS_FLAT_MAP_IMPL_H_
 
+#include <cstddef>
 #include <optional>
 #include <tuple>
 #include <type_traits>
 
 #include "absl/random/bit_gen_ref.h"
 #include "absl/random/distributions.h"
+#include "absl/status/status.h"
+#include "absl/strings/str_format.h"
 #include "absl/types/span.h"
 #include "./fuzztest/internal/domains/domain_base.h"
 #include "./fuzztest/internal/domains/serialization_helpers.h"
+#include "./fuzztest/internal/logging.h"
 #include "./fuzztest/internal/meta.h"
 #include "./fuzztest/internal/serialization.h"
 #include "./fuzztest/internal/status.h"
@@ -68,10 +72,7 @@
           return std::make_tuple(input_domains.Init(prng)...);
         },
         input_domains_);
-    auto output_domain = ApplyIndex<sizeof...(InputDomain)>([&](auto... I) {
-      return flat_mapper_(
-          std::get<I>(input_domains_).GetValue(std::get<I>(input_corpus))...);
-    });
+    auto output_domain = GetOutputDomain(input_corpus);
     return std::tuple_cat(std::make_tuple(output_domain.Init(prng)),
                           input_corpus);
   }
@@ -84,7 +85,7 @@
     // re-initializing would lose the "still crashing" output value.
     bool mutate_inputs = !only_shrink && absl::Bernoulli(prng, 0.1);
     if (mutate_inputs) {
-      ApplyIndex<sizeof...(InputDomain)>([&](auto... I) {
+      ApplyIndex<kNumInputValues>([&](auto... I) {
         // The first field of `val` is the output corpus value, so skip it.
         (std::get<I>(input_domains_)
              .Mutate(std::get<I + 1>(val), prng, only_shrink),
@@ -121,10 +122,12 @@
     if (!input_corpus.has_value()) {
       return std::nullopt;
     }
-    auto output_domain = ApplyIndex<sizeof...(InputDomain)>([&](auto... I) {
-      return flat_mapper_(
-          std::get<I>(input_domains_).GetValue(std::get<I>(*input_corpus))...);
-    });
+    absl::Status input_values_validity = ValidateInputValues(*input_corpus);
+    if (!input_values_validity.ok()) {
+      absl::FPrintF(GetStderr(), "[!] %s", input_values_validity.message());
+      return std::nullopt;
+    }
+    auto output_domain = GetOutputDomain(*input_corpus);
     // We know obj.Subs()[0] exists because ParseWithDomainTuple succeeded.
     auto output_corpus = output_domain.ParseCorpus((*obj.Subs())[0]);
     if (!output_corpus.has_value()) {
@@ -141,19 +144,7 @@
 
   absl::Status ValidateCorpusValue(const corpus_type& corpus_value) const {
     // Check input values first.
-    absl::Status input_values_validity = absl::OkStatus();
-    ApplyIndex<sizeof...(InputDomain)>([&](auto... I) {
-      (
-          [&] {
-            if (!input_values_validity.ok()) return;
-            const absl::Status s =
-                std::get<I>(input_domains_)
-                    .ValidateCorpusValue(std::get<I + 1>(corpus_value));
-            input_values_validity =
-                Prefix(s, "Invalid value for FlatMap()-ed domain");
-          }(),
-          ...);
-    });
+    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)
@@ -161,16 +152,49 @@
   }
 
  private:
-  using output_domain_t = std::decay_t<
-      std::invoke_result_t<FlatMapper, value_type_t<InputDomain>...>>;
-  output_domain_t GetOutputDomain(const corpus_type& val) const {
-    return ApplyIndex<sizeof...(InputDomain)>([&](auto... I) {
-      // The first field of `val` is the output corpus value, so skip it.
-      return flat_mapper_(
-          std::get<I>(input_domains_).GetValue(std::get<I + 1>(val))...);
+  // 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
+  // `kNumInputValues` elements.
+  template <typename Tuple>
+  FlatMapOutputDomain<FlatMapper, InputDomain...> GetOutputDomain(
+      const Tuple& tuple) const {
+    static_assert(is_tuple_v<Tuple> &&
+                  std::tuple_size_v<Tuple> >= kNumInputValues);
+    static constexpr size_t kOffset =
+        std::tuple_size_v<Tuple> - kNumInputValues;
+    return ApplyIndex<kNumInputValues>([&](auto... I) {
+      // The first field of `tuple` may be the output corpus value, so skip it.
+      return flat_mapper_(std::get<I>(input_domains_)
+                              .GetValue(std::get<kOffset + I>(tuple))...);
     });
   }
 
+  // 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
+  // `kNumInputValues` elements.
+  template <typename Tuple>
+  absl::Status ValidateInputValues(const Tuple& tuple) const {
+    static_assert(is_tuple_v<Tuple> &&
+                  std::tuple_size_v<Tuple> >= kNumInputValues);
+    static constexpr size_t kOffset =
+        std::tuple_size_v<Tuple> - kNumInputValues;
+    return ApplyIndex<kNumInputValues>([&](auto... I) {
+      absl::Status input_values_validity = absl::OkStatus();
+      (
+          [&] {
+            if (!input_values_validity.ok()) return;
+            const absl::Status s =
+                std::get<I>(input_domains_)
+                    .ValidateCorpusValue(std::get<kOffset + I>(tuple));
+            input_values_validity =
+                Prefix(s, "Invalid value for FlatMap()-ed domain");
+          }(),
+          ...);
+      return input_values_validity;
+    });
+  }
+
+  static constexpr size_t kNumInputValues = sizeof...(InputDomain);
   FlatMapper flat_mapper_;
   std::tuple<InputDomain...> input_domains_;
 };