// 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 "services/on_device_model/android/backend_session_impl_android.h"

#include <algorithm>
#include <iterator>
#include <memory>
#include <string>
#include <variant>
#include <vector>

#include "base/android/jni_android.h"
#include "base/android/jni_array.h"
#include "base/android/jni_string.h"
#include "base/functional/bind.h"
#include "base/functional/callback.h"
#include "base/memory/ptr_util.h"
#include "base/metrics/histogram_functions.h"
#include "base/notimplemented.h"
#include "base/sequence_checker.h"
#include "base/strings/strcat.h"
#include "components/optimization_guide/core/optimization_guide_util.h"
#include "mojo/public/cpp/bindings/pending_remote.h"
#include "services/on_device_model/android/on_device_model_bridge.h"
#include "services/on_device_model/ml/chrome_ml_types.h"
#include "services/on_device_model/public/mojom/on_device_model.mojom.h"

// Must come after all headers that specialize FromJniType() / ToJniType().
#include "services/on_device_model/android/jni_headers/AiCoreSessionWrapper_jni.h"
#include "services/on_device_model/android/jni_headers/GenerateOptionsHelper_jni.h"
#include "services/on_device_model/android/jni_headers/InputPieceHelper_jni.h"

namespace on_device_model {

namespace {

// Converts mojom input pieces to Java InputPiece objects. Only token and text
// pieces are supported on Android.
std::vector<base::android::ScopedJavaLocalRef<jobject>>
ConvertInputPiecesToJava(
    JNIEnv* env,
    const std::vector<on_device_model::mojom::InputPiecePtr>& pieces) {
  std::vector<base::android::ScopedJavaLocalRef<jobject>> java_inputs;
  using Tag = on_device_model::mojom::InputPiece::Tag;
  for (const auto& piece : pieces) {
    switch (piece->which()) {
      case Tag::kToken:
        java_inputs.push_back(Java_InputPieceHelper_fromToken(
            env, static_cast<int>(piece->get_token())));
        break;
      case Tag::kText:
        java_inputs.push_back(Java_InputPieceHelper_fromText(
            env,
            base::android::ConvertUTF8ToJavaString(env, piece->get_text())));
        break;
      case Tag::kBitmap:
      case Tag::kAudio:
      case Tag::kToolDeclaration:
      case Tag::kToolCall:
      case Tag::kToolResponse:
      case Tag::kUnknownType:
        // TODO(crbug.com/425408635): Support image, audio, and other input
        // types.
        NOTREACHED();
    }
  }
  return java_inputs;
}

}  // namespace

BackendSessionImplAndroid::BackendSessionImplAndroid(
    optimization_guide::proto::ModelExecutionFeature feature,
    on_device_model::mojom::SessionParamsPtr params,
    std::vector<on_device_model::mojom::InputPiecePtr> context_input_pieces)
    : java_session_(
          OnDeviceModelBridge::CreateSession(feature, params.Clone())),
      context_input_pieces_(std::move(context_input_pieces)),
      feature_(feature),
      params_(std::move(params)) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  weak_ptr_ = weak_factory_.GetWeakPtr();
}

BackendSessionImplAndroid::BackendSessionImplAndroid(
    optimization_guide::proto::ModelExecutionFeature feature,
    on_device_model::mojom::SessionParamsPtr params)
    : BackendSessionImplAndroid(feature,
                                std::move(params),
                                /*context_input_pieces=*/{}) {}

BackendSessionImplAndroid::~BackendSessionImplAndroid() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  JNIEnv* env = base::android::AttachCurrentThread();
  Java_AiCoreSessionWrapper_onNativeDestroyed(env, java_session_);
}

void BackendSessionImplAndroid::Append(
    on_device_model::mojom::AppendOptionsPtr options,
    mojo::PendingRemote<on_device_model::mojom::ContextClient> client,
    mojo::ReportBadMessageCallback bad_message_callback,
    base::OnceClosure on_complete) {
  context_input_pieces_.insert(
      context_input_pieces_.end(),
      std::make_move_iterator(options->input->pieces.begin()),
      std::make_move_iterator(options->input->pieces.end()));
  if (client) {
    // Bind the context client and signal completion to prevent the caller's
    // mojo pipe disconnect handler which invokes OnError from firing.
    // Temporarily pass 0 for tokens_processed which is used for UMA histograms
    // and context window bookkeeping, neither of which affects model execution
    // correctness on Android where each Generate call sends the full input
    // independently.
    // TODO(crbug.com/477033510): Report actual token count if Android supports
    // stateful context window for prompt API in the future.
    mojo::Remote<on_device_model::mojom::ContextClient> context_client(
        std::move(client));
    context_client->OnComplete(/*tokens_processed=*/0);
  }
  std::move(on_complete).Run();
}

void BackendSessionImplAndroid::Generate(
    on_device_model::mojom::GenerateOptionsPtr input,
    mojo::PendingRemote<on_device_model::mojom::StreamingResponder> response,
    base::OnceClosure on_complete) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  CHECK(!responder_.is_bound()) << "Caller should not call Generate() again "
                                   "before OnComplete() is received.";
  responder_.Bind(std::move(response));

  JNIEnv* env = base::android::AttachCurrentThread();
  // There isn't a generic mojo utility for converting c++ mojo struct to java,
  // so disassemble the struct here and reassemble it in java.
  // Only passing the parameters that are supported on Android.
  base::android::ScopedJavaLocalRef<jobject> java_generate_options =
      Java_GenerateOptionsHelper_create(env, input->max_output_tokens);

  std::vector<base::android::ScopedJavaLocalRef<jobject>> java_inputs =
      ConvertInputPiecesToJava(env, context_input_pieces_);

  Java_AiCoreSessionWrapper_generate(
      env, java_session_, reinterpret_cast<intptr_t>(this),
      java_generate_options,
      base::android::ToJavaArrayOfObjects(env, java_inputs));
  std::move(on_complete).Run();
}

void BackendSessionImplAndroid::SizeInTokens(
    on_device_model::mojom::InputPtr input,
    mojo::ReportBadMessageCallback bad_message_callback,
    base::OnceCallback<void(uint32_t)> callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(!size_in_tokens_callback_)
      << "Caller should not call SizeInTokens() again before the previous "
         "callback is invoked.";
  // Store the callback to be invoked from Java.
  size_in_tokens_callback_ = std::move(callback);

  JNIEnv* env = base::android::AttachCurrentThread();

  // Convert input pieces to Java objects.
  std::vector<base::android::ScopedJavaLocalRef<jobject>> java_inputs =
      ConvertInputPiecesToJava(env, input->pieces);

  // Call Java method to count tokens. The result will be returned via
  // OnSizeInTokensResult callback.
  Java_AiCoreSessionWrapper_getSizeInTokens(
      env, java_session_, reinterpret_cast<int64_t>(this),
      base::android::ToJavaArrayOfObjects(env, java_inputs));
}

void BackendSessionImplAndroid::Score(
    const std::string& text,
    base::OnceCallback<void(float)> callback) {
  NOTIMPLEMENTED();
  std::move(callback).Run(0.0f);
}

void BackendSessionImplAndroid::GetProbabilitiesBlocking(
    const std::string& input,
    base::OnceCallback<void(const std::vector<float>&)> callback) {
  NOTIMPLEMENTED();
  std::move(callback).Run({});
}

std::unique_ptr<BackendSession> BackendSessionImplAndroid::Clone() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  // AiCore doesn't support cloning natively yet. If it does in the future, we
  // should copy the Java object and call the native Clone function here.
  // Use `base::WrapUnique` because the constructor is private and
  // `std::make_unique` cannot access it.
  return base::WrapUnique(new BackendSessionImplAndroid(
      feature_, params_.Clone(), mojo::Clone(context_input_pieces_)));
}

void BackendSessionImplAndroid::AsrStream(
    on_device_model::mojom::AsrStreamOptionsPtr options,
    mojo::PendingRemote<on_device_model::mojom::AsrStreamResponder> responder) {
  NOTIMPLEMENTED();
}

void BackendSessionImplAndroid::AsrAddAudioChunk(
    on_device_model::mojom::AudioDataPtr data) {
  NOTIMPLEMENTED();
}

void BackendSessionImplAndroid::Hint(mojom::HintOptionsPtr options) {}

void BackendSessionImplAndroid::OnResponse(const std::string& response) {
  sequence_checker_helper_.PostTask(
      FROM_HERE,
      base::BindOnce(&BackendSessionImplAndroid::OnResponseOnSequence,
                     weak_ptr_, response));
}

void BackendSessionImplAndroid::OnResponseOnSequence(
    const std::string& response) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  auto chunk = on_device_model::mojom::ResponseChunk::New();
  chunk->text = response;
  responder_->OnResponse(std::move(chunk));
}

void BackendSessionImplAndroid::OnComplete(GenerateResult generate_result) {
  sequence_checker_helper_.PostTask(
      FROM_HERE,
      base::BindOnce(&BackendSessionImplAndroid::OnCompleteOnSequence,
                     weak_ptr_, generate_result));
}

void BackendSessionImplAndroid::OnCompleteOnSequence(
    GenerateResult generate_result) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  base::UmaHistogramEnumeration("OnDeviceModel.Android.GenerateResult",
                                generate_result);
  base::UmaHistogramEnumeration(
      base::StrCat({"OnDeviceModel.Android.GenerateResult.",
                    optimization_guide::GetStringNameForModelExecutionFeature(
                        feature_)}),
      generate_result);
  responder_->OnComplete(on_device_model::mojom::ResponseSummary::New());
  responder_.reset();
}

void BackendSessionImplAndroid::OnSizeInTokensResult(uint32_t size) {
  sequence_checker_helper_.PostTask(
      FROM_HERE,
      base::BindOnce(&BackendSessionImplAndroid::OnSizeInTokensResultOnSequence,
                     weak_ptr_, size));
}

void BackendSessionImplAndroid::OnSizeInTokensResultOnSequence(uint32_t size) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (size_in_tokens_callback_) {
    std::move(size_in_tokens_callback_).Run(size);
  }
}

static void JNI_AiCoreSessionWrapper_OnComplete(JNIEnv* env,
                                                int64_t backend_session,
                                                int32_t j_generate_result) {
  reinterpret_cast<BackendSessionImplAndroid*>(backend_session)
      ->OnComplete(static_cast<BackendSessionImplAndroid::GenerateResult>(
          j_generate_result));
}

static void JNI_AiCoreSessionWrapper_OnResponse(
    JNIEnv* env,
    int64_t backend_session,
    const jni_zero::JavaRef<jstring>& j_response) {
  reinterpret_cast<BackendSessionImplAndroid*>(backend_session)
      ->OnResponse(base::android::ConvertJavaStringToUTF8(env, j_response));
}

static void JNI_AiCoreSessionWrapper_OnSizeInTokensResult(
    JNIEnv* env,
    int64_t backend_session,
    int32_t j_token_count) {
  // j_token_count cannot be negative, but just in case, clamp it to 0 as
  // OnSizeInTokensResult expects a uint32_t value.
  reinterpret_cast<BackendSessionImplAndroid*>(backend_session)
      ->OnSizeInTokensResult(std::max(0, j_token_count));
}

}  // namespace on_device_model

DEFINE_JNI(AiCoreSessionWrapper)
DEFINE_JNI(GenerateOptionsHelper)
DEFINE_JNI(InputPieceHelper)
