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

#ifndef THIRD_PARTY_BLINK_RENDERER_MODULES_AI_LANGUAGE_MODEL_H_
#define THIRD_PARTY_BLINK_RENDERER_MODULES_AI_LANGUAGE_MODEL_H_

#include <variant>

#include "base/types/pass_key.h"
#include "third_party/blink/public/mojom/ai/ai_language_model.mojom-blink.h"
#include "third_party/blink/public/mojom/ai/ai_manager.mojom-blink.h"
#include "third_party/blink/public/mojom/ai/model_streaming_responder.mojom-blink.h"
#include "third_party/blink/renderer/bindings/core/v8/idl_types.h"
#include "third_party/blink/renderer/bindings/core/v8/script_promise.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_availability.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_append_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_clone_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_create_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_message_role.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_prompt_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_language_model_sampling_mode.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_typedefs.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_union_languagemodelmessagecontentsequence_string.h"
#include "third_party/blink/renderer/core/dom/events/event_target.h"
#include "third_party/blink/renderer/core/event_type_names.h"
#include "third_party/blink/renderer/core/execution_context/execution_context_lifecycle_observer.h"
#include "third_party/blink/renderer/core/streams/readable_stream.h"
#include "third_party/blink/renderer/modules/ai/ai_utils.h"
#include "third_party/blink/renderer/modules/ai/language_model_params.h"
#include "third_party/blink/renderer/modules/modules_export.h"
#include "third_party/blink/renderer/platform/bindings/script_wrappable.h"
#include "third_party/blink/renderer/platform/mojo/heap_mojo_remote.h"

namespace blink {

// The class that represents a `LanguageModel` object.
class MODULES_EXPORT LanguageModel final : public EventTarget,
                                           public ExecutionContextClient {
  DEFINE_WRAPPERTYPEINFO();

 public:
  // Get the mojo enum value for the given V8 `role` enum value.
  static mojom::blink::AILanguageModelPromptRole ConvertRoleToMojo(
      V8LanguageModelMessageRole role);

  LanguageModel(
      ExecutionContext* execution_context,
      mojo::PendingRemote<mojom::blink::AILanguageModel> pending_remote,
      scoped_refptr<base::SequencedTaskRunner> task_runner,
      mojom::blink::AILanguageModelInstanceInfoPtr info,
      bool has_context = false);
  ~LanguageModel() override = default;

  bool has_context() const { return has_context_; }

  void Trace(Visitor* visitor) const override;

  // EventTarget implementation
  const AtomicString& InterfaceName() const override;
  ExecutionContext* GetExecutionContext() const override;

  DEFINE_ATTRIBUTE_EVENT_LISTENER(quotaoverflow, kQuotaoverflow)
  DEFINE_ATTRIBUTE_EVENT_LISTENER(contextoverflow, kContextoverflow)

  // language_model.idl implementation.
  static ScriptPromise<LanguageModel> create(
      ScriptState* script_state,
      LanguageModelCreateOptions* options,
      ExceptionState& exception_state);
  static ScriptPromise<V8Availability> availability(
      ScriptState* script_state,
      LanguageModelCreateCoreOptions* options,
      ExceptionState& exception_state);
  static ScriptPromise<IDLNullable<LanguageModelParams>> params(
      ScriptState* script_state,
      ExceptionState& exception_state);

  ScriptPromise<V8LanguageModelPromptResult> prompt(
      ScriptState* script_state,
      const V8LanguageModelPrompt* input,
      const LanguageModelPromptOptions* options,
      ExceptionState& exception_state);
  ReadableStream* promptStreaming(ScriptState* script_state,
                                  const V8LanguageModelPrompt* input,
                                  const LanguageModelPromptOptions* options,
                                  ExceptionState& exception_state);
  ScriptPromise<IDLUndefined> append(ScriptState* script_state,
                                     const V8LanguageModelPrompt* input,
                                     const LanguageModelAppendOptions* options,
                                     ExceptionState& exception_state);
  ScriptPromise<IDLDouble> measureContextUsage(
      ScriptState* script_state,
      const V8LanguageModelPrompt* input,
      const LanguageModelPromptOptions* options,
      ExceptionState& exception_state);
  ScriptPromise<IDLDouble> measureInputUsage(
      ScriptState* script_state,
      const V8LanguageModelPrompt* input,
      const LanguageModelPromptOptions* options,
      ExceptionState& exception_state);
  double contextUsage() const { return info_->input_usage; }
  double contextWindow() const { return info_->input_quota; }
  double inputQuota() const { return info_->input_quota; }
  double inputUsage() const { return info_->input_usage; }
  uint32_t topK() const { return info_->sampling_params->top_k; }
  float temperature() const { return info_->sampling_params->temperature; }
  std::optional<V8LanguageModelSamplingMode> samplingMode() const {
    return ConvertSamplingModeToV8(info_->sampling_mode);
  }

  ScriptPromise<LanguageModel> clone(ScriptState* script_state,
                                     const LanguageModelCloneOptions* options,
                                     ExceptionState& exception_state);
  void destroy(ScriptState* script_state, ExceptionState& exception_state);

  static void ExecuteAvailability(
      HeapMojoRemote<mojom::blink::AIManager>& ai_manager_remote,
      const LanguageModelCreateCoreOptions* options,
      mojom::blink::AILanguageModelSamplingParamsPtr resolved_sampling_params,
      base::OnceCallback<void(mojom::blink::ModelAvailabilityCheckResult)>
          callback);
  HeapMojoRemote<mojom::blink::AILanguageModel>& GetAILanguageModelRemote();
  scoped_refptr<base::SequencedTaskRunner> GetTaskRunner();

 private:
  void ResolvePromiseOnComplete(
      ScriptPromiseResolver<V8LanguageModelPromptResult>* resolver,
      const String& response,
      mojom::blink::ModelExecutionContextInfoPtr context_info);
  void HandleToolCalls(Vector<mojom::blink::ToolCallPtr> tool_calls);
  void OnResponseComplete(
      mojom::blink::ModelExecutionContextInfoPtr context_info);
  void OnContextOverflow();

  using ResolverOrStream =
      std::variant<ScriptPromiseResolverBase*, ReadableStream*>;
  // Helper to make AILanguageModelProxy::Prompt compatible with
  // ConvertPromptInputsToMojo callback.
  void ExecutePrompt(
      ScriptState* script_state,
      ResolverOrStream resolver_or_stream,
      on_device_model::mojom::blink::ResponseConstraintPtr constraint,
      mojo::PendingRemote<mojom::blink::ModelStreamingResponder>
          pending_responder,
      Vector<mojom::blink::AILanguageModelPromptPtr> prompts);

  // Helper to make AILanguageModelProxy::MeasureInputUsage compatible with
  // ConvertPromptInputsToMojo callback.
  void ExecuteMeasureInputUsage(
      ScriptPromiseResolver<IDLDouble>* resolver,
      AbortSignal* signal,
      Vector<mojom::blink::AILanguageModelPromptPtr> prompts);

  // Validates common prompt input and returns true on success.
  // Returns false and throws exceptions on failure.
  bool ValidateInput(ScriptState* script_state,
                     const V8LanguageModelPrompt* input,
                     AbortSignal* signal,
                     ExceptionState& exception_state);

  // Validates and processes prompt input and returns the processed constraints.
  // Returns std::nullopt on failure.
  std::optional<on_device_model::mojom::blink::ResponseConstraintPtr>
  ValidateAndProcessPromptInput(ScriptState* script_state,
                                const V8LanguageModelPrompt* input,
                                const LanguageModelPromptOptions* options,
                                ExceptionState& exception_state);

  scoped_refptr<base::SequencedTaskRunner> task_runner_;
  HeapMojoRemote<mojom::blink::AILanguageModel> language_model_remote_;
  // Info about the language model's supported types and instance state.
  blink::mojom::blink::AILanguageModelInstanceInfoPtr info_;
  // Tool calls from the current response, populated by HandleToolCalls.
  Vector<mojom::blink::ToolCallPtr> pending_tool_calls_;
  // Whether the session has any context, including pending requests.
  bool has_context_ = false;
};

}  // namespace blink

#endif  // THIRD_PARTY_BLINK_RENDERER_MODULES_AI_LANGUAGE_MODEL_H_
