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

#include "services/on_device_model/on_device_model_service.h"

#include "base/files/scoped_temp_file.h"
#include "base/json/json_reader.h"
#include "base/strings/strcat.h"
#include "base/test/bind.h"
#include "base/test/metrics/histogram_tester.h"
#include "base/test/scoped_feature_list.h"
#include "base/test/task_environment.h"
#include "base/test/test_future.h"
#include "base/threading/thread_restrictions.h"
#include "mojo/public/cpp/bindings/remote.h"
#include "mojo/public/cpp/test_support/test_utils.h"
#include "services/on_device_model/fake/fake_chrome_ml_api.h"
#include "services/on_device_model/fake/on_device_model_fake.h"
#include "services/on_device_model/ml/chrome_ml_types.h"
#include "services/on_device_model/ml/gpu_blocklist.h"
#include "services/on_device_model/on_device_model_mojom_impl.h"
#include "services/on_device_model/public/cpp/model_assets.h"
#include "services/on_device_model/public/cpp/service_client.h"
#include "services/on_device_model/public/cpp/test_support/test_response_holder.h"
#include "services/on_device_model/public/cpp/text_safety_assets.h"
#include "services/on_device_model/public/mojom/on_device_model.mojom.h"
#include "testing/gmock/include/gmock/gmock.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/abseil-cpp/absl/functional/overload.h"

namespace on_device_model {
namespace {

using ::testing::ElementsAre;

// Creates a test tool declaration matching the fake's tool call name.
ml::ToolDeclaration MakeToolDeclaration() {
  ml::ToolDeclaration decl;
  decl.name = fake_ml::kFakeToolName;
  decl.description = "A test tool";
  decl.input_schema_json =
      R"({"type":"object","properties":{"input":{"type":"string"}}})";
  return decl;
}

ml::ToolCall MakeToolCall(
    std::string arguments_json = R"({"location":{"city":"Paris"}})") {
  ml::ToolCall call;
  call.call_id = fake_ml::kFakeToolCallId;
  call.name = fake_ml::kFakeToolName;
  call.arguments_json = std::move(arguments_json);
  return call;
}

mojom::InputPiecePtr MakeMojomInputPiece(ml::InputPiece piece) {
  return std::visit(
      absl::Overload{
          [](ml::Token token) { return mojom::InputPiece::NewToken(token); },
          [](std::string text) {
            return mojom::InputPiece::NewText(std::move(text));
          },
          [](SkBitmap bitmap) {
            return mojom::InputPiece::NewBitmap(std::move(bitmap));
          },
          [](ml::AudioBuffer audio) {
            return mojom::InputPiece::NewAudio(
                mojom::AudioData::New(audio.sample_rate_hz, audio.num_channels,
                                      audio.num_frames, std::move(audio.data)));
          },
          [](ml::ToolDeclaration decl) {
            auto parsed_schema = base::JSONReader::ReadDict(
                decl.input_schema_json, base::JSON_PARSE_RFC);
            CHECK(parsed_schema.has_value());
            return mojom::InputPiece::NewToolDeclaration(
                mojom::ToolDeclaration::New(std::move(decl.name),
                                            std::move(decl.description),
                                            std::move(*parsed_schema)));
          },
          [](ml::ToolCall call) {
            auto parsed_arguments = base::JSONReader::ReadDict(
                call.arguments_json, base::JSON_PARSE_RFC);
            CHECK(parsed_arguments.has_value());
            return mojom::InputPiece::NewToolCall(mojom::ToolCall::New(
                std::move(call.call_id), std::move(call.name),
                std::move(*parsed_arguments)));
          },
          [](ml::ToolResponse response) {
            std::optional<base::Value> result;
            if (!response.result_json.empty()) {
              result = base::JSONReader::Read(response.result_json,
                                              base::JSON_PARSE_RFC);
              CHECK(result.has_value());
            }
            std::optional<std::string> error_message;
            if (!response.error_message.empty()) {
              error_message = std::move(response.error_message);
            }
            return mojom::InputPiece::NewToolResponse(mojom::ToolResponse::New(
                std::move(response.call_id), std::move(response.name),
                std::move(result), std::move(error_message)));
          },
          [](bool unknown_type) {
            return mojom::InputPiece::NewUnknownType(unknown_type);
          },
      },
      std::move(piece));
}

mojom::InputPtr MakeMojomInput(std::vector<ml::InputPiece> input) {
  auto mojom_input = mojom::Input::New();
  mojom_input->pieces.reserve(input.size());
  for (auto& piece : input) {
    mojom_input->pieces.push_back(MakeMojomInputPiece(std::move(piece)));
  }
  return mojom_input;
}

mojom::InputPtr MakeMojomInput(mojom::InputPiecePtr piece) {
  auto mojom_input = mojom::Input::New();
  mojom_input->pieces.push_back(std::move(piece));
  return mojom_input;
}

mojom::AppendOptionsPtr MakeAppendOptions(mojom::InputPiecePtr piece) {
  auto options = mojom::AppendOptions::New();
  options->input = MakeMojomInput(std::move(piece));
  return options;
}

mojom::InputPiecePtr MakeInvalidToolResponseInputPiece() {
  base::DictValue result_dict;
  result_dict.Set("output", "42");
  std::optional<base::Value> result;
  result.emplace(std::move(result_dict));
  return mojom::InputPiece::NewToolResponse(mojom::ToolResponse::New(
      fake_ml::kFakeToolCallId, fake_ml::kFakeToolName, std::move(result),
      std::make_optional<std::string>("tool failed")));
}

class ContextClientWaiter : public mojom::ContextClient {
 public:
  mojo::PendingRemote<mojom::ContextClient> BindRemote() {
    return receiver_.BindNewPipeAndPassRemote();
  }

  void OnComplete(uint32_t tokens_processed) override {
    tokens_processed_ = tokens_processed;
    run_loop_.Quit();
  }

  int WaitForCompletion() {
    run_loop_.Run();
    return *tokens_processed_;
  }

  bool IsComplete() const { return tokens_processed_.has_value(); }

 private:
  base::RunLoop run_loop_;
  mojo::Receiver<mojom::ContextClient> receiver_{this};
  std::optional<int> tokens_processed_;
};

class FakeFile {
 public:
  explicit FakeFile(const std::string& content) {
    base::ScopedAllowBlockingForTesting allow_blocking;
    CHECK(temp_file_.Create());
    base::File file(temp_file_.path(), base::File::FLAG_OPEN |
                                           base::File::FLAG_WRITE |
                                           base::File::FLAG_READ);
    CHECK(file.IsValid());
    file.WriteAtCurrentPos(base::as_byte_span(content));
  }
  ~FakeFile() = default;

  base::File Open() {
    base::ScopedAllowBlockingForTesting allow_blocking;
    return base::File(temp_file_.path(), base::File::FLAG_OPEN |
                                             base::File::FLAG_WRITE |
                                             base::File::FLAG_READ);
  }

  base::FilePath Path() { return temp_file_.path(); }

 private:
  base::ScopedTempFile temp_file_;
};

class OnDeviceModelServiceTest : public testing::Test {
 public:
  OnDeviceModelServiceTest()
      : service_impl_(service_.BindNewPipeAndPassReceiver(),
                      *fake_ml::GetFakeChromeML()) {}

  mojo::Remote<mojom::OnDeviceModelService>& service() { return service_; }

  mojo::Remote<mojom::OnDeviceModel> LoadModel(
      ml::ModelBackendType backend_type = ml::ModelBackendType::kGpuBackend,
      ml::ModelPerformanceHint performance_hint =
          ml::ModelPerformanceHint::kHighestQuality) {
    mojo::Remote<mojom::OnDeviceModel> remote;
    auto params = mojom::LoadModelParams::New();
    params->backend_type = backend_type;
    params->performance_hint = performance_hint;
    params->max_tokens = 8000;
    params->assets = ModelAssets::FromPath(base::FilePath());
    base::test::TestFuture<mojom::LoadModelResult> future;
    service()->LoadModel(std::move(params), remote.BindNewPipeAndPassReceiver(),
                         future.GetCallback());
    EXPECT_EQ(future.Get(), mojom::LoadModelResult::kSuccess);
    return remote;
  }

  mojo::Remote<mojom::OnDeviceModel> LoadAdaptationWithParams(
      mojom::OnDeviceModel& model,
      mojom::LoadAdaptationParamsPtr adaptation_params) {
    mojo::Remote<mojom::OnDeviceModel> remote;
    base::test::TestFuture<mojom::LoadModelResult> future;
    model.LoadAdaptation(std::move(adaptation_params),
                         remote.BindNewPipeAndPassReceiver(),
                         future.GetCallback());
    EXPECT_EQ(future.Get(), mojom::LoadModelResult::kSuccess);
    return remote;
  }

  mojo::Remote<mojom::OnDeviceModel> LoadAdaptation(
      mojom::OnDeviceModel& model,
      base::File adaptation_data) {
    auto params = mojom::LoadAdaptationParams::New();
    params->assets.weights = std::move(adaptation_data);
    return LoadAdaptationWithParams(model, std::move(params));
  }

  mojo::Remote<mojom::OnDeviceModel> LoadAdaptation(
      mojom::OnDeviceModel& model,
      base::FilePath adaptation_path) {
    auto params = mojom::LoadAdaptationParams::New();
    params->assets.weights_path = std::move(adaptation_path);
    return LoadAdaptationWithParams(model, std::move(params));
  }

  mojom::AppendOptionsPtr MakeInput(const std::string& input) {
    return MakeInput({ml::InputPiece(input)});
  }

  mojom::AppendOptionsPtr MakeInput(std::vector<ml::InputPiece> input) {
    auto options = mojom::AppendOptions::New();
    options->input = MakeMojomInput(std::move(input));
    return options;
  }

  std::vector<std::string> GetResponses(mojom::OnDeviceModel& model,
                                        const std::string& input) {
    TestResponseHolder response;
    mojo::Remote<mojom::Session> session;
    model.StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
    auto options = mojom::AppendOptions::New();
    options->input =
        MakeMojomInput(std::vector<ml::InputPiece>{ml::InputPiece(input)});
    session->Append(std::move(options), {});
    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
    return response.responses();
  }

  std::unique_ptr<ContextClientWaiter> AppendAndFlush(
      mojo::Remote<mojom::Session>& session,
      const std::string& input) {
    auto client = std::make_unique<ContextClientWaiter>();
    session->Append(MakeInput(input), client->BindRemote());
    session.FlushForTesting();
    return client;
  }

  // Creates a session with tool declarations in the system prompt.
  void SetupToolSession(mojom::OnDeviceModel& model,
                        mojo::Remote<mojom::Session>& session) {
    model.StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
    session->Append(MakeInput({ml::Token::kSystem, MakeToolDeclaration(),
                               "Tools could be used.", ml::Token::kEnd}),
                    {});
  }

  size_t GetNumModels() { return service_impl_.NumModelsForTesting(); }

  void ForceQueueing(bool force) {
    service_impl_.SetForceQueueingForTesting(force);
  }

  void FlushService() { service_.FlushForTesting(); }

 protected:
  base::test::TaskEnvironment task_environment_{
      base::test::TaskEnvironment::TimeSource::MOCK_TIME};

 private:
  mojo::Remote<mojom::OnDeviceModelService> service_;
  OnDeviceModelService service_impl_;
  base::test::ScopedFeatureList feature_list_{
      ml::kOnDeviceModelAllowGpuForTesting};
};

TEST_F(OnDeviceModelServiceTest, IdleTimeout) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  base::test::TestFuture<uint32_t, const std::string&> model_future;
  base::test::TestFuture<uint32_t, const std::string&> session_future;
  model.set_disconnect_with_reason_handler(model_future.GetCallback());
  session.set_disconnect_with_reason_handler(session_future.GetCallback());

  session->Append(MakeInput("foo"), {});
  task_environment_.FastForwardBy(kDefaultModelIdleTimeout - base::Seconds(1));
  EXPECT_FALSE(model_future.IsReady());
  EXPECT_FALSE(session_future.IsReady());

  // Another call to the session should reset timeout.
  session->Append(MakeInput("bar"), {});
  task_environment_.FastForwardBy(kDefaultModelIdleTimeout - base::Seconds(1));
  EXPECT_FALSE(model_future.IsReady());
  EXPECT_FALSE(session_future.IsReady());

  // A new session should reset timeout.
  mojo::Remote<mojom::Session> session2;
  model->StartSession(session2.BindNewPipeAndPassReceiver(), nullptr);
  task_environment_.FastForwardBy(kDefaultModelIdleTimeout - base::Seconds(1));
  EXPECT_FALSE(model_future.IsReady());
  EXPECT_FALSE(session_future.IsReady());

  task_environment_.FastForwardBy(base::Seconds(1));
  EXPECT_EQ(std::get<0>(model_future.Get()),
            static_cast<uint32_t>(ModelDisconnectReason::kIdleShutdown));
  EXPECT_EQ(std::get<0>(session_future.Get()),
            static_cast<uint32_t>(ModelDisconnectReason::kIdleShutdown));
}

TEST_F(OnDeviceModelServiceTest, AsrStreamIdleTimeout) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  base::test::TestFuture<uint32_t, const std::string&> model_future;
  base::test::TestFuture<uint32_t, const std::string&> session_future;
  model.set_disconnect_with_reason_handler(model_future.GetCallback());
  session.set_disconnect_with_reason_handler(session_future.GetCallback());

  class DummyResponder : public mojom::AsrStreamResponder {
   public:
    void OnResponse(
        std::vector<mojom::SpeechRecognitionResultPtr> result) override {}
  };
  DummyResponder responder_impl;
  mojo::PendingRemote<mojom::AsrStreamResponder> responder_remote;
  mojo::Receiver<mojom::AsrStreamResponder> receiver(
      &responder_impl, responder_remote.InitWithNewPipeAndPassReceiver());

  auto options = mojom::AsrStreamOptions::New();
  options->sample_rate_hz = 16000;
  mojo::Remote<mojom::AsrStreamInput> asr_input;
  session->AsrStream(std::move(options), asr_input.BindNewPipeAndPassReceiver(),
                     std::move(responder_remote));

  task_environment_.FastForwardBy(kDefaultModelIdleTimeout - base::Seconds(1));
  EXPECT_FALSE(model_future.IsReady());
  EXPECT_FALSE(session_future.IsReady());

  // An ASR chunk should reset timeout.
  auto audio_data = mojom::AudioData::New();
  audio_data->sample_rate = 16000;
  audio_data->channel_count = 1;
  audio_data->frame_count = 1;
  audio_data->data = {0};
  asr_input->AddAudioChunk(std::move(audio_data));
  task_environment_.RunUntilIdle();

  task_environment_.FastForwardBy(kDefaultModelIdleTimeout - base::Seconds(1));
  EXPECT_FALSE(model_future.IsReady());
  EXPECT_FALSE(session_future.IsReady());

  task_environment_.FastForwardBy(base::Seconds(1));
  EXPECT_EQ(std::get<0>(model_future.Get()),
            static_cast<uint32_t>(ModelDisconnectReason::kIdleShutdown));
  EXPECT_EQ(std::get<0>(session_future.Get()),
            static_cast<uint32_t>(ModelDisconnectReason::kIdleShutdown));
}

TEST_F(OnDeviceModelServiceTest, AsrStreamDisconnectDoesNotDisconnectSession) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  class DummyResponder : public mojom::AsrStreamResponder {
   public:
    void OnResponse(
        std::vector<mojom::SpeechRecognitionResultPtr> result) override {}
  };
  DummyResponder responder_impl;
  mojo::PendingRemote<mojom::AsrStreamResponder> responder_remote;
  mojo::Receiver<mojom::AsrStreamResponder> receiver(
      &responder_impl, responder_remote.InitWithNewPipeAndPassReceiver());

  auto options = mojom::AsrStreamOptions::New();
  options->sample_rate_hz = 16000;
  mojo::Remote<mojom::AsrStreamInput> asr_input;
  session->AsrStream(std::move(options), asr_input.BindNewPipeAndPassReceiver(),
                     std::move(responder_remote));
  task_environment_.RunUntilIdle();

  // Disconnect the ASR stream.
  asr_input.reset();
  task_environment_.RunUntilIdle();

  // The session remote should remain connected and functional.
  EXPECT_TRUE(session.is_connected());

  TestResponseHolder response;
  session->Append(MakeInput("test"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();
  EXPECT_THAT(response.responses(), ElementsAre("test"));
}

TEST_F(OnDeviceModelServiceTest, AsrStreamReuseOnExistingSession) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  class DummyResponder : public mojom::AsrStreamResponder {
   public:
    void OnResponse(
        std::vector<mojom::SpeechRecognitionResultPtr> result) override {}
  };

  // First ASR stream on session.
  DummyResponder responder_impl1;
  mojo::PendingRemote<mojom::AsrStreamResponder> responder_remote1;
  mojo::Receiver<mojom::AsrStreamResponder> receiver1(
      &responder_impl1, responder_remote1.InitWithNewPipeAndPassReceiver());
  base::test::TestFuture<void> receiver1_disconnect;
  receiver1.set_disconnect_handler(receiver1_disconnect.GetCallback());

  auto options1 = mojom::AsrStreamOptions::New();
  options1->sample_rate_hz = 16000;
  mojo::Remote<mojom::AsrStreamInput> asr_input1;
  session->AsrStream(std::move(options1),
                     asr_input1.BindNewPipeAndPassReceiver(),
                     std::move(responder_remote1));
  task_environment_.RunUntilIdle();

  auto audio_data1 = mojom::AudioData::New();
  audio_data1->sample_rate = 16000;
  audio_data1->channel_count = 1;
  audio_data1->frame_count = 1;
  audio_data1->data = {0};
  asr_input1->AddAudioChunk(std::move(audio_data1));
  task_environment_.RunUntilIdle();

  // Second ASR stream on the same session replacing the previous stream while
  // the first stream is still open (hot replacement).
  DummyResponder responder_impl2;
  mojo::PendingRemote<mojom::AsrStreamResponder> responder_remote2;
  mojo::Receiver<mojom::AsrStreamResponder> receiver2(
      &responder_impl2, responder_remote2.InitWithNewPipeAndPassReceiver());

  auto options2 = mojom::AsrStreamOptions::New();
  options2->sample_rate_hz = 16000;
  mojo::Remote<mojom::AsrStreamInput> asr_input2;
  session->AsrStream(std::move(options2),
                     asr_input2.BindNewPipeAndPassReceiver(),
                     std::move(responder_remote2));
  task_environment_.RunUntilIdle();

  EXPECT_TRUE(session.is_connected());
  EXPECT_TRUE(asr_input2.is_connected());
  EXPECT_FALSE(asr_input1.is_connected());
  EXPECT_TRUE(receiver1_disconnect.IsReady());

  auto audio_data2 = mojom::AudioData::New();
  audio_data2->sample_rate = 16000;
  audio_data2->channel_count = 1;
  audio_data2->frame_count = 1;
  audio_data2->data = {0};
  asr_input2->AddAudioChunk(std::move(audio_data2));
  task_environment_.RunUntilIdle();
  EXPECT_TRUE(session.is_connected());
}

TEST_F(OnDeviceModelServiceTest, Responds) {
  auto model = LoadModel();
  EXPECT_THAT(GetResponses(*model, "bar"), ElementsAre("bar"));
  // Try another input on  the same model.
  EXPECT_THAT(GetResponses(*model, "cat"), ElementsAre("cat"));
}

TEST_F(OnDeviceModelServiceTest, Append) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput("cheese"), {});
  session->Append(MakeInput("more"), {});
  session->Append(MakeInput("cheddar"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre("cheese", "more", "cheddar"));
}

TEST_F(OnDeviceModelServiceTest, PerSessionSamplingParams) {
  auto model = LoadModel();

  // Sampling params passed at session creation are used during Generate().
  auto session_params = mojom::SessionParams::New();
  session_params->top_k = 2;
  session_params->temperature = 0.5;

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(),
                      std::move(session_params));

  session->Append(MakeInput("cheese"), {});
  session->Append(MakeInput("more"), {});
  session->Append(MakeInput("cheddar"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(),
              ElementsAre("TopK: 2, Temp: 0.5", "cheese", "more", "cheddar"));
}

TEST_F(OnDeviceModelServiceTest, ClampedSamplingParams) {
  auto model = LoadModel();

  // top_k = 0 should be clamped to kMinTopK (1). We use temperature = 0.5 to
  // ensure the params block is printed by the fake engine.
  {
    auto session_params = mojom::SessionParams::New();
    session_params->top_k = 0;
    session_params->temperature = 0.5;

    TestResponseHolder response;
    mojo::Remote<mojom::Session> session;
    model->StartSession(session.BindNewPipeAndPassReceiver(),
                        std::move(session_params));

    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();

    EXPECT_THAT(response.responses(), ElementsAre("TopK: 1, Temp: 0.5"));
  }

  // temperature = -1 should be clamped to kMinTemperature (0.0f). We use
  // top_k = 2 to ensure the params block is printed by the fake engine.
  {
    auto session_params = mojom::SessionParams::New();
    session_params->top_k = 2;
    session_params->temperature = -1;

    TestResponseHolder response;
    mojo::Remote<mojom::Session> session;
    model->StartSession(session.BindNewPipeAndPassReceiver(),
                        std::move(session_params));

    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();

    EXPECT_THAT(response.responses(), ElementsAre("TopK: 2, Temp: 0"));
  }

  // top_k = 1000 should be clamped to MaxTopK (128).
  {
    auto session_params = mojom::SessionParams::New();
    session_params->top_k = 1000;
    session_params->temperature = 0.5;

    TestResponseHolder response;
    mojo::Remote<mojom::Session> session;
    model->StartSession(session.BindNewPipeAndPassReceiver(),
                        std::move(session_params));

    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();

    EXPECT_THAT(response.responses(), ElementsAre("TopK: 128, Temp: 0.5"));
  }
}

TEST_F(OnDeviceModelServiceTest, CloneContextAndContinue) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput("cheese"), {});
  session->Append(MakeInput("more"), {});

  mojo::Remote<mojom::Session> cloned;
  session->Clone(cloned.BindNewPipeAndPassReceiver());

  {
    TestResponseHolder response;
    cloned->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
    EXPECT_THAT(response.responses(), ElementsAre("cheese", "more"));
  }
  {
    TestResponseHolder response;
    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
    EXPECT_THAT(response.responses(), ElementsAre("cheese", "more"));
  }

  session->Append(MakeInput("foo"), {});
  cloned->Append(MakeInput("bar"), {});
  {
    TestResponseHolder response;
    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
    EXPECT_THAT(response.responses(), ElementsAre("cheese", "more", "foo"));
  }
  {
    TestResponseHolder response;
    cloned->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
    EXPECT_THAT(response.responses(), ElementsAre("cheese", "more", "bar"));
  }
}

TEST_F(OnDeviceModelServiceTest, MultipleSessionsAppend) {
  auto model = LoadModel();

  TestResponseHolder response1, response2, response3, response4, response5;
  mojo::Remote<mojom::Session> session1, session2, session3, session4, session5;

  model->StartSession(session1.BindNewPipeAndPassReceiver(), nullptr);
  model->StartSession(session2.BindNewPipeAndPassReceiver(), nullptr);

  session1->Append(MakeInput("cheese"), {});
  session1->Append(MakeInput("more"), {});
  session2->Append(MakeInput("apple"), {});

  session1->Clone(session3.BindNewPipeAndPassReceiver());
  session1->Append(MakeInput("cheddar"), {});
  session1->Generate(mojom::GenerateOptions::New(), response1.BindRemote());

  session2->Append(MakeInput("banana"), {});

  session2->Clone(session4.BindNewPipeAndPassReceiver());
  session2->Append(MakeInput("candy"), {});
  session2->Generate(mojom::GenerateOptions::New(), response2.BindRemote());

  session4->Clone(session5.BindNewPipeAndPassReceiver());
  session4->Append(MakeInput("chip"), {});
  session4->Generate(mojom::GenerateOptions::New(), response3.BindRemote());

  session3->Append(MakeInput("choco"), {});
  session3->Generate(mojom::GenerateOptions::New(), response4.BindRemote());

  session5->Append(MakeInput("orange"), {});
  session5->Generate(mojom::GenerateOptions::New(), response5.BindRemote());

  response1.WaitForCompletion();
  response2.WaitForCompletion();
  response3.WaitForCompletion();
  response4.WaitForCompletion();
  response5.WaitForCompletion();

  EXPECT_THAT(response1.responses(), ElementsAre("cheese", "more", "cheddar"));
  EXPECT_THAT(response2.responses(), ElementsAre("apple", "banana", "candy"));
  EXPECT_THAT(response3.responses(), ElementsAre("apple", "banana", "chip"));
  EXPECT_THAT(response4.responses(), ElementsAre("cheese", "more", "choco"));
  EXPECT_THAT(response5.responses(), ElementsAre("apple", "banana", "orange"));
}

TEST_F(OnDeviceModelServiceTest, CountTokens) {
  auto model = LoadModel();

  std::vector<std::string> inputs = {"cheese", "more", "cheddar"};

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput(inputs.at(0)), {});
  session->Append(MakeInput(inputs.at(1)), {});

  session->Append(MakeInput(inputs.at(2)), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  constexpr int kEosTokenCount = 1;
  EXPECT_THAT(response.output_token_count(), inputs.size() + kEosTokenCount);
}

// TODO(crbug.com/540118700): Remove once the legacy engine has been removed and
// all of these unittests use context_usage by default.
TEST_F(OnDeviceModelServiceTest, CountTokensWithTokenDecodedSet) {
  base::AutoReset<bool> calculate_tokens_decoded =
      fake_ml::EnableCalculateTokensDecodedForTesting();
  auto model = LoadModel();

  std::vector<std::string> inputs = {"cheese", "more", "cheddar"};

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput(inputs.at(0)), {});
  session->Append(MakeInput(inputs.at(1)), {});

  session->Append(MakeInput(inputs.at(2)), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  constexpr int kEosTokenCount = 1;
  EXPECT_THAT(response.output_token_count(), inputs.size() + kEosTokenCount);
}

TEST_F(OnDeviceModelServiceTest, AppendWithTokenLimits) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  std::string input = "big cheese";
  ContextClientWaiter client1;
  auto max_input = MakeInput("big cheese");
  max_input->max_tokens = 4;
  session->Append(std::move(max_input), client1.BindRemote());
  EXPECT_EQ(client1.WaitForCompletion(), 4);

  ContextClientWaiter client2;
  auto offset_input = MakeInput("big cheese");
  session->Append(std::move(offset_input), client2.BindRemote());
  EXPECT_EQ(client2.WaitForCompletion(), 10);

  session->Append(MakeInput("cheddar"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(),
              ElementsAre("big ", "big cheese", "cheddar"));
}

TEST_F(OnDeviceModelServiceTest, MultipleSessionsWaitPreviousSession) {
  auto model = LoadModel();

  TestResponseHolder response1;
  mojo::Remote<mojom::Session> session1;
  model->StartSession(session1.BindNewPipeAndPassReceiver(), nullptr);
  session1->Append(MakeInput("1"), {});
  session1->Generate(mojom::GenerateOptions::New(), response1.BindRemote());

  mojo::Remote<mojom::Session> session2;
  model->StartSession(session2.BindNewPipeAndPassReceiver(), nullptr);

  // First session should not get canceled.
  session1.reset_on_disconnect();
  FlushService();
  EXPECT_TRUE(session1);

  // Response from first session should still work.
  response1.WaitForCompletion();
  EXPECT_THAT(response1.responses(), ElementsAre("1"));

  // Second session still works.
  TestResponseHolder response2;
  session2->Append(MakeInput("2"), {});
  session2->Generate(mojom::GenerateOptions::New(), response2.BindRemote());
  response2.WaitForCompletion();
  EXPECT_THAT(response2.responses(), ElementsAre("2"));
}

TEST_F(OnDeviceModelServiceTest, LoadsAdaptation) {
  FakeFile weights1("Adapt1");
  FakeFile weights2("Adapt2");
  auto model = LoadModel();
  auto adaptation1 = LoadAdaptation(*model, weights1.Open());
  EXPECT_THAT(GetResponses(*model, "foo"), ElementsAre("foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));

  auto adaptation2 = LoadAdaptation(*model, weights2.Open());
  EXPECT_THAT(GetResponses(*model, "foo"), ElementsAre("foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));
  EXPECT_THAT(GetResponses(*adaptation2, "foo"),
              ElementsAre("Adaptation: Adapt2 (1)", "foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));
}

TEST_F(OnDeviceModelServiceTest, LoadsAdaptationWithPath) {
  FakeFile weights1("Adapt1");
  FakeFile weights2("Adapt2");
  auto model = LoadModel(ml::ModelBackendType::kApuBackend);
  auto adaptation1 = LoadAdaptation(*model, weights1.Path());
  EXPECT_THAT(GetResponses(*model, "foo"), ElementsAre("foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));

  auto adaptation2 = LoadAdaptation(*model, weights2.Path());
  EXPECT_THAT(GetResponses(*model, "foo"), ElementsAre("foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));
  EXPECT_THAT(GetResponses(*adaptation2, "foo"),
              ElementsAre("Adaptation: Adapt2 (1)", "foo"));
  EXPECT_THAT(GetResponses(*adaptation1, "foo"),
              ElementsAre("Adaptation: Adapt1 (0)", "foo"));
}

TEST_F(OnDeviceModelServiceTest, LoadingAdaptationDoesNotCancelSession) {
  FakeFile weights1("Adapt1");
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session.reset_on_disconnect();

  LoadAdaptation(*model, weights1.Open());
  FlushService();
  EXPECT_TRUE(session);
}

TEST_F(OnDeviceModelServiceTest, DeletesModel) {
  FakeFile weights1("Adapt1");
  FakeFile weights2("Adapt2");
  FakeFile weights3("Adapt3");
  auto model1 = LoadModel();
  auto adaptation1 = LoadAdaptation(*model1, weights1.Open());
  auto adaptation2 = LoadAdaptation(*model1, weights2.Open());
  EXPECT_EQ(GetNumModels(), 1u);

  auto model2 = LoadModel();
  auto adaptation3 = LoadAdaptation(*model2, weights3.Open());
  EXPECT_EQ(GetNumModels(), 2u);

  adaptation1.reset();
  adaptation2.reset();
  FlushService();
  EXPECT_EQ(GetNumModels(), 2u);

  model1.reset();
  FlushService();
  EXPECT_EQ(GetNumModels(), 1u);

  model2.reset();
  FlushService();
  EXPECT_EQ(GetNumModels(), 1u);

  adaptation3.reset();
  FlushService();
  EXPECT_EQ(GetNumModels(), 0u);
}

TEST_F(OnDeviceModelServiceTest, Score) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput("hi"), {});

  {
    base::test::TestFuture<float> future;
    session->Score("x", future.GetCallback());
    EXPECT_EQ(future.Get(), float('x'));
  }
  {
    base::test::TestFuture<float> future;
    session->Score("y", future.GetCallback());
    EXPECT_EQ(future.Get(), float('y'));
  }
}

TEST_F(OnDeviceModelServiceTest, AppendWithTokens) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  {
    std::vector<ml::InputPiece> pieces;
    pieces.push_back(ml::Token::kSystem);
    pieces.push_back("hi");
    pieces.push_back(ml::Token::kEnd);
    session->Append(MakeInput(std::move(pieces)), {});
  }
  {
    std::vector<ml::InputPiece> pieces;
    pieces.push_back(ml::Token::kModel);
    pieces.push_back("hello");
    pieces.push_back(ml::Token::kEnd);
    session->Append(MakeInput(std::move(pieces)), {});
  }
  {
    std::vector<ml::InputPiece> pieces;
    pieces.push_back(ml::Token::kUser);
    pieces.push_back("bye");
    session->Append(MakeInput(std::move(pieces)), {});
    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  }
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(),
              ElementsAre("System: hi End.", "Model: hello End.", "User: bye"));
}

TEST_F(OnDeviceModelServiceTest, AppendWithImages) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  auto params = mojom::SessionParams::New();
  params->capabilities.Put(CapabilityFlags::kImageInput);
  model->StartSession(session.BindNewPipeAndPassReceiver(), std::move(params));

  {
    std::vector<ml::InputPiece> pieces;
    pieces.push_back("cheddar");

    SkBitmap cheesy_bitmap;
    cheesy_bitmap.allocPixels(
        SkImageInfo::Make(7, 21, kRGBA_8888_SkColorType, kOpaque_SkAlphaType),
        0);
    cheesy_bitmap.eraseColor(SK_ColorYELLOW);
    pieces.push_back(cheesy_bitmap);

    pieces.push_back("cheese");

    session->Append(MakeInput(std::move(pieces)), {});
  }

  TestResponseHolder response;
  {
    std::vector<ml::InputPiece> pieces;
    pieces.push_back("bleu");

    SkBitmap moldy_cheese;
    moldy_cheese.allocPixels(
        SkImageInfo::Make(63, 42, kRGBA_8888_SkColorType, kOpaque_SkAlphaType),
        0);
    moldy_cheese.eraseColor(SK_ColorBLUE);
    pieces.push_back(moldy_cheese);

    pieces.push_back("cheese");

    session->Append(MakeInput(std::move(pieces)), {});
    session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
    response.WaitForCompletion();
  }

  EXPECT_THAT(response.responses(),
              ElementsAre("cheddar[Bitmap of size 7x21]cheese",
                          "bleu[Bitmap of size 63x42]cheese"));
}

TEST_F(OnDeviceModelServiceTest, GpuBlocked) {
  // The fake implementation of ChromeML always blocks GPU by default.
  base::test::ScopedFeatureList feature_list;
  feature_list.InitAndDisableFeature(ml::kOnDeviceModelAllowGpuForTesting);

  mojo::Remote<mojom::OnDeviceModel> remote;
  auto params = mojom::LoadModelParams::New();
  params->backend_type = ml::ModelBackendType::kGpuBackend;
  params->max_tokens = 8000;
  params->assets = ModelAssets::FromPath(base::FilePath());
  base::test::TestFuture<mojom::LoadModelResult> future;
  service()->LoadModel(std::move(params), remote.BindNewPipeAndPassReceiver(),
                       future.GetCallback());
  EXPECT_EQ(future.Get(), mojom::LoadModelResult::kGpuBlocked);
}

TEST_F(OnDeviceModelServiceTest, CpuModel) {
  // The fake implementation of ChromeML always blocks GPU by default.
  base::test::ScopedFeatureList feature_list;
  feature_list.InitAndDisableFeature(ml::kOnDeviceModelAllowGpuForTesting);

  auto model = LoadModel(ml::ModelBackendType::kCpuBackend);
  EXPECT_THAT(GetResponses(*model, "foo"), ElementsAre("CPU backend", "foo"));
}

TEST_F(OnDeviceModelServiceTest, PerformanceHint) {
  auto model = LoadModel(ml::ModelBackendType::kGpuBackend,
                         ml::ModelPerformanceHint::kFastestInference);
  EXPECT_THAT(GetResponses(*model, "foo"),
              ElementsAre("Fastest inference", "foo"));
}

TEST_F(OnDeviceModelServiceTest, Capabilities) {
  auto expect_capabilities = [&](const std::string& data,
                                 const Capabilities& expected) {
    FakeFile file(data);
    ModelFile model_file(file.Open());
    base::test::TestFuture<const Capabilities&> future;
    service()->GetCapabilities(std::move(model_file), future.GetCallback());
    EXPECT_EQ(expected, future.Take());
  };
  expect_capabilities("none", {});
  expect_capabilities("image", {CapabilityFlags::kImageInput});
  expect_capabilities("audio", {CapabilityFlags::kAudioInput});
  expect_capabilities("image audio", {CapabilityFlags::kImageInput,
                                      CapabilityFlags::kAudioInput});
}

TEST_F(OnDeviceModelServiceTest, CapabilitiesFromFilePath) {
  auto expect_capabilities = [&](const std::string& data,
                                 const Capabilities& expected) {
    FakeFile file(data);
    ModelFile model_file(file.Path());
    base::test::TestFuture<const Capabilities&> future;
    service()->GetCapabilities(std::move(model_file), future.GetCallback());
    EXPECT_EQ(expected, future.Take());
  };
  expect_capabilities("none", {});
  expect_capabilities("image", {CapabilityFlags::kImageInput});
  expect_capabilities("audio", {CapabilityFlags::kAudioInput});
  expect_capabilities("image audio", {CapabilityFlags::kImageInput,
                                      CapabilityFlags::kAudioInput});
}

TEST_F(OnDeviceModelServiceTest, SetPriority) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> background;
  model->StartSession(background.BindNewPipeAndPassReceiver(), nullptr);
  background->SetPriority(mojom::Priority::kBackground);

  mojo::Remote<mojom::Session> foreground;
  model->StartSession(foreground.BindNewPipeAndPassReceiver(), nullptr);

  base::HistogramTester histogram_tester;

  ForceQueueing(true);
  auto bg_waiter = AppendAndFlush(background, "bg");
  auto fg_waiter = AppendAndFlush(foreground, "fg");

  constexpr char kForegroundHistogram[] = "OnDeviceModel.QueueTime.Foreground";
  constexpr char kBackgroundHistogram[] = "OnDeviceModel.QueueTime.Background";
  histogram_tester.ExpectTotalCount(kForegroundHistogram, 0);
  histogram_tester.ExpectTotalCount(kBackgroundHistogram, 0);
  ForceQueueing(false);

  fg_waiter->WaitForCompletion();
  EXPECT_FALSE(bg_waiter->IsComplete());
  histogram_tester.ExpectTotalCount(kForegroundHistogram, 1);
  histogram_tester.ExpectTotalCount(kBackgroundHistogram, 0);

  ForceQueueing(true);

  // Add another call to fg client, should jump ahead of bg again.
  fg_waiter = AppendAndFlush(foreground, "fg");
  ForceQueueing(false);

  fg_waiter->WaitForCompletion();
  EXPECT_FALSE(bg_waiter->IsComplete());
  histogram_tester.ExpectTotalCount(kForegroundHistogram, 2);
  histogram_tester.ExpectTotalCount(kBackgroundHistogram, 0);

  bg_waiter->WaitForCompletion();
  histogram_tester.ExpectTotalCount(kForegroundHistogram, 2);
  histogram_tester.ExpectTotalCount(kBackgroundHistogram, 1);
}

TEST_F(OnDeviceModelServiceTest, SetPriorityAfterQueue) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> background;
  model->StartSession(background.BindNewPipeAndPassReceiver(), nullptr);

  mojo::Remote<mojom::Session> foreground;
  model->StartSession(foreground.BindNewPipeAndPassReceiver(), nullptr);

  ForceQueueing(true);
  auto bg_waiter = AppendAndFlush(background, "bg");
  auto fg_waiter = AppendAndFlush(foreground, "fg");

  background->SetPriority(mojom::Priority::kBackground);
  background.FlushForTesting();
  ForceQueueing(false);

  fg_waiter->WaitForCompletion();
  EXPECT_FALSE(bg_waiter->IsComplete());
  bg_waiter->WaitForCompletion();
}

TEST_F(OnDeviceModelServiceTest, SetPriorityBackToForeground) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> background;
  model->StartSession(background.BindNewPipeAndPassReceiver(), nullptr);
  background->SetPriority(mojom::Priority::kBackground);

  mojo::Remote<mojom::Session> foreground;
  model->StartSession(foreground.BindNewPipeAndPassReceiver(), nullptr);

  ForceQueueing(true);

  auto bg_waiter = AppendAndFlush(background, "bg");
  auto fg_waiter = AppendAndFlush(foreground, "fg");

  ForceQueueing(false);
  fg_waiter->WaitForCompletion();
  EXPECT_FALSE(bg_waiter->IsComplete());

  ForceQueueing(true);

  fg_waiter = AppendAndFlush(foreground, "fg");

  background->SetPriority(mojom::Priority::kForeground);
  background.FlushForTesting();

  ForceQueueing(false);
  bg_waiter->WaitForCompletion();

  EXPECT_FALSE(fg_waiter->IsComplete());
  fg_waiter->WaitForCompletion();
}

TEST_F(OnDeviceModelServiceTest, SetPriorityMultipleSessions) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> background1;
  model->StartSession(background1.BindNewPipeAndPassReceiver(), nullptr);
  background1->SetPriority(mojom::Priority::kBackground);

  mojo::Remote<mojom::Session> background2;
  model->StartSession(background2.BindNewPipeAndPassReceiver(), nullptr);
  background2->SetPriority(mojom::Priority::kBackground);

  mojo::Remote<mojom::Session> foreground1;
  model->StartSession(foreground1.BindNewPipeAndPassReceiver(), nullptr);

  mojo::Remote<mojom::Session> foreground2;
  model->StartSession(foreground2.BindNewPipeAndPassReceiver(), nullptr);

  std::set<ContextClientWaiter*> all;
  auto append = [&](mojo::Remote<mojom::Session>& session) {
    std::unique_ptr<ContextClientWaiter> waiter = AppendAndFlush(session, "in");
    all.insert(waiter.get());
    return waiter;
  };
  ForceQueueing(true);
  auto bg1_waiter1 = append(background1);
  auto bg2_waiter1 = append(background2);
  auto fg1_waiter1 = append(foreground1);
  auto fg2_waiter1 = append(foreground2);
  auto fg1_waiter2 = append(foreground1);
  auto bg2_waiter2 = append(background2);
  auto fg2_waiter2 = append(foreground2);
  auto bg1_waiter2 = append(background1);
  ForceQueueing(false);

  auto wait_for_next = [&](ContextClientWaiter* next) {
    next->WaitForCompletion();
    all.erase(next);
    for (auto* waiter : all) {
      EXPECT_FALSE(waiter->IsComplete());
    }
  };
  wait_for_next(fg1_waiter1.get());

  // Add another item, should be added at the end of fg items.
  ForceQueueing(true);
  fg1_waiter1 = append(foreground1);
  ForceQueueing(false);

  wait_for_next(fg2_waiter1.get());
  wait_for_next(fg1_waiter2.get());
  wait_for_next(fg2_waiter2.get());
  wait_for_next(fg1_waiter1.get());
  wait_for_next(bg1_waiter1.get());

  // Add a few fg and bg items, fg should run immediately, bg should run last.
  ForceQueueing(true);
  bg1_waiter1 = append(background1);
  fg1_waiter1 = append(foreground1);
  fg2_waiter1 = append(foreground2);
  ForceQueueing(false);

  wait_for_next(fg1_waiter1.get());
  wait_for_next(fg2_waiter1.get());
  wait_for_next(bg2_waiter1.get());
  wait_for_next(bg2_waiter2.get());
  wait_for_next(bg1_waiter2.get());

  // Add another bg item, but bump priority to fg, should run immediately.
  ForceQueueing(true);
  bg2_waiter1 = append(background2);
  background2->SetPriority(mojom::Priority::kForeground);
  background2.FlushForTesting();
  ForceQueueing(false);

  wait_for_next(bg2_waiter1.get());
  wait_for_next(bg1_waiter1.get());
}

TEST_F(OnDeviceModelServiceTest, SetPriorityCloneInherits) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> background;
  model->StartSession(background.BindNewPipeAndPassReceiver(), nullptr);
  background->SetPriority(mojom::Priority::kBackground);

  mojo::Remote<mojom::Session> foreground;
  model->StartSession(foreground.BindNewPipeAndPassReceiver(), nullptr);

  mojo::Remote<mojom::Session> clone;
  background->Clone(clone.BindNewPipeAndPassReceiver());
  background.FlushForTesting();

  ForceQueueing(true);
  auto bg_waiter = AppendAndFlush(background, "bg");
  auto clone_waiter = AppendAndFlush(clone, "clone");
  auto fg_waiter = AppendAndFlush(foreground, "fg");
  ForceQueueing(false);

  fg_waiter->WaitForCompletion();
  EXPECT_FALSE(bg_waiter->IsComplete());
  EXPECT_FALSE(clone_waiter->IsComplete());

  bg_waiter->WaitForCompletion();
  EXPECT_FALSE(clone_waiter->IsComplete());

  clone_waiter->WaitForCompletion();
}

TEST_F(OnDeviceModelServiceTest, ToolCallsNotGeneratedWithoutDeclarations) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  // Append without tool declarations hint.
  session->Append(MakeInput("some input"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  // Should complete normally without tool calls.
  EXPECT_FALSE(response.has_tool_calls());
  EXPECT_TRUE(response.complete());
}

TEST_F(OnDeviceModelServiceTest, ToolDeclarationAndCallGeneration) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  SetupToolSession(*model, session);
  session->Append(MakeInput({ml::Token::kUser, "What is the weather?"}), {});

  TestResponseHolder response;
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_TRUE(response.complete());
  // Verify the fake echoed back tool declarations from the system prompt.
  EXPECT_THAT(response.responses(),
              testing::Contains(testing::HasSubstr(base::StrCat(
                  {fake_ml::kToolDeclPrefix, fake_ml::kFakeToolName, "]"}))));
  ASSERT_TRUE(response.has_tool_calls());
  ASSERT_EQ(response.tool_calls().size(), 1u);
  EXPECT_EQ(response.tool_calls()[0]->call_id, fake_ml::kFakeToolCallId);
  EXPECT_EQ(response.tool_calls()[0]->name, fake_ml::kFakeToolName);
}

TEST_F(OnDeviceModelServiceTest, ToolResponseProcessing) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  SetupToolSession(*model, session);
  session->Append(MakeInput({ml::Token::kUser, "Please use the test tool"}),
                  {});

  // First generation triggers tool calls and completes.
  TestResponseHolder response;
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_TRUE(response.complete());
  ASSERT_TRUE(response.has_tool_calls());

  // Send tool responses back via Append.
  ml::ToolResponse tool_resp;
  tool_resp.call_id = fake_ml::kFakeToolCallId;
  tool_resp.name = fake_ml::kFakeToolName;
  tool_resp.result_json = R"({"output":"42"})";
  session->Append(MakeInput({ml::Token::kToolResponse, std::move(tool_resp),
                             ml::Token::kEnd}),
                  {});

  // Second generation incorporates tool response context.
  TestResponseHolder response2;
  session->Generate(mojom::GenerateOptions::New(), response2.BindRemote());
  response2.WaitForCompletion();

  EXPECT_TRUE(response2.complete());
  EXPECT_THAT(response2.responses(),
              testing::Contains(testing::HasSubstr(base::StrCat(
                  {fake_ml::kToolRespPrefix, fake_ml::kFakeToolName, "="}))));
}

TEST_F(OnDeviceModelServiceTest, ToolCallInputProcessing) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(
      MakeInput({ml::Token::kModel, MakeToolCall(), ml::Token::kEnd}), {});

  TestResponseHolder response;
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_TRUE(response.complete());
  EXPECT_THAT(
      response.responses(),
      testing::Contains(testing::HasSubstr(base::StrCat(
          {fake_ml::kToolCallPrefix, fake_ml::kFakeToolCallId, ":",
           fake_ml::kFakeToolName, R"(={"location":{"city":"Paris"}}])"}))));
}

TEST_F(OnDeviceModelServiceTest, ToolCallInputSizeInTokens) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  base::test::TestFuture<uint32_t> future;
  session->GetSizeInTokens(
      MakeMojomInput(std::vector<ml::InputPiece>{MakeToolCall("{}")}),
      future.GetCallback());

  EXPECT_EQ(future.Get(),
            base::StrCat({fake_ml::kToolCallPrefix, fake_ml::kFakeToolCallId,
                          ":", fake_ml::kFakeToolName, "={}]"})
                .size());
}

TEST_F(OnDeviceModelServiceTest, ToolCallInputPreservedInClone) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  auto append_waiter = std::make_unique<ContextClientWaiter>();
  session->Append(
      MakeInput({ml::Token::kModel, MakeToolCall(), ml::Token::kEnd}),
      append_waiter->BindRemote());
  append_waiter->WaitForCompletion();

  mojo::Remote<mojom::Session> cloned;
  session->Clone(cloned.BindNewPipeAndPassReceiver());
  TestResponseHolder response;
  cloned->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_TRUE(response.complete());
  EXPECT_THAT(
      response.responses(),
      testing::Contains(testing::HasSubstr(base::StrCat(
          {fake_ml::kToolCallPrefix, fake_ml::kFakeToolCallId, ":",
           fake_ml::kFakeToolName, R"(={"location":{"city":"Paris"}}])"}))));
}

TEST_F(OnDeviceModelServiceTest, InvalidToolResponseReportsBadMessageOnAppend) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  mojo::test::BadMessageObserver observer;
  session->Append(MakeAppendOptions(MakeInvalidToolResponseInputPiece()), {});

  EXPECT_EQ(observer.WaitForBadMessage(),
            "SessionAccessor::AppendInternal: failed to convert input");
}

TEST_F(OnDeviceModelServiceTest,
       InvalidToolResponseReportsBadMessageOnSizeInTokens) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  mojo::test::BadMessageObserver observer;
  session->GetSizeInTokens(MakeMojomInput(MakeInvalidToolResponseInputPiece()),
                           base::DoNothing());

  EXPECT_EQ(observer.WaitForBadMessage(),
            "SessionAccessor::SizeInTokensInternal: failed to convert input");
}

TEST_F(OnDeviceModelServiceTest, ToolDeclarationsIgnoredOutsideSystemPrompt) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  // Tool declarations outside the system prompt.
  session->Append(MakeInput({ml::Token::kUser, "Please help me. ",
                             MakeToolDeclaration(), ml::Token::kEnd}),
                  {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_FALSE(response.has_tool_calls());
  EXPECT_TRUE(response.complete());
}

TEST_F(OnDeviceModelServiceTest, ToolCallsWithClonedSession) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  SetupToolSession(*model, session);

  // Clone the session.
  mojo::Remote<mojom::Session> cloned;
  session->Clone(cloned.BindNewPipeAndPassReceiver());

  // Generate on the cloned session.
  TestResponseHolder response;
  cloned->Append(MakeInput({ml::Token::kUser, "Calculate 2+2"}), {});
  cloned->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();

  // Cloned session should still generate tool calls.
  EXPECT_TRUE(response.complete());
  ASSERT_TRUE(response.has_tool_calls());
  ASSERT_EQ(response.tool_calls().size(), 1u);
  EXPECT_EQ(response.tool_calls()[0]->call_id, fake_ml::kFakeToolCallId);
  EXPECT_EQ(response.tool_calls()[0]->name, fake_ml::kFakeToolName);
}

TEST_F(OnDeviceModelServiceTest, GenerateRejectedWhileAwaitingToolResponses) {
  auto model = LoadModel();

  mojo::Remote<mojom::Session> session;
  SetupToolSession(*model, session);
  session->Append(MakeInput({ml::Token::kUser, "Please use the test tool"}),
                  {});

  // First generation triggers tool calls.
  TestResponseHolder response1;
  session->Generate(mojom::GenerateOptions::New(), response1.BindRemote());
  response1.WaitForCompletion();
  EXPECT_TRUE(response1.complete());
  ASSERT_TRUE(response1.has_tool_calls());

  // Second generation without tool response should be rejected by closing
  // the responder pipe.
  TestResponseHolder response2;
  session->Generate(mojom::GenerateOptions::New(), response2.BindRemote());
  response2.WaitForCompletion();
  EXPECT_TRUE(response2.disconnected());
  EXPECT_FALSE(response2.complete());

  // After providing tool responses, generation should succeed again.
  ml::ToolResponse tool_resp;
  tool_resp.call_id = fake_ml::kFakeToolCallId;
  tool_resp.name = fake_ml::kFakeToolName;
  tool_resp.result_json = R"({"output":"42"})";
  session->Append(MakeInput({ml::Token::kToolResponse, std::move(tool_resp),
                             ml::Token::kEnd}),
                  {});

  TestResponseHolder response3;
  session->Generate(mojom::GenerateOptions::New(), response3.BindRemote());
  response3.WaitForCompletion();
  EXPECT_TRUE(response3.complete());
  EXPECT_THAT(response3.responses(),
              testing::Contains(testing::HasSubstr(base::StrCat(
                  {fake_ml::kToolRespPrefix, fake_ml::kFakeToolName, "="}))));
}

#if defined(ENABLE_ON_DEVICE_CONSTRAINTS)
TEST_F(OnDeviceModelServiceTest, JSONSchemaConstraint) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewJsonSchema(R"({
    "type": "object",
    "required": ["Rating"],
    "additionalProperties": false,
    "properties": {
      "Rating": {
        "type": "number",
        "minimum": 1,
        "maximum": 5
      }
    }
  })");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre(R"({"Rating":1})"));
}

TEST_F(OnDeviceModelServiceTest, JSONSchemaConstraintWithPrefix) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput({"a", ml::Token::kModel, "{\"Rating\""}), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewJsonSchema(R"({
    "type": "object",
    "required": ["Rating"],
    "additionalProperties": false,
    "properties": {
      "Rating": {
        "type": "number",
        "minimum": 1,
        "maximum": 5
      }
    }
  })");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre("aModel: {\"Rating\"", ":1}"));
}

TEST_F(OnDeviceModelServiceTest, JSONSchemaConstraintWithInvalidPrefix) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput({"a", ml::Token::kModel, "{\"bad\""}), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewJsonSchema(R"({
    "type": "object",
    "required": ["Rating"],
    "additionalProperties": false,
    "properties": {
      "Rating": {
        "type": "number",
        "minimum": 1,
        "maximum": 5
      }
    }
  })");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  // For now invalid prefix will cause a disconnect.
  // TODO:crbug.com/434766400 - Add better error messages.
  EXPECT_THAT(response.responses(), ElementsAre());
  EXPECT_TRUE(response.disconnected());
}

TEST_F(OnDeviceModelServiceTest, JSONSchemaConstraintInvalid) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput("hi"), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewJsonSchema("blah");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  // For now invalid schema will cause a disconnect.
  EXPECT_THAT(response.responses(), ElementsAre());
  EXPECT_TRUE(response.disconnected());
}

TEST_F(OnDeviceModelServiceTest, RegexConstraint) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewRegex("hello");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre("hello"));
}

TEST_F(OnDeviceModelServiceTest, RegexConstraintIgnoresUserPrefix) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput({ml::Token::kUser, "hel"}), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewRegex("hello");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre("User: hel", "hello"));
}

TEST_F(OnDeviceModelServiceTest, RegexConstraintWithPrefix) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput({"a", ml::Token::kModel, "hel"}), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewRegex("hello");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  EXPECT_THAT(response.responses(), ElementsAre("aModel: hel", "lo"));
}

TEST_F(OnDeviceModelServiceTest, RegexConstraintWithInvalidPrefix) {
  auto model = LoadModel();

  TestResponseHolder response;
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);
  session->Append(MakeInput({ml::Token::kModel, "boo"}), {});

  auto options = mojom::GenerateOptions::New();
  options->constraint = mojom::ResponseConstraint::NewRegex("^hello$");
  session->Generate(std::move(options), response.BindRemote());
  response.WaitForCompletion();

  // For now invalid prefix will cause a disconnect.
  // TODO:crbug.com/434766400 - Add better error messages.
  EXPECT_THAT(response.responses(), ElementsAre());
  EXPECT_TRUE(response.disconnected());
}

#endif

TEST_F(OnDeviceModelServiceTest, AsrStreamInitializationFailure) {
  auto model = LoadModel();
  mojo::Remote<mojom::Session> session;
  model->StartSession(session.BindNewPipeAndPassReceiver(), nullptr);

  class DummyResponder : public mojom::AsrStreamResponder {
   public:
    void OnResponse(
        std::vector<mojom::SpeechRecognitionResultPtr> result) override {}
  };
  DummyResponder responder_impl;
  mojo::PendingRemote<mojom::AsrStreamResponder> responder_remote;
  mojo::Receiver<mojom::AsrStreamResponder> receiver(
      &responder_impl, responder_remote.InitWithNewPipeAndPassReceiver());

  base::test::TestFuture<uint32_t, const std::string&> received_reason_future;
  receiver.set_disconnect_with_reason_handler(
      received_reason_future.GetCallback());

  auto options = mojom::AsrStreamOptions::New();
  options->sample_rate_hz = 0;
  mojo::PendingRemote<mojom::AsrStreamInput> asr_input;
  session->AsrStream(std::move(options),
                     asr_input.InitWithNewPipeAndPassReceiver(),
                     std::move(responder_remote));

  EXPECT_TRUE(received_reason_future.Wait());
  EXPECT_EQ(std::get<0>(received_reason_future.Take()),
            static_cast<uint32_t>(mojom::AsrError::kInitializationFailed));

  // The session remote should remain connected and functional after failure.
  EXPECT_TRUE(session.is_connected());

  TestResponseHolder response;
  session->Append(MakeInput("test"), {});
  session->Generate(mojom::GenerateOptions::New(), response.BindRemote());
  response.WaitForCompletion();
  EXPECT_THAT(response.responses(), ElementsAre("test"));
}

}  // namespace
}  // namespace on_device_model
