// Copyright 2026 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include <stdint.h>

#include <memory>
#include <utility>

#include "base/command_line.h"
#include "base/functional/bind.h"
#include "base/run_loop.h"
#include "base/task/single_thread_task_runner.h"
#include "base/test/task_environment.h"
#include "base/test/test_timeouts.h"
#include "components/viz/service/layers/layer_context_impl.h"
#include "components/viz/service/layers/layer_context_impl_base_unittest.h"
#include "components/viz/service/layers/layer_context_impl_mojolpm_fuzzer.pb.h"
#include "mojo/core/embedder/embedder.h"
#include "mojo/public/tools/fuzzers/mojolpm.h"
#include "services/viz/public/mojom/compositing/layer_context.mojom-mojolpm.h"
#include "services/viz/public/mojom/compositing/tiling.mojom-mojolpm.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/libprotobuf-mutator/src/src/libfuzzer/libfuzzer_macro.h"

namespace {

struct FuzzerEnvironment {
  FuzzerEnvironment(int* argc, char** argv) {
    base::CommandLine::Init(*argc, argv);
    TestTimeouts::Initialize();
    testing::InitGoogleTest(argc, argv);
    mojo::core::Init();
  }
  base::test::TaskEnvironment task_environment{
      base::test::TaskEnvironment::MainThreadType::DEFAULT,
      base::test::TaskEnvironment::TimeSource::MOCK_TIME};
};

scoped_refptr<base::SequencedTaskRunner> GetFuzzerTaskRunner() {
  return base::SingleThreadTaskRunner::GetCurrentDefault();
}

class LayerContextImplTestForFuzzing : public viz::LayerContextImplTest {
 public:
  void TestBody() override {}
  viz::LayerContextImpl* layer_context_impl() {
    return layer_context_impl_.get();
  }
};

void FixUpSharedImageFormat(mojolpm::viz::mojom::SharedImageFormat* format,
                            int index) {
  if (format->instance_case() == mojolpm::viz::mojom::SharedImageFormat::kOld) {
    return;
  }
  auto* union_ptr = format->mutable_new_();
  if (!union_ptr->has_id()) {
    union_ptr->set_id(0);
  }

  bool needs_fixup = true;
  if (union_ptr->union_member_case() ==
      mojolpm::viz::mojom::SharedImageFormat_ProtoUnion::kMSingleplanarFormat) {
    auto single_format = union_ptr->m_singleplanar_format();
    if (single_format ==
            mojolpm::viz::mojom::MojoLPM_SingleplanarFormat_RGBA_8888 ||
        single_format ==
            mojolpm::viz::mojom::MojoLPM_SingleplanarFormat_BGRA_8888) {
      needs_fixup = false;
    }
  }

  if (needs_fixup) {
    if (index % 2 == 0) {
      union_ptr->set_m_singleplanar_format(
          mojolpm::viz::mojom::MojoLPM_SingleplanarFormat_RGBA_8888);
    } else {
      union_ptr->set_m_singleplanar_format(
          mojolpm::viz::mojom::MojoLPM_SingleplanarFormat_BGRA_8888);
    }
  }
}

void FixUpLayerTreeUpdate(mojolpm::viz::mojom::LayerTreeUpdate* update) {
  // If update is based on an old instance, it's unnecessary to
  // fix up the LayerTreeUpdate proto since it's already been done.
  if (update->instance_case() == mojolpm::viz::mojom::LayerTreeUpdate::kOld) {
    return;
  }
  // Default construct a LayerTreeUpdate mojo struct and convert it to a proto
  // to get a valid proto with all fields initialized to their default values.
  // Then merge the generated proto into the default proto to copy all
  // initialized fields from the generated proto. The resulting proto contains
  // all valid fields and retains the fields generated by the fuzzing engine.
  mojolpm::viz::mojom::LayerTreeUpdate default_update_proto;
  if (mojolpm::ToProto(viz::mojom::LayerTreeUpdate::New(),
                       default_update_proto)) {
    default_update_proto.MergeFrom(*update);
    *update = std::move(default_update_proto);
  }

  auto* new_update = update->mutable_new_();

  // If display_color_spaces is old, skip fixing up SharedImageFormats.
  if (new_update->mutable_m_display_color_spaces()->instance_case() ==
      mojolpm::gfx::mojom::DisplayColorSpaces::kOld) {
    return;
  }
  auto* formats = new_update->mutable_m_display_color_spaces()
                      ->mutable_new_()
                      ->mutable_m_formats();
  int i = 0;
  for (auto* format :
       {formats->mutable_value_0(), formats->mutable_value_1(),
        formats->mutable_value_2(), formats->mutable_value_3(),
        formats->mutable_value_4(), formats->mutable_value_5()}) {
    FixUpSharedImageFormat(format->mutable_value(), i++);
  }
}

constexpr uint32_t kMaxPropertyTreeNodes = 1000u;
constexpr size_t kMaxLayers = 1000;
constexpr size_t kMaxTilings = 100;
constexpr size_t kMaxTilesPerTiling = 100;
constexpr size_t kMaxUIResourceRequests = 100;
constexpr size_t kMaxSurfaceRanges = 100;
constexpr size_t kMaxLatencyInfo = 100;
constexpr size_t kMaxViewTransitionRequests = 100;
constexpr size_t kMaxTrackedElementRects = 100;
constexpr size_t kMaxCopyOutputRequests = 100;
constexpr size_t kMaxAnimationTimelines = 100;
constexpr size_t kMaxAnimationsPerTimeline = 100;
constexpr size_t kMaxKeyframeModelsPerAnimation = 10;
constexpr size_t kMaxKeyframesPerAnimationCurve = 10;

void ClampTiling(viz::mojom::Tiling* tiling) {
  if (!tiling) {
    return;
  }
  if (tiling->tiles.size() > kMaxTilesPerTiling) {
    tiling->tiles.resize(kMaxTilesPerTiling);
  }
}

void ClampLayerTreeUpdate(viz::mojom::LayerTreeUpdate* update) {
  if (!update) {
    return;
  }

  update->num_transform_nodes =
      std::min(update->num_transform_nodes, kMaxPropertyTreeNodes);
  update->num_clip_nodes =
      std::min(update->num_clip_nodes, kMaxPropertyTreeNodes);
  update->num_effect_nodes =
      std::min(update->num_effect_nodes, kMaxPropertyTreeNodes);
  update->num_scroll_nodes =
      std::min(update->num_scroll_nodes, kMaxPropertyTreeNodes);

  if (update->transform_nodes.size() > kMaxPropertyTreeNodes) {
    update->transform_nodes.resize(kMaxPropertyTreeNodes);
  }
  if (update->clip_nodes.size() > kMaxPropertyTreeNodes) {
    update->clip_nodes.resize(kMaxPropertyTreeNodes);
  }
  if (update->effect_nodes.size() > kMaxPropertyTreeNodes) {
    update->effect_nodes.resize(kMaxPropertyTreeNodes);
  }
  if (update->scroll_nodes.size() > kMaxPropertyTreeNodes) {
    update->scroll_nodes.resize(kMaxPropertyTreeNodes);
  }

  if (update->layers.size() > kMaxLayers) {
    update->layers.resize(kMaxLayers);
  }
  if (update->layer_order && update->layer_order->size() > kMaxLayers) {
    update->layer_order->resize(kMaxLayers);
  }

  if (update->tilings.size() > kMaxTilings) {
    update->tilings.resize(kMaxTilings);
  }
  for (auto& tiling : update->tilings) {
    ClampTiling(tiling.get());
  }

  if (update->ui_resource_requests.size() > kMaxUIResourceRequests) {
    update->ui_resource_requests.resize(kMaxUIResourceRequests);
  }
  if (update->surface_ranges &&
      update->surface_ranges->size() > kMaxSurfaceRanges) {
    update->surface_ranges->resize(kMaxSurfaceRanges);
  }
  if (update->latency_info.size() > kMaxLatencyInfo) {
    update->latency_info.resize(kMaxLatencyInfo);
  }

  if (update->view_transition_requests &&
      update->view_transition_requests->size() > kMaxViewTransitionRequests) {
    update->view_transition_requests->resize(kMaxViewTransitionRequests);
  }

  if (update->tracked_element_rects.size() > kMaxTrackedElementRects) {
    auto iter = update->tracked_element_rects.begin();
    std::advance(iter, kMaxTrackedElementRects);
    update->tracked_element_rects.erase(iter,
                                        update->tracked_element_rects.end());
  }
  for (auto& [feature, rects] : update->tracked_element_rects) {
    if (rects.size() > kMaxTrackedElementRects) {
      rects.resize(kMaxTrackedElementRects);
    }
  }

  for (auto& node : update->effect_nodes) {
    if (node->copy_output_requests.size() > kMaxCopyOutputRequests) {
      node->copy_output_requests.resize(kMaxCopyOutputRequests);
    }
  }

  if (update->removed_animation_timelines &&
      update->removed_animation_timelines->size() > kMaxAnimationTimelines) {
    update->removed_animation_timelines->resize(kMaxAnimationTimelines);
  }
  if (update->animation_timelines) {
    if (update->animation_timelines->size() > kMaxAnimationTimelines) {
      update->animation_timelines->resize(kMaxAnimationTimelines);
    }
    for (auto& timeline : *update->animation_timelines) {
      if (timeline->removed_animations.size() > kMaxAnimationsPerTimeline) {
        timeline->removed_animations.resize(kMaxAnimationsPerTimeline);
      }
      if (timeline->new_animations.size() > kMaxAnimationsPerTimeline) {
        timeline->new_animations.resize(kMaxAnimationsPerTimeline);
      }
      for (auto& animation : timeline->new_animations) {
        if (animation->keyframe_models.size() >
            kMaxKeyframeModelsPerAnimation) {
          animation->keyframe_models.resize(kMaxKeyframeModelsPerAnimation);
        }
        for (auto& model : animation->keyframe_models) {
          if (model->keyframes.size() > kMaxKeyframesPerAnimationCurve) {
            model->keyframes.resize(kMaxKeyframesPerAnimationCurve);
          }
        }
      }
    }
  }
}

class LayerContextTestcase
    : public mojolpm::Testcase<viz::fuzzing::layer_context::proto::Testcase,
                               viz::fuzzing::layer_context::proto::Action> {
 public:
  using ProtoTestcase = viz::fuzzing::layer_context::proto::Testcase;
  using ProtoAction = viz::fuzzing::layer_context::proto::Action;

  explicit LayerContextTestcase(const ProtoTestcase& testcase)
      : mojolpm::Testcase<ProtoTestcase, ProtoAction>(testcase) {}

  ~LayerContextTestcase() = default;

  void SetUp(base::OnceClosure done_closure) override {
    impl_test_ = std::make_unique<LayerContextImplTestForFuzzing>();
    impl_test_->SetUp();

    // Give it a valid initial state just like the unit tests do.
    auto default_update = impl_test_->CreateDefaultUpdate();
    (void)impl_test_->layer_context_impl()->DoUpdateDisplayTree(
        std::move(default_update));

    GetFuzzerTaskRunner()->PostTask(FROM_HERE, std::move(done_closure));
  }

  void TearDown(base::OnceClosure done_closure) override {
    impl_test_.reset();
    GetFuzzerTaskRunner()->PostTask(FROM_HERE, std::move(done_closure));
  }

  void RunAction(const ProtoAction& action,
                 base::OnceClosure run_closure) override {
    switch (action.action_case()) {
      case ProtoAction::kRunThread:
        base::SingleThreadTaskRunner::GetCurrentDefault()->PostTaskAndReply(
            FROM_HERE, base::DoNothing(), std::move(run_closure));
        return;

      case ProtoAction::kUpdateDisplayTree: {
        viz::mojom::LayerTreeUpdatePtr update;
        auto proto_update = action.update_display_tree();
        FixUpLayerTreeUpdate(&proto_update);
        if (mojolpm::FromProto(proto_update, update) && update) {
          ClampLayerTreeUpdate(update.get());
          // Call the implementation directly, bypassing the Mojo pipe to avoid
          // ReportBadMessage disconnecting the endpoint on invalid fuzz data.
          (void)impl_test_->layer_context_impl()->DoUpdateDisplayTree(
              std::move(update));
        }
        break;
      }

      case ProtoAction::kUpdateDisplayTiling: {
        viz::mojom::TilingPtr tiling;
        if (mojolpm::FromProto(action.update_display_tiling(), tiling) &&
            tiling) {
          ClampTiling(tiling.get());
          (void)impl_test_->layer_context_impl()->DoUpdateDisplayTiling(
              std::move(tiling));
        }
        break;
      }

      case ProtoAction::ACTION_NOT_SET:
        break;
    }

    GetFuzzerTaskRunner()->PostTask(FROM_HERE, std::move(run_closure));
  }

 private:
  std::unique_ptr<LayerContextImplTestForFuzzing> impl_test_;
};

}  // namespace

DEFINE_BINARY_PROTO_FUZZER(
    const viz::fuzzing::layer_context::proto::Testcase& testcase) {
  if (!testcase.actions_size() && !testcase.sequences_size()) {
    return;
  }

  int argc = 1;
  const char* argv_array[] = {"layer_context_impl_mojolpm_fuzzer"};
  char** argv = const_cast<char**>(argv_array);
  static FuzzerEnvironment env(&argc, argv);

  LayerContextTestcase testcase_runner(testcase);

  base::RunLoop main_run_loop;
  GetFuzzerTaskRunner()->PostTask(
      FROM_HERE,
      base::BindOnce(&mojolpm::RunTestcase<LayerContextTestcase>,
                     base::Unretained(&testcase_runner), GetFuzzerTaskRunner(),
                     main_run_loop.QuitClosure()));
  main_run_loop.Run();
}
