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

#include "chrome/browser/ai/ai_proofreader.h"

#include <memory>

#include "base/run_loop.h"
#include "base/test/bind.h"
#include "base/test/gmock_expected_support.h"
#include "base/test/metrics/histogram_tester.h"
#include "base/test/mock_callback.h"
#include "base/test/protobuf_matchers.h"
#include "base/test/scoped_feature_list.h"
#include "base/test/task_environment.h"
#include "base/test/test_future.h"
#include "base/types/expected.h"
#include "chrome/browser/ai/ai_test_utils.h"
#include "chrome/browser/ai/features.h"
#include "chrome/browser/optimization_guide/mock_optimization_guide_keyed_service.h"
#include "components/optimization_guide/core/model_execution/test/mock_on_device_capability.h"
#include "components/optimization_guide/core/model_execution/test/substitution_builder.h"
#include "components/optimization_guide/core/optimization_guide_proto_util.h"
#include "components/optimization_guide/core/optimization_guide_switches.h"
#include "components/optimization_guide/core/optimization_guide_util.h"
#include "components/optimization_guide/proto/features/proofreader_api.pb.h"
#include "components/optimization_guide/proto/string_value.pb.h"
#include "content/public/browser/web_contents.h"
#include "mojo/public/cpp/test_support/test_utils.h"
#include "services/on_device_model/public/cpp/features.h"
#include "testing/gmock/include/gmock/gmock.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/blink/public/common/features_generated.h"
#include "third_party/blink/public/mojom/ai/ai_manager.mojom.h"
#include "third_party/blink/public/mojom/ai/model_streaming_responder.mojom.h"

namespace {

namespace proto = ::optimization_guide::proto;

using ::base::test::ErrorIs;
using ::base::test::TestFuture;
using ::blink::mojom::AILanguageCode;
using ::blink::mojom::AILanguageCodePtr;
using ::on_device_model::mojom::PerformanceClass;
using ::optimization_guide::FieldSubstitution;
using ::optimization_guide::ForbidUnsafe;
using ::optimization_guide::ProtoField;
using ::optimization_guide::StringValueField;
using ::optimization_guide::proto::ProofreaderApiRequest;
using ::optimization_guide::proto::ProofreaderApiResponse;
using ::testing::_;
using ::testing::ElementsAre;
using ::testing::ElementsAreArray;

constexpr char kInputString[] = "input string";
constexpr char kCorrections[] = "[\"From `input` to `Input`\"]";

using CreateProofreaderResult =
    base::expected<mojo::PendingRemote<blink::mojom::AIProofreader>,
                   blink::mojom::AIManagerCreateClientError>;

class TestCreateProofreaderClient
    : public blink::mojom::AIManagerCreateProofreaderClient {
 public:
  TestCreateProofreaderClient() = default;
  ~TestCreateProofreaderClient() override = default;
  TestCreateProofreaderClient(const TestCreateProofreaderClient&) = delete;
  TestCreateProofreaderClient& operator=(const TestCreateProofreaderClient&) =
      delete;

  mojo::PendingRemote<blink::mojom::AIManagerCreateProofreaderClient>
  BindNewPipeAndPassRemote() {
    return receiver_.BindNewPipeAndPassRemote();
  }

  void OnResult(
      mojo::PendingRemote<::blink::mojom::AIProofreader> proofreader) override {
    result_.SetValue(std::move(proofreader));
  }

  void OnError(blink::mojom::AIManagerCreateClientError error,
               blink::mojom::QuotaErrorInfoPtr quota_error_info) override {
    result_.SetValue(base::unexpected(error));
  }

  TestFuture<CreateProofreaderResult>& result() { return result_; }

 private:
  TestFuture<CreateProofreaderResult> result_;
  mojo::Receiver<blink::mojom::AIManagerCreateProofreaderClient> receiver_{
      this};
};

blink::mojom::AIProofreaderCreateOptionsPtr GetDefaultOptions() {
  return blink::mojom::AIProofreaderCreateOptions::New(
      /*include_correction_types=*/false,
      /*include_correction_explanations=*/false,
      /*correction_explanation_language=*/AILanguageCode::New(""),
      /*expected_input_languages=*/std::vector<AILanguageCodePtr>());
}

optimization_guide::proto::FeatureTextSafetyConfiguration CreateSafetyConfig() {
  optimization_guide::proto::FeatureTextSafetyConfiguration safety_config;
  safety_config.set_feature(
      optimization_guide::proto::MODEL_EXECUTION_FEATURE_PROOFREADER_API);
  safety_config.mutable_safety_category_thresholds()->Add(ForbidUnsafe());

  {
    auto* check = safety_config.add_request_check();
    check->mutable_input_template()->Add(FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kTextFieldNumber})));
  }
  {
    auto* check = safety_config.add_request_check();
    check->mutable_input_template()->Add(FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kCorrectedTextFieldNumber})));
  }
  {
    auto* check = safety_config.add_request_check();
    check->mutable_input_template()->Add(FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kCorrectionFieldNumber})));
  }

  return safety_config;
}

class AIProofreaderTest : public AITestUtils::AITestBase {
 public:
  AIProofreaderTest() {
    scoped_feature_list_.InitAndEnableFeature(
        blink::features::kAIProofreadingAPI);
  }

 protected:
  proto::SolutionConfig CreateSolution() override {
    proto::OnDeviceModelExecutionFeatureConfig config;
    config.set_can_skip_text_safety(true);
    config.set_feature(proto::ModelExecutionFeature::
                           MODEL_EXECUTION_FEATURE_PROOFREADER_API);

    auto& input_config = *config.mutable_input_config();
    input_config.set_request_base_name(ProofreaderApiRequest().GetTypeName());

    *input_config.add_execute_substitutions() = FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kTextFieldNumber}));
    *input_config.add_execute_substitutions() = FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kCorrectedTextFieldNumber}));
    *input_config.add_execute_substitutions() = FieldSubstitution(
        "%s", ProtoField({ProofreaderApiRequest::kCorrectionFieldNumber}));

    auto& output_config = *config.mutable_output_config();
    output_config.set_proto_type(ProofreaderApiResponse().GetTypeName());
    *output_config.mutable_proto_field() = StringValueField();

    proto::SolutionConfig solution_config;
    *solution_config.mutable_feature() = config;
    *solution_config.mutable_safety() = CreateSafetyConfig();
    return solution_config;
  }

  mojo::Remote<blink::mojom::AIProofreader> GetAIProofreaderRemote(
      blink::mojom::AIProofreaderCreateOptionsPtr options =
          GetDefaultOptions()) {
    TestCreateProofreaderClient create_proofreader_client;
    GetAIManagerRemote()->CreateProofreader(
        create_proofreader_client.BindNewPipeAndPassRemote(),
        std::move(options), mojo::NullRemote());

    CreateProofreaderResult result = create_proofreader_client.result().Take();
    EXPECT_OK(result);
    return mojo::Remote<blink::mojom::AIProofreader>(std::move(result.value()));
  }

  void RunSimpleProofreadTest(bool include_correction_types,
                              bool include_correction_explanations) {
    fake_broker_->settings().set_execute_result({"Result text"});

    const auto options = blink::mojom::AIProofreaderCreateOptions::New(
        include_correction_types, include_correction_explanations,
        /*correction_explanation_language=*/AILanguageCode::New(""),
        /*expected_input_languages=*/std::vector<AILanguageCodePtr>());

    mojo::Remote<blink::mojom::AIProofreader> proofreader_remote =
        GetAIProofreaderRemote(options.Clone());

    EXPECT_THAT(Proofread(*proofreader_remote, kInputString),
                ElementsAreArray({"Result text"}));
  }

  std::vector<std::string> Proofread(blink::mojom::AIProofreader& proofreader,
                                     const std::string& input) {
    AITestUtils::TestStreamingResponder responder;
    proofreader.Proofread(kInputString, responder.BindRemote());
    EXPECT_TRUE(responder.WaitForCompletion());
    // Return Proofreader's response without the final empty string chunk.
    return responder.responses_without_last();
  }

  void EnsureModelIsReady() {
    TestCreateProofreaderClient proofreader_client;
    GetAIManagerRemote()->CreateProofreader(
        proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
        mojo::NullRemote());

    auto result = proofreader_client.result().Take();
    EXPECT_OK(result);
  }

 private:
  base::test::ScopedFeatureList scoped_feature_list_;
};

TEST_F(AIProofreaderTest, CreateProofreaderNoService) {
  SetupNullOptimizationGuideKeyedService();

  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());

  EXPECT_THAT(
      create_proofreader_client.result().Take(),
      ErrorIs(
          blink::mojom::AIManagerCreateClientError::kUnableToCreateSession));
}

TEST_F(AIProofreaderTest, ProofreaderTelemetry) {
  base::HistogramTester histogram_tester;
  EXPECT_CALL(*mock_optimization_guide_keyed_service_,
              GetOnDeviceModelEligibility(
                  optimization_guide::mojom::OnDeviceFeature::kProofreaderApi))
      .WillRepeatedly(testing::Return(
          optimization_guide::OnDeviceModelEligibilityReason::kSuccess));
  EnsureModelIsReady();
  GetAIProofreaderRemote();

  histogram_tester.ExpectUniqueSample(
      "OptimizationGuide.ModelExecution."
      "OnDeviceModelEligibilityReason.ProofreaderApi",
      optimization_guide::OnDeviceModelEligibilityReason::kSuccess, 2);
}

TEST_F(AIProofreaderTest, CanCreateDefaultOptions) {
  {
    base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
    GetAIManagerInterface()->CanCreateProofreader(GetDefaultOptions(),
                                                  future.GetCallback());
    EXPECT_EQ(future.Get(),
              blink::mojom::ModelAvailabilityCheckResult::kDownloadable);
  }

  // After model is ready, `CanCreateProofreader` should return available.
  EnsureModelIsReady();

  {
    base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
    GetAIManagerInterface()->CanCreateProofreader(GetDefaultOptions(),
                                                  future.GetCallback());
    EXPECT_EQ(future.Get(),
              blink::mojom::ModelAvailabilityCheckResult::kAvailable);
  }
}

TEST_F(AIProofreaderTest, CanCreateIsLanguagesSupported) {
  EnsureModelIsReady();

  auto options = GetDefaultOptions();
  options->correction_explanation_language = AILanguageCode::New("en");
  options->expected_input_languages =
      AITestUtils::ToMojoLanguageCodes({"en-US", ""});

  base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
  GetAIManagerInterface()->CanCreateProofreader(std::move(options),
                                                future.GetCallback());
  EXPECT_EQ(future.Get(),
            blink::mojom::ModelAvailabilityCheckResult::kAvailable);
}

TEST_F(AIProofreaderTest, CanCreateUnIsLanguagesSupported) {
  auto options = GetDefaultOptions();
  options->correction_explanation_language = AILanguageCode::New("es-ES");
  options->expected_input_languages =
      AITestUtils::ToMojoLanguageCodes({"en", "fr", "ja"});
  base::MockCallback<AIManager::CanCreateProofreaderCallback> callback;
  EXPECT_CALL(callback, Run(blink::mojom::ModelAvailabilityCheckResult::
                                kUnavailableUnsupportedLanguage));
  GetAIManagerInterface()->CanCreateProofreader(std::move(options),
                                                callback.Get());
}

TEST_F(AIProofreaderTest, CreateProofreaderModelNotEligible) {
  base::test::ScopedFeatureList feature_list;
  feature_list.InitWithFeaturesAndParameters(
      {{optimization_guide::features::kOnDeviceModelPerformanceParams,
        {{"compatible_on_device_performance_classes", "3,4,5,6"}}}},
      {{on_device_model::features::kOnDeviceModelCpuBackend}});

  fake_broker_->settings().performance_class =
      on_device_model::mojom::PerformanceClass::kVeryLow;

  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());

  EXPECT_THAT(
      create_proofreader_client.result().Take(),
      ErrorIs(
          blink::mojom::AIManagerCreateClientError::kUnableToCreateSession));
}

#if BUILDFLAG(IS_ANDROID)
TEST_F(AIProofreaderTest, CreateProofreaderSafetyConfigNotAvailable) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    solution_config.mutable_feature()->set_can_skip_text_safety(false);
    // Provide a safety asset that does not support proofreader.
    solution_config.mutable_safety()->set_feature(
        optimization_guide::proto::MODEL_EXECUTION_FEATURE_TEST);
    return solution_config;
  }());

  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());

  EXPECT_THAT(
      create_proofreader_client.result().Take(),
      ErrorIs(
          blink::mojom::AIManagerCreateClientError::kUnableToCreateSession));
}
#endif

TEST_F(AIProofreaderTest, ProofreadUnableToCalculateTokenSize) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    solution_config.mutable_feature()
        ->mutable_input_config()
        ->set_request_base_name("InvalidRequestBaseName");
    return solution_config;
  }());

  auto proofreader_remote = GetAIProofreaderRemote();
  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread(kInputString, responder.BindRemote());
  EXPECT_FALSE(responder.WaitForCompletion());
  EXPECT_EQ(
      responder.error_status(),
      blink::mojom::ModelStreamingResponseStatus::kErrorFailedToCountTokens);
}

TEST_F(AIProofreaderTest, ProofreadDefault) {
  RunSimpleProofreadTest(false, false);
}

TEST_F(AIProofreaderTest, ProofreadWithOptions) {
  bool types[]{false, true};
  bool explanations[]{false, true};
  for (const auto& include_correction_types : types) {
    for (const auto& include_correction_explanations : explanations) {
      SCOPED_TRACE(testing::Message() << include_correction_types << " "
                                      << include_correction_explanations);
      RunSimpleProofreadTest(include_correction_types,
                             include_correction_explanations);
    }
  }
}

TEST_F(AIProofreaderTest, InputLimitExceededError) {
  auto proofreader_remote = GetAIProofreaderRemote();

  fake_broker_->settings().set_size_in_tokens(
      blink::mojom::kTinyModelMaxInputTokenSize + 1);

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread(kInputString, responder.BindRemote());
  EXPECT_FALSE(responder.WaitForCompletion());
  EXPECT_EQ(responder.error_status(),
            blink::mojom::ModelStreamingResponseStatus::kErrorInputTooLarge);
  ASSERT_EQ(responder.quota_error_info().requested,
            blink::mojom::kTinyModelMaxInputTokenSize + 1);
  ASSERT_EQ(responder.quota_error_info().quota,
            blink::mojom::kTinyModelMaxInputTokenSize);
}

TEST_F(AIProofreaderTest, ProofreadMultipleResponse) {
  auto proofreader_remote = GetAIProofreaderRemote();

  std::vector<std::string> result = {"Result ", "text"};
  fake_broker_->settings().set_execute_result(result);
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString),
              ElementsAreArray(result));
}

TEST_F(AIProofreaderTest, MultipleProofread) {
  auto proofreader_remote = GetAIProofreaderRemote();

  std::vector<std::string> result = {"Result ", "text"};
  fake_broker_->settings().set_execute_result(result);
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString),
              ElementsAreArray(result));

  std::vector<std::string> result2 = {"Result ", "text ", "2"};
  fake_broker_->settings().set_execute_result(result2);
  EXPECT_THAT(Proofread(*proofreader_remote, "input string 2"),
              ElementsAreArray(result2));
}

TEST_F(AIProofreaderTest, GetCorretionTypeDefault) {
  fake_broker_->settings().set_execute_result({"Correction type"});

  const auto options = blink::mojom::AIProofreaderCreateOptions::New(
      /*include_correction_types=*/true,
      /*include_correction_explanations=*/false,
      /*correction_explanation_language=*/AILanguageCode::New(""),
      /*expected_input_languages=*/std::vector<AILanguageCodePtr>());

  mojo::Remote<blink::mojom::AIProofreader> proofreader_remote =
      GetAIProofreaderRemote(options.Clone());
  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->GetCorrectionsTypes(kCorrections, responder.BindRemote());
  EXPECT_TRUE(responder.WaitForCompletion());
  EXPECT_THAT(responder.responses_without_last(),
              ElementsAreArray({"Correction type"}));
}

TEST_F(AIProofreaderTest, Priority) {
  fake_broker_->settings().set_execute_result({"hi"});
  auto proofreader_remote = GetAIProofreaderRemote();

  EXPECT_THAT(Proofread(*proofreader_remote, kInputString), ElementsAre("hi"));

  web_contents()->WasHidden();
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString),
              ElementsAre("Priority: background", "hi"));

  web_contents()->WasShown();
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString), ElementsAre("hi"));
}

TEST_F(AIProofreaderTest, TextSafetyInput) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    solution_config.mutable_feature()->set_can_skip_text_safety(false);
    return solution_config;
  }());

  fake_broker_->settings().set_execute_result({"hi"});
  auto proofreader_remote = GetAIProofreaderRemote();
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString), ElementsAre("hi"));

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread("unsafe", responder.BindRemote());
  EXPECT_FALSE(responder.WaitForCompletion());
  EXPECT_EQ(responder.error_status(),
            blink::mojom::ModelStreamingResponseStatus::kErrorFiltered);
}

TEST_F(AIProofreaderTest, TextSafetyOutput) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    solution_config.mutable_feature()->set_can_skip_text_safety(false);
    solution_config.mutable_safety()
        ->mutable_partial_output_checks()
        ->set_minimum_tokens(1000);
    return solution_config;
  }());

  // Fake text safety checker looks for the string "unsafe".
  fake_broker_->settings().set_execute_result(
      {"a", "b", "c", "d", "e", "f", "g", "unsafe", "h"});
  auto proofreader_remote = GetAIProofreaderRemote();
  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread(kInputString, responder.BindRemote());
  EXPECT_FALSE(responder.WaitForCompletion());
  EXPECT_EQ(responder.error_status(),
            blink::mojom::ModelStreamingResponseStatus::kErrorFiltered);
  EXPECT_TRUE(responder.responses().empty());
}

TEST_F(AIProofreaderTest, TextSafetyOutputPartial) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    solution_config.mutable_feature()->set_can_skip_text_safety(false);
    solution_config.mutable_safety()
        ->mutable_partial_output_checks()
        ->set_minimum_tokens(3);
    solution_config.mutable_safety()
        ->mutable_partial_output_checks()
        ->set_token_interval(2);
    return solution_config;
  }());

  // Fake text safety checker looks for the string "unsafe".
  fake_broker_->settings().set_execute_result(
      {"a", "b", "c", "d", "e", "f", "g", "unsafe", "h"});
  auto proofreader_remote = GetAIProofreaderRemote();
  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread(kInputString, responder.BindRemote());
  EXPECT_FALSE(responder.WaitForCompletion());
  EXPECT_EQ(responder.error_status(),
            blink::mojom::ModelStreamingResponseStatus::kErrorFiltered);
  // Partial checks should still allow some output to stream.
  EXPECT_THAT(responder.responses(), ElementsAre("abc", "de", "fg"));
}

TEST_F(AIProofreaderTest, ServiceCrash) {
  fake_broker_->settings().set_execute_result({"hi"});

  auto proofreader_remote = GetAIProofreaderRemote();
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString), ElementsAre("hi"));

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->Proofread(kInputString, responder.BindRemote());
  fake_broker_->launcher().CrashService();

  EXPECT_FALSE(responder.WaitForCompletion());
  // TODO(crbug.com/494980521): Crashes should be yield kErrorSessionDestroyed.
  EXPECT_EQ(
      responder.error_status(),
      blink::mojom::ModelStreamingResponseStatus::kErrorFailedToCountTokens);

  proofreader_remote = GetAIProofreaderRemote();
  EXPECT_THAT(Proofread(*proofreader_remote, kInputString), ElementsAre("hi"));
}

TEST_F(AIProofreaderTest, DynamicConstraints) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    proto::ProofreaderApiMetadata metadata;
    metadata.mutable_constraints()->mutable_label_mode_constraint()->set_regex(
        "^Correction type.*");

    *solution_config.mutable_feature()->mutable_feature_metadata() =
        optimization_guide::AnyWrapProto(metadata);
    return solution_config;
  }());

  fake_broker_->settings().set_execute_result(
      {"Correction type: Spelling"});

  auto options = GetDefaultOptions();
  options->include_correction_types = true;

  mojo::Remote<blink::mojom::AIProofreader> proofreader_remote =
      GetAIProofreaderRemote(std::move(options));

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->GetCorrectionsTypes(kCorrections, responder.BindRemote());
  EXPECT_TRUE(responder.WaitForCompletion());
  EXPECT_THAT(responder.responses_without_last(),
              ElementsAreArray({"Hint: constrained_decoding ",
                                "Constraint: regex ^Correction type.*",
                                "Correction type: Spelling"}));
}

TEST_F(AIProofreaderTest, NoConstraints) {
  SetSolutionConfig([&]() {
    auto solution_config = CreateSolution();
    proto::ProofreaderApiMetadata metadata;

    *solution_config.mutable_feature()->mutable_feature_metadata() =
        optimization_guide::AnyWrapProto(metadata);
    return solution_config;
  }());

  fake_broker_->settings().set_execute_result(
      {"Correction type: Spelling"});

  auto options = GetDefaultOptions();
  options->include_correction_types = true;

  mojo::Remote<blink::mojom::AIProofreader> proofreader_remote =
      GetAIProofreaderRemote(std::move(options));

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->GetCorrectionsTypes(kCorrections, responder.BindRemote());
  EXPECT_TRUE(responder.WaitForCompletion());
  EXPECT_THAT(responder.responses_without_last(),
              ElementsAreArray({"Correction type: Spelling"}));
}

TEST_F(AIProofreaderTest, NoMetadata) {
  SetSolutionConfig([&]() {
    return CreateSolution();
  }());

  fake_broker_->settings().set_execute_result(
      {"Correction type: Spelling"});

  auto options = GetDefaultOptions();
  options->include_correction_types = true;

  mojo::Remote<blink::mojom::AIProofreader> proofreader_remote =
      GetAIProofreaderRemote(std::move(options));

  AITestUtils::TestStreamingResponder responder;
  proofreader_remote->GetCorrectionsTypes(kCorrections, responder.BindRemote());
  EXPECT_TRUE(responder.WaitForCompletion());
  EXPECT_THAT(responder.responses_without_last(),
              ElementsAreArray({"Correction type: Spelling"}));
}

TEST_F(AIProofreaderTest, CreateBuiltInAIAPIsEnterprisePolicyDisabled) {
  SetBuiltInAIAPIsEnterprisePolicy(false);
  base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
  GetAIManagerInterface()->CanCreateProofreader(GetDefaultOptions(),
                                                future.GetCallback());
  EXPECT_EQ(future.Get(), blink::mojom::ModelAvailabilityCheckResult::
                              kUnavailableEnterprisePolicyDisabled);

  mojo::test::BadMessageObserver observer;
  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());
  EXPECT_EQ(observer.WaitForBadMessage(), "Policy or user setting disabled");
  SetBuiltInAIAPIsEnterprisePolicy(true);
}

TEST_F(AIProofreaderTest, CreateGenAILocalEnterprisePolicyDisabled) {
  SetGenAILocalEnterprisePolicy(false);
  base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
  GetAIManagerInterface()->CanCreateProofreader(GetDefaultOptions(),
                                                future.GetCallback());
  EXPECT_EQ(future.Get(), blink::mojom::ModelAvailabilityCheckResult::
                              kUnavailableEnterprisePolicyDisabled);

  mojo::test::BadMessageObserver observer;
  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());
  EXPECT_EQ(observer.WaitForBadMessage(), "Policy or user setting disabled");
  SetGenAILocalEnterprisePolicy(true);
}

TEST_F(AIProofreaderTest, CreateOnDeviceAiUserSettingDisabled) {
  SetOnDeviceAiUserSetting(false);
  base::test::TestFuture<blink::mojom::ModelAvailabilityCheckResult> future;
  GetAIManagerInterface()->CanCreateProofreader(GetDefaultOptions(),
                                                future.GetCallback());
  EXPECT_EQ(future.Get(), blink::mojom::ModelAvailabilityCheckResult::
                              kUnavailableFeatureNotEnabled);

  mojo::test::BadMessageObserver observer;
  TestCreateProofreaderClient create_proofreader_client;
  GetAIManagerRemote()->CreateProofreader(
      create_proofreader_client.BindNewPipeAndPassRemote(), GetDefaultOptions(),
      mojo::NullRemote());
  EXPECT_EQ(observer.WaitForBadMessage(), "Policy or user setting disabled");
  SetOnDeviceAiUserSetting(true);
}

}  // namespace
