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

#ifndef CHROME_BROWSER_AI_AI_TEST_UTILS_H_
#define CHROME_BROWSER_AI_AI_TEST_UTILS_H_

#include <cstdint>
#include <vector>

#include "base/functional/callback_forward.h"
#include "base/run_loop.h"
#include "base/supports_user_data.h"
#include "base/test/scoped_feature_list.h"
#include "build/build_config.h"
#include "chrome/browser/ai/ai_manager.h"
#include "chrome/browser/optimization_guide/mock_optimization_guide_keyed_service.h"
#include "chrome/test/base/chrome_render_view_host_test_harness.h"
#if BUILDFLAG(IS_ANDROID)
#include "components/optimization_guide/core/model_execution/test/fake_model_assets.h"
#include "components/optimization_guide/core/model_execution/test/fake_model_broker_android.h"
#else
#include "components/optimization_guide/core/model_execution/manifest_broker/test/fake_manifest_broker.h"
#include "components/optimization_guide/core/model_execution/manifest_broker/test/scenario_builder.h"
#endif
#include "components/optimization_guide/proto/manifest.pb.h"
#include "components/optimization_guide/proto/on_device_model_execution_config.pb.h"
#include "components/update_client/crx_update_item.h"
#include "mojo/public/cpp/bindings/receiver.h"
#include "mojo/public/cpp/bindings/remote.h"
#include "services/network/public/mojom/permissions_policy/permissions_policy_feature.mojom-forward.h"
#include "services/on_device_model/public/mojom/download_observer.mojom.h"
#include "testing/gmock/include/gmock/gmock.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/blink/public/mojom/ai/ai_common.mojom.h"
#include "third_party/blink/public/mojom/ai/ai_manager.mojom.h"
#include "third_party/blink/public/mojom/ai/model_streaming_responder.mojom.h"

class AITestUtils {
 public:
  class TestStreamingResponder : public blink::mojom::ModelStreamingResponder {
   public:
    TestStreamingResponder();
    ~TestStreamingResponder() override;

    mojo::PendingRemote<blink::mojom::ModelStreamingResponder> BindRemote();

    // Returns true on successful completion and false on error.
    bool WaitForCompletion();

    // Returns true after tool calls are received.
    bool WaitForToolCalls();

    void WaitForContextOverflow();

    blink::mojom::ModelStreamingResponseStatus error_status() const {
      EXPECT_TRUE(error_status_.has_value());
      return *error_status_;
    }

    blink::mojom::QuotaErrorInfo quota_error_info() const {
      return *quota_error_info_;
    }

    const std::vector<std::string> responses() const { return responses_; }
    const std::vector<std::string> responses_without_last() const {
      EXPECT_TRUE(responses_.size() > 1);
      EXPECT_EQ(responses_.back(), "");
      return std::vector<std::string>(responses_.begin(), responses_.end() - 1);
    }

    uint64_t current_tokens() const { return current_tokens_; }

    const std::vector<blink::mojom::ToolCallPtr>& tool_calls() const {
      return tool_calls_;
    }

   private:
    // blink::mojom::ModelStreamingResponder:
    void OnError(blink::mojom::ModelStreamingResponseStatus status,
                 blink::mojom::QuotaErrorInfoPtr quota_error_info) override;
    void OnStreaming(const std::string& text) override;
    void OnCompletion(
        blink::mojom::ModelExecutionContextInfoPtr context_info) override;
    void OnToolCalls(
        std::vector<blink::mojom::ToolCallPtr> tool_calls) override;
    void OnContextOverflow() override;

    std::optional<blink::mojom::ModelStreamingResponseStatus> error_status_;
    blink::mojom::QuotaErrorInfoPtr quota_error_info_;
    std::vector<std::string> responses_;
    std::vector<blink::mojom::ToolCallPtr> tool_calls_;
    uint64_t current_tokens_ = 0;
    base::RunLoop run_loop_;
    base::RunLoop tool_calls_run_loop_;
    base::RunLoop context_overflow_run_loop_;
    mojo::Receiver<blink::mojom::ModelStreamingResponder> receiver_{this};
  };

  class AITestBase : public ChromeRenderViewHostTestHarness {
   public:
    AITestBase();
    ~AITestBase() override;

    void SetUp() override;
    void TearDown() override;

   protected:
    virtual void SetupBroker();
    virtual void SetupMockOptimizationGuideKeyedService();
    virtual void SetupNullOptimizationGuideKeyedService();

    virtual optimization_guide::proto::SolutionConfig CreateSolution() = 0;

    void SetSolutionConfig(
        optimization_guide::proto::SolutionConfig solution_config);

    blink::mojom::AIManager* GetAIManagerInterface();
    mojo::Remote<blink::mojom::AIManager> GetAIManagerRemote();
    size_t GetAIManagerContextBoundObjectSetSize();

    // Navigates to disable the specified policy and recreates `ai_manager_`.
    void DisablePolicy(network::mojom::PermissionsPolicyFeature feature);

    void InstallBaseModel();
    void UnInstallBaseModel();
    void SetSizeInTokens(uint32_t size);
    void SetExecuteResult(const std::vector<std::string>& result);

    // Helpers to set enterprise policies and user settings for testing.
    void SetBuiltInAIAPIsEnterprisePolicy(bool allowed);
    void SetGenAILocalEnterprisePolicy(bool allowed);
    void SetOnDeviceAiUserSetting(bool allowed);

    raw_ptr<MockOptimizationGuideKeyedService>
        mock_optimization_guide_keyed_service_;

#if BUILDFLAG(IS_ANDROID)
    base::test::ScopedFeatureList scoped_feature_list_;
    std::unique_ptr<optimization_guide::FakeModelBrokerAndroid> fake_broker_;
    std::vector<std::unique_ptr<optimization_guide::FakeAdaptationAsset>>
        fake_assets_;
#else
    std::unique_ptr<optimization_guide::FakeManifestBroker> fake_broker_;
#endif

    std::unique_ptr<AIManager> ai_manager_;
  };

  // Converts string language codes to AILanguageCode mojo struct.
  static std::vector<blink::mojom::AILanguageCodePtr> ToMojoLanguageCodes(
      const std::vector<std::string>& language_codes);
};

#endif  // CHROME_BROWSER_AI_AI_TEST_UTILS_H_
