// 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/webnn/webnn_context_impl.h"

#include <memory>
#include <set>
#include <utility>

#include "base/atomic_sequence_num.h"
#include "base/files/file_util.h"
#include "base/functional/callback_helpers.h"
#include "base/logging.h"
#include "base/metrics/histogram_functions.h"
#include "base/sequence_checker.h"
#include "base/strings/stringprintf.h"
#include "base/task/bind_post_task.h"
#include "base/task/thread_pool.h"
#include "base/trace_event/memory_dump_manager.h"
#include "gpu/command_buffer/service/shared_image/shared_image_manager.h"
#include "services/webnn/error.h"
#include "services/webnn/gpu_task_scheduler.h"
#include "services/webnn/public/cpp/data_type_limits.h"
#include "services/webnn/public/cpp/graph_validation_utils.h"
#include "services/webnn/public/cpp/ml_tensor_usage.h"
#include "services/webnn/public/cpp/operand_descriptor.h"
#include "services/webnn/public/cpp/supported_data_types.h"
#include "services/webnn/public/cpp/supported_tensors.h"
#include "services/webnn/public/cpp/webnn_trace.h"
#include "services/webnn/public/mojom/webnn_context.mojom.h"
#include "services/webnn/public/mojom/webnn_context_provider.mojom.h"
#include "services/webnn/public/mojom/webnn_error.mojom.h"
#include "services/webnn/public/mojom/webnn_graph.mojom.h"
#include "services/webnn/public/mojom/webnn_graph_builder.mojom.h"
#include "services/webnn/public/mojom/webnn_tensor.mojom.h"
#include "services/webnn/webnn_context_provider_impl.h"
#include "services/webnn/webnn_context_provider_in_renderer.h"
#include "services/webnn/webnn_tensor_impl.h"
#include "third_party/tflite/buildflags.h"

#if BUILDFLAG(BUILD_TFLITE_WITH_XNNPACK)
#include "third_party/xnnpack/src/include/xnnpack.h"  // nogncheck
#endif  // BUILD_TFLITE_WITH_XNNPACK

namespace {
// Generates process-unique IDs to use for tracing resources.
base::AtomicSequenceNumber g_next_webnn_context_tracing_id;

// Return false if the named tensors for dispatch don't match the built
// graph's expectation.
bool ValidateWebNNTensors(
    const base::flat_map<std::string, scoped_refptr<webnn::WebNNTensorImpl>>&
        named_tensors,
    const base::flat_map<std::string, webnn::OperandDescriptor>&
        names_to_descriptors) {
  return std::ranges::equal(
      named_tensors, names_to_descriptors,
      [](const auto& named_tensor, const auto& tensor_spec) {
        const auto& [tensor_name, tensor_impl] = named_tensor;
        const auto& [tensor_spec_name, tensor_spec_descriptor] = tensor_spec;
        return tensor_name == tensor_spec_name &&
               tensor_impl->data_type() == tensor_spec_descriptor.data_type() &&
               tensor_impl->shape() == tensor_spec_descriptor.shape();
      });
}

// Return false if the same tensor was specified in inputs and outputs.
bool ValidateWebNNTensorsUsage(
    const base::flat_map<std::string, blink::WebNNTensorToken>& named_inputs,
    const base::flat_map<std::string, blink::WebNNTensorToken>& named_outputs) {
  // Validate that output tensors are unique.
  std::set<blink::WebNNTensorToken> output_tensors;
  for (const auto& named_output : named_outputs) {
    output_tensors.insert(named_output.second);
  }

  if (output_tensors.size() != named_outputs.size()) {
    return false;
  }

  // Validate tensors used for input and output are unique.
  for (const auto& named_input : named_inputs) {
    if (output_tensors.contains(named_input.second)) {
      return false;
    }
  }

  return true;
}

}  // namespace

namespace webnn {

WebNNContextImpl::WebNNContextImpl(
    mojo::PendingReceiver<mojom::WebNNContext> receiver,
    base::WeakPtr<WebNNContextProviderImpl> context_provider,
    WebNNContextImpl::ContextBackendUma backend_uma,
    ContextProperties properties,
    mojom::CreateContextOptionsPtr options,
    mojo::ScopedDataPipeConsumerHandle write_tensor_consumer,
    mojo::ScopedDataPipeProducerHandle read_tensor_producer,
    std::unique_ptr<GpuTaskScheduler> gpu_task_scheduler,
    scoped_refptr<gpu::MemoryTracker> memory_tracker,
    scoped_refptr<base::SingleThreadTaskRunner> owning_task_runner,
    gpu::SharedImageManager* shared_image_manager,
    scoped_refptr<base::SingleThreadTaskRunner> main_task_runner)
    : WebNNObjectBase<mojom::WebNNContext,
                      blink::WebNNContextToken,
                      mojo::Receiver<mojom::WebNNContext>>(
          std::move(receiver),
          gpu_task_scheduler->scheduler_task_runner()),
      has_context_provider_(true),
      context_provider_(std::move(context_provider)),
      properties_(IntersectWithBaseProperties(std::move(properties))),
      options_(std::move(options)),
      memory_type_tracker_(std::move(memory_tracker)),
      gpu_task_scheduler_(std::move(gpu_task_scheduler)),
      write_tensor_consumer_(std::move(write_tensor_consumer)),
      read_tensor_producer_(std::move(read_tensor_producer)),
      shared_image_manager_(shared_image_manager),
      main_task_runner_(std::move(main_task_runner)),
      owning_task_runner_(std::move(owning_task_runner)),
      tracing_id_(g_next_webnn_context_tracing_id.GetNext()) {
  InitializeContext(backend_uma);
}

WebNNContextImpl::WebNNContextImpl(
    mojo::PendingReceiver<mojom::WebNNContext> receiver,
    base::WeakPtr<WebNNContextProviderInRenderer> context_provider_in_renderer,
    WebNNContextImpl::ContextBackendUma backend_uma,
    ContextProperties properties,
    mojom::CreateContextOptionsPtr options,
    scoped_refptr<base::SingleThreadTaskRunner> owning_task_runner,
    scoped_refptr<base::SingleThreadTaskRunner> main_task_runner)
    : WebNNObjectBase<mojom::WebNNContext,
                      blink::WebNNContextToken,
                      mojo::Receiver<mojom::WebNNContext>>(std::move(receiver),
                                                           owning_task_runner),
      context_provider_in_renderer_(std::move(context_provider_in_renderer)),
      is_context_provider_in_renderer_(true),
      properties_(IntersectWithBaseProperties(std::move(properties))),
      options_(std::move(options)),
      memory_type_tracker_(base::MakeRefCounted<gpu::MemoryTracker>()),
      main_task_runner_(std::move(main_task_runner)),
      owning_task_runner_(std::move(owning_task_runner)),
      tracing_id_(g_next_webnn_context_tracing_id.GetNext()) {
  InitializeContext(backend_uma);
}

void WebNNContextImpl::InitializeContext(ContextBackendUma backend_uma) {
  RecordContextBackendUma(backend_uma);
#if BUILDFLAG(BUILD_TFLITE_WITH_XNNPACK)
  const xnn_status status = xnn_initialize(/*allocator=*/nullptr);
  if (status == xnn_status_success) {
    is_xnnpack_initialized_ = true;
  } else {
    LOG(WARNING) << "Failed to initialize XNNPACK (status=" << status
                 << "); falling back to built-in TFLite kernels for CPU "
                    "inference.";
  }
#endif  // BUILDFLAG(BUILD_TFLITE_WITH_XNNPACK)
  base::trace_event::MemoryDumpManager::GetInstance()->RegisterDumpProvider(
      this, "WebNN", owning_task_runner_);
}

// static
base::RepeatingClosure* g_destruction_callback_for_testing = nullptr;

WebNNContextImpl::~WebNNContextImpl() {
  CHECK(!has_graph_builders())
      << "Graph builders must be cleared in OnDisconnect().";

  for (auto impl : tensor_impls_) {
    // Non-interop tensors require manual tracking cleanup since the memory
    // is owned by the context and unlike interop, cannot be released by shared
    // image.
    if (impl->has_shared_image()) {
      impl->DestroyAccessAndRepresentationAndWait();
    } else {
      memory_type_tracker_.TrackMemFree(impl->PackedByteLength());
    }
  }

  base::trace_event::MemoryDumpManager::GetInstance()->UnregisterDumpProvider(
      this);

#if BUILDFLAG(BUILD_TFLITE_WITH_XNNPACK)
  // Deinitialize XNNPACK only if it was successfully initialized; otherwise
  // calling other XNNPACK APIs is unsafe.
  if (is_xnnpack_initialized_) {
    const xnn_status status = xnn_deinitialize();
    CHECK_EQ(status, xnn_status_success);
  }
#endif  // BUILDFLAG(BUILD_TFLITE_WITH_XNNPACK)

  // Destroy the GPU task scheduler before signaling destruction. This releases
  // any pending MultiplexRouter refs held by queued tasks (via WrapRefCounted).
  // The callback must fire after this to guarantee the router is fully cleaned
  // up.
  gpu_task_scheduler_.reset();

  // Sequence destruction must happen on the provider main sequence, and only
  // after this context has fully torn down (including scheduler shutdown).
  if (has_context_provider_) {
    main_task_runner_->PostTask(
        FROM_HERE,
        base::BindOnce(&WebNNContextProviderImpl::DestroyAndRemoveGpuSequence,
                       context_provider_, handle()));
  }

  if (g_destruction_callback_for_testing) {
    g_destruction_callback_for_testing->Run();
  }
}

// static
void WebNNContextImpl::SetDestructionCallbackForTesting(  // IN-TEST
    base::RepeatingClosure* callback) {
  g_destruction_callback_for_testing = callback;
}

// static
void WebNNContextImpl::RecordContextBackendUma(ContextBackendUma backend_uma) {
  base::UmaHistogramEnumeration("WebNN.Context.Backend", backend_uma);
}

void WebNNContextImpl::OnDisconnect() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  // Explicitly reset all tensor and graph receivers before destruction since
  // destroying bound receivers can cause Mojo to DCHECK due to pending
  // callbacks or if destruction occurs on a different runner than the bound
  // runner.
  for (auto impl : tensor_impls_) {
    impl->ResetMojoReceiver();
  }


  // Close the primary pipe before clearing builders. Closing the pipe detaches
  // all endpoint clients on the router, so the subsequent Clear() won't trigger
  // MaybePostToProcessTasks() and leak a router ref.
  ResetMojoReceiver();
  ClearGraphBuilders();

  base::OnceClosure remove_task;
  if (is_context_provider_in_renderer_) {
    remove_task =
        base::BindOnce(&WebNNContextProviderInRenderer::RemoveWebNNContextImpl,
                       context_provider_in_renderer_, handle());
  }

  if (!remove_task) {
    remove_task =
        base::BindOnce(&WebNNContextProviderImpl::RemoveWebNNContextImpl,
                       context_provider_, handle());
  }

  main_task_runner_->PostTask(FROM_HERE, std::move(remove_task));
}

void WebNNContextImpl::ReportBadMessageAndDisconnect(std::string_view message) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  GetMojoReceiver().ReportBadMessage(message);
  OnDisconnect();
}

#if BUILDFLAG(IS_WIN)
void WebNNContextImpl::DestroyAllContextsAndKillGpuProcess() {
  if (!main_task_runner_->RunsTasksInCurrentSequence()) {
    main_task_runner_->PostTask(
        FROM_HERE,
        base::BindOnce(
            &WebNNContextProviderImpl::DestroyAllContextsAndKillGpuProcess,
            context_provider_));
    return;
  }

  context_provider_->DestroyAllContextsAndKillGpuProcess();
}
#endif  // BUILDFLAG(IS_WIN)

void WebNNContextImpl::AddGraphImpl(scoped_refptr<WebNNGraphImpl> graph_impl) {
  graph_impls_.emplace(std::move(graph_impl));
}

void WebNNContextImpl::OpenWeightsFile(
    base::OnceCallback<void(base::File,
                            mojo::PendingRemote<mojom::WeightsFileSession>)>
        callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(is_context_provider_in_renderer_);

  base::OnceClosure task =
      base::BindOnce(&WebNNContextProviderInRenderer::OpenWeightsFile,
                     context_provider_in_renderer_,
                     base::BindPostTaskToCurrentDefault(std::move(callback)));
  if (!main_task_runner_->RunsTasksInCurrentSequence()) {
    main_task_runner_->PostTask(FROM_HERE, std::move(task));
  } else {
    std::move(task).Run();
  }
}

void WebNNContextImpl::CreateWeightsFile(
    base::OnceCallback<void(base::File)> callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  // CreateWeightsFile is only used by the GPU-process path.
  DCHECK(!is_context_provider_in_renderer_);

  base::OnceClosure create_task = base::BindOnce(
      &WebNNContextProviderImpl::CreateWeightsFile, context_provider_,
      base::BindPostTaskToCurrentDefault(std::move(callback)));

  if (!main_task_runner_->RunsTasksInCurrentSequence()) {
    main_task_runner_->PostTask(FROM_HERE, std::move(create_task));
  } else {
    std::move(create_task).Run();
  }
}

void WebNNContextImpl::BuildGraph(
    mojom::GraphInfoPtr graph_info,
    WebNNGraphImpl::ComputeResourceInfo compute_resource_info,
    base::flat_map<OperandId, std::unique_ptr<WebNNConstantOperand>>
        constant_operands,
    BuildGraphCallback callback) {
  CreateGraphImpl(std::move(graph_info), std::move(compute_resource_info),
                  std::move(constant_operands),
                  base::BindOnce(&WebNNContextImpl::OnGraphBuilt, AsWeakPtr(),
                                 std::move(callback)));
}

void WebNNContextImpl::OnGraphBuilt(
    BuildGraphCallback callback,
    base::expected<scoped_refptr<WebNNGraphImpl>, mojom::ErrorPtr> result) {
  if (!result.has_value()) {
    std::move(callback).Run(base::unexpected(std::move(result.error())));
    return;
  }

  GraphCreationResult creation_result(result.value()->handle(),
                                      result.value()->devices());
  graph_impls_.emplace(std::move(result.value()));

  std::move(callback).Run(std::move(creation_result));
}

void WebNNContextImpl::CreateGraphBuilder(
    mojo::PendingReceiver<mojom::WebNNGraphBuilder> receiver) {
  CreateGraphBuilderImpl(std::move(receiver));
}

void WebNNContextImpl::CreateTensor(
    mojom::TensorInfoPtr tensor_info,
    mojom::WebNNContext::CreateTensorCallback callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  ScopedTrace scoped_trace("WebNNContextImpl::CreateTensor");

  if (!ValidateTensor(properties_, tensor_info->descriptor).has_value()) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  mojo::PendingAssociatedRemote<mojom::WebNNTensor> remote;
  auto receiver = remote.InitWithNewEndpointAndPassReceiver();

  auto result = CreateTensorImpl(std::move(receiver), std::move(tensor_info));
  if (!result.has_value()) {
    std::move(callback).Run(
        mojom::CreateTensorResult::NewError(std::move(result.error())));
    return;
  }

  auto success = mojom::CreateTensorSuccess::New(std::move(remote),
                                                 result.value()->handle());
  std::move(callback).Run(
      mojom::CreateTensorResult::NewSuccess(std::move(success)));

  memory_type_tracker_.TrackMemAlloc(result.value()->PackedByteLength());

  // Associates a `WebNNTensor` instance with this context so the WebNN service
  // can access the implementation.
  tensor_impls_.emplace(*std::move(result));
}

GpuTaskScheduler* WebNNContextImpl::gpu_task_scheduler() const {
  return gpu_task_scheduler_.get();
}

bool WebNNContextImpl::HasValidWriteTensorConsumer() const {
  return write_tensor_consumer_.is_valid();
}

bool WebNNContextImpl::HasValidReadTensorProducer() const {
  return read_tensor_producer_.is_valid();
}

void WebNNContextImpl::ReadDataFromBigBufferOrDataPipe(
    mojo_base::BigBuffer src_buffer,
    base::span<uint8_t> dst_span) {
  if (src_buffer.size() == 0) {
    CHECK(write_tensor_consumer_);
    size_t bytes_read = 0;
    if (write_tensor_consumer_->ReadData(MOJO_READ_DATA_FLAG_ALL_OR_NONE,
                                         dst_span,
                                         bytes_read) != MOJO_RESULT_OK) {
      OnLost("WriteTensor(): Failed to read tensor data from data pipe.");
    }
  } else {
    dst_span.copy_from(src_buffer);
  }
}

mojo_base::BigBuffer WebNNContextImpl::WriteDataToDataPipeOrBigBuffer(
    base::span<const uint8_t> src_span) {
  if (read_tensor_producer_ &&
      src_span.size() > mojo_base::BigBuffer::kMaxInlineBytes &&
      read_tensor_producer_->WriteAllData(src_span) == MOJO_RESULT_OK) {
    return mojo_base::BigBuffer();
  }
  return mojo_base::BigBuffer(src_span);
}

void WebNNContextImpl::CreateTensorFromMailbox(mojom::TensorInfoPtr tensor_info,
                                               const gpu::Mailbox& mailbox,
                                               const gpu::SyncToken& fence,
                                               CreateTensorCallback callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  ScopedTrace scoped_trace("WebNNContextImpl::CreateTensorFromMailbox");

  if (!tensor_info->usage.Has(MLTensorUsageFlags::kWebGpuInterop)) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  if (!ValidateTensor(properties_, tensor_info->descriptor).has_value()) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  // SharedImageManager is not available when running without GPU
  // dependencies. WebGPU interop requires GPU process resources.
  if (!shared_image_manager_) {
    std::move(callback).Run(ToError<mojom::CreateTensorResult>(
        mojom::Error::Code::kNotSupportedError,
        "WebGPU interop is not supported in this context."));
    return;
  }

  if (!gpu_task_scheduler_) {
    std::move(callback).Run(ToError<mojom::CreateTensorResult>(
        mojom::Error::Code::kNotSupportedError,
        "WebGPU interop is not supported without a GPU sequence."));
    return;
  }

  // Ensure the Mojo callback is posted back to the task runner. Running
  // it directly on the GPU sequence can violate Mojo's sequence checks,
  // even if executing on the same thread.
  auto mojo_callback_wrapper =
      base::BindPostTask(mojo_task_runner(), std::move(callback));

  // Must be a scheduled task since this depends on shared image creation task.
  RunOrScheduleTaskWithThisContext(
      base::BindOnce(
          [](mojom::TensorInfoPtr tensor_info, const gpu::Mailbox& mailbox,
             CreateTensorCallback callback, ScopedTrace scoped_trace,
             WebNNContextImpl& self) {
            CHECK(self.shared_image_manager_);

            constexpr char kWebNNCreateTensorErrorMessage[] =
                "Failed to create tensor.";

            // Tensor will own the representation.
            // TODO(https://crbug.com/481747252): When SharedImageBacking memory
            // tracking is fixed memory tracking for interop should work.
            WebNNTensorImpl::RepresentationPtr representation(
                self.shared_image_manager_
                    ->ProduceWebNNTensor(mailbox, &self.memory_type_tracker_)
                    .release(),
                WebNNTensorImpl::OnTaskRunnerDeleterWithWait(
                    self.main_task_runner()));
            if (!representation) {
              std::move(callback).Run(ToError<mojom::CreateTensorResult>(
                  mojom::Error::Code::kUnknownError,
                  kWebNNCreateTensorErrorMessage));
              return;
            }

            mojo::PendingAssociatedRemote<mojom::WebNNTensor> remote;
            auto receiver = remote.InitWithNewEndpointAndPassReceiver();

            auto result = self.CreateTensorFromSharedImageImpl(
                std::move(receiver), std::move(tensor_info),
                std::move(representation));
            if (!result.has_value()) {
              std::move(callback).Run(mojom::CreateTensorResult::NewError(
                  std::move(result.error())));
              return;
            }

            if (!result.value()->ImportTensorInternal()) {
              std::move(callback).Run(ToError<mojom::CreateTensorResult>(
                  mojom::Error::Code::kUnknownError,
                  kWebNNCreateTensorErrorMessage));
              return;
            }

            auto success = mojom::CreateTensorSuccess::New(
                std::move(remote), result.value()->handle());
            std::move(callback).Run(
                mojom::CreateTensorResult::NewSuccess(std::move(success)));
            self.tensor_impls_.emplace(*std::move(result));
          },
          std::move(tensor_info), mailbox, std::move(mojo_callback_wrapper),
          std::move(scoped_trace)),
      fence);
}

void WebNNContextImpl::Dispatch(
    const blink::WebNNGraphToken& graph_token,
    const base::flat_map<std::string, blink::WebNNTensorToken>& named_inputs,
    const base::flat_map<std::string, blink::WebNNTensorToken>& named_outputs) {
  ScopedTrace scoped_trace("WebNNContextImpl::Dispatch");

  if (!ValidateWebNNTensorsUsage(named_inputs, named_outputs)) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  // Resolve graph token to graph impl.
  auto graph_it = graph_impls_.find(graph_token);
  if (graph_it == graph_impls_.end()) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidGraph);
    return;
  }
  scoped_refptr<WebNNGraphImpl> graph_impl = *graph_it;

  // Resolve the token of an input MLTensor to the corresponding `WebNNTensor`
  // instance.
  std::vector<std::pair<std::string, scoped_refptr<WebNNTensorImpl>>>
      name_to_input_tensors;
  name_to_input_tensors.reserve(named_inputs.size());
  for (const auto& [name, tensor_handle] : named_inputs) {
    scoped_refptr<WebNNTensorImpl> input_tensor =
        GetWebNNTensorImpl(tensor_handle);
    if (!input_tensor) {
      return;
    }

    name_to_input_tensors.emplace_back(name, std::move(input_tensor));
  }
  base::flat_map<std::string, scoped_refptr<WebNNTensorImpl>>
      name_to_input_tensor_map(std::move(name_to_input_tensors));
  if (!ValidateWebNNTensors(
          name_to_input_tensor_map,
          graph_impl->compute_resource_info().input_names_to_descriptors)) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  // Resolve the token of an output MLTensor to the corresponding `WebNNTensor`
  // instance.
  std::vector<std::pair<std::string, scoped_refptr<WebNNTensorImpl>>>
      name_to_output_tensors;
  name_to_output_tensors.reserve(named_outputs.size());
  for (const auto& [name, tensor_handle] : named_outputs) {
    scoped_refptr<WebNNTensorImpl> output_tensor =
        GetWebNNTensorImpl(tensor_handle);
    if (!output_tensor) {
      return;
    }

    name_to_output_tensors.emplace_back(name, std::move(output_tensor));
  }

  base::flat_map<std::string, scoped_refptr<WebNNTensorImpl>>
      name_to_output_tensor_map(std::move(name_to_output_tensors));
  if (!ValidateWebNNTensors(
          name_to_output_tensor_map,
          graph_impl->compute_resource_info().output_names_to_descriptors)) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return;
  }

  graph_impl->RunDispatch(
      std::move(name_to_input_tensor_map), std::move(name_to_output_tensor_map),
      std::move(scoped_trace), GetMojoReceiver().GetBadMessageCallback());
}

void WebNNContextImpl::DestroyGraph(
    const blink::WebNNGraphToken& graph_handle) {
  auto it = graph_impls_.find(graph_handle);
  if (it == graph_impls_.end()) {
    GetMojoReceiver().ReportBadMessage(kBadMessageInvalidGraph);
    return;
  }
  graph_impls_.erase(it);
}

void WebNNContextImpl::RequestCompilerContext(
    mojo::PendingReceiver<mojom::WebNNCompilerContext>
        compiler_context_receiver) {
  // The base class drops the receiver (pipe disconnects). Only
  // DispatchContextImplOrt overrides this with a real implementation that
  // reconnects to the Compiler process.
  LOG(WARNING) << "[WebNN] RequestCompilerContext() is not implemented for "
                  "this context.";
}

void WebNNContextImpl::RemoveWebNNTensorImpl(
    const blink::WebNNTensorToken& handle) {
  const auto it = tensor_impls_.find(handle);
  CHECK(it != tensor_impls_.end());
  if (it->get()->has_shared_image()) {
    it->get()->DestroyAccessAndRepresentationAndWait();
  } else {
    memory_type_tracker_.TrackMemFree(it->get()->PackedByteLength());
  }
  // Upon calling erase, the handle will no longer refer to a valid
  // `WebNNTensorImpl`.
  tensor_impls_.erase(it);
}

const ContextProperties& WebNNContextImpl::properties() const {
  return properties_;
}

const mojom::CreateContextOptions& WebNNContextImpl::options() const {
  return *options_;
}

void WebNNContextImpl::OnLost(const std::string& reason) {
  if (!mojo_task_runner()->RunsTasksInCurrentSequence()) {
    mojo_task_runner()->PostTask(
        FROM_HERE,
        base::BindOnce(&WebNNContextImpl::OnLost, AsWeakPtr(), reason));
    return;
  }

  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  ResetMojoReceiver(reason);
  OnDisconnect();
}

void WebNNContextImpl::RunOrScheduleTaskWithThisContext(
    RunOrScheduleTaskCallback task,
    const gpu::SyncToken& fence) {
  // Safe to use std::ref because `this` owns gpu_task_scheduler_ and
  // its deletion drops all pending tasks before the context is destroyed.
  RunOrScheduleTask(base::BindOnce(std::move(task), std::ref(*this)), fence);
}

void WebNNContextImpl::RunOrScheduleTask(base::OnceClosure task,
                                         const gpu::SyncToken& fence,
                                         const gpu::SyncToken& release) {
  if (gpu_task_scheduler_) {
    gpu_task_scheduler_->ScheduleGpuTask(std::move(task), fence, release);
    return;
  }

  DCHECK(!fence.HasData());
  DCHECK(!release.HasData());
  DCHECK(owning_task_runner()->RunsTasksInCurrentSequence());
  std::move(task).Run();
}

scoped_refptr<WebNNTensorImpl> WebNNContextImpl::GetWebNNTensorImpl(
    const blink::WebNNTensorToken& tensor_handle) {
  const auto it = tensor_impls_.find(tensor_handle);
  if (it == tensor_impls_.end()) {
    ReportBadMessageAndDisconnect(kBadMessageInvalidTensor);
    return nullptr;
  }
  return it->get();
}

bool WebNNContextImpl::OnMemoryDump(
    const base::trace_event::MemoryDumpArgs& args,
    base::trace_event::ProcessMemoryDump* pmd) {
  std::string dump_name = base::StringPrintf("webnn/context_0x%x", tracing_id_);
  auto* const dump = pmd->CreateAllocatorDump(dump_name);
  dump->AddScalar(base::trace_event::MemoryAllocatorDump::kNameSize,
                  base::trace_event::MemoryAllocatorDump::kUnitsBytes,
                  memory_type_tracker_.memory_tracker()->GetSize());
  return true;
}

ContextProperties WebNNContextImpl::IntersectWithBaseProperties(
    ContextProperties backend_context_properties) {
  // A specific maximum rank is still under discussion, but 8 is the highest
  // supported by any backend.
  constexpr SupportedRanks kNonScalarMaxRank = SupportedRanks::NonScalarUpTo(8);
  constexpr SupportedRanks kAtLeast2D{2, 8};

  // Only intersects for ones that have limits defined in the specification.
  // For ones that has no limit, no need to intersect with
  // `SupportedDataTypes::All()`.
  backend_context_properties.data_type_limits.arg_min_max_input.ranks
      .IntersectWith(kNonScalarMaxRank);
  backend_context_properties.data_type_limits.arg_min_max_output.data_types
      .RetainAll(DataTypeConstraint::kInt32To64);
  backend_context_properties.data_type_limits.batch_normalization_input
      .IntersectWith({DataTypeConstraint::kFloat16To32, kNonScalarMaxRank});
  backend_context_properties.data_type_limits.batch_normalization_mean
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.concat_inputs.ranks.IntersectWith(
      kNonScalarMaxRank);
  backend_context_properties.data_type_limits.conv2d_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.conv2d_bias.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.conv_transpose2d_input
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.conv_transpose2d_bias
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.cumulative_sum_input
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32Ints32To64, kNonScalarMaxRank});
  backend_context_properties.data_type_limits.dequantize_linear_input.data_types
      .RetainAll(DataTypeConstraint::kInts4Ints8Ints32);
  backend_context_properties.data_type_limits.dequantize_linear_scale.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.logical_and_input.data_types
      .RetainAll(DataTypeConstraint::kUint8);
  backend_context_properties.data_type_limits.logical_or_input.data_types
      .RetainAll(DataTypeConstraint::kUint8);
  backend_context_properties.data_type_limits.logical_xor_input.data_types
      .RetainAll(DataTypeConstraint::kUint8);
  backend_context_properties.data_type_limits.logical_not_input.data_types
      .RetainAll(DataTypeConstraint::kUint8);
  backend_context_properties.data_type_limits.is_nan_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.is_infinite_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.logical_output.RetainAll(
      DataTypeConstraint::kUint8);
  backend_context_properties.data_type_limits.abs_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32Int8To64);
  backend_context_properties.data_type_limits.ceil_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.cos_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.erf_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.exp_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.floor_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.log_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.neg_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32Int8To64);
  backend_context_properties.data_type_limits.reciprocal_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.round_even_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.sign_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32Int8To64);
  backend_context_properties.data_type_limits.sin_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.sqrt_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.tan_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.elu_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.gather_input.ranks.IntersectWith(
      kNonScalarMaxRank);
  backend_context_properties.data_type_limits.gather_indices.data_types
      .RetainAll(DataTypeConstraint::kGatherScatterIndicesSupportedDataTypes);
  backend_context_properties.data_type_limits.gather_elements_input.ranks
      .IntersectWith(kNonScalarMaxRank);
  backend_context_properties.data_type_limits.gather_elements_indices
      .IntersectWith(
          {DataTypeConstraint::kGatherScatterIndicesSupportedDataTypes,
           kNonScalarMaxRank});
  backend_context_properties.data_type_limits.gather_nd_input.ranks
      .IntersectWith(kNonScalarMaxRank);
  backend_context_properties.data_type_limits.gather_nd_indices.IntersectWith(
      {DataTypeConstraint::kGatherScatterIndicesSupportedDataTypes,
       kNonScalarMaxRank});
  backend_context_properties.data_type_limits.gelu_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.gemm_a.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(2)});
  backend_context_properties.data_type_limits.gemm_c.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::UpTo(2)});
  backend_context_properties.data_type_limits.gru_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(3)});
  backend_context_properties.data_type_limits.gru_bias.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(2)});
  backend_context_properties.data_type_limits.gru_output_sequence.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.gru_cell_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(2)});
  backend_context_properties.data_type_limits.gru_cell_bias.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.hard_sigmoid_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.hard_swish_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.instance_normalization_input
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.instance_normalization_scale
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.layer_normalization_input
      .data_types.RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.leaky_relu_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.linear_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.lstm_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(3)});
  backend_context_properties.data_type_limits.lstm_bias.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(2)});
  backend_context_properties.data_type_limits.lstm_output_sequence
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.lstm_cell_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(2)});
  backend_context_properties.data_type_limits.lstm_cell_bias.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(1)});
  backend_context_properties.data_type_limits.matmul_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, kAtLeast2D});
  backend_context_properties.data_type_limits.average_pool2d_input
      .IntersectWith(
          {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.l2_pool2d_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.max_pool2d_input.ranks
      .IntersectWith(SupportedRanks::Exactly(4));
  backend_context_properties.data_type_limits.prelu_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32Int8To64);
  backend_context_properties.data_type_limits.quantize_linear_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.quantize_linear_zero_point
      .data_types.RetainAll(DataTypeConstraint::kInts4Ints8Ints32);
  backend_context_properties.data_type_limits.reduce_l1_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32Ints32To64);
  backend_context_properties.data_type_limits.reduce_l2_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.reduce_log_sum_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.reduce_log_sum_exp_input
      .data_types.RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.reduce_mean_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.reduce_product_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32Ints32To64);
  backend_context_properties.data_type_limits.reduce_sum_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32Ints32To64);
  backend_context_properties.data_type_limits.reduce_sum_square_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32Ints32To64);
  backend_context_properties.data_type_limits.relu_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32Int8To64);
  backend_context_properties.data_type_limits.resample2d_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, SupportedRanks::Exactly(4)});
  backend_context_properties.data_type_limits.scatter_elements_input.ranks
      .IntersectWith(kNonScalarMaxRank);
  backend_context_properties.data_type_limits.scatter_elements_indices
      .data_types.RetainAll(
          DataTypeConstraint::kGatherScatterIndicesSupportedDataTypes);
  backend_context_properties.data_type_limits.scatter_nd_input.ranks
      .IntersectWith(kNonScalarMaxRank);
  backend_context_properties.data_type_limits.scatter_nd_indices.IntersectWith(
      {DataTypeConstraint::kGatherScatterIndicesSupportedDataTypes,
       kNonScalarMaxRank});
  backend_context_properties.data_type_limits.sigmoid_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.softmax_input.IntersectWith(
      {DataTypeConstraint::kFloat16To32, kNonScalarMaxRank});
  backend_context_properties.data_type_limits.softplus_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.softsign_input.data_types
      .RetainAll(DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.split_input.ranks.IntersectWith(
      kNonScalarMaxRank);
  backend_context_properties.data_type_limits.tanh_input.data_types.RetainAll(
      DataTypeConstraint::kFloat16To32);
  backend_context_properties.data_type_limits.triangular_input.ranks
      .IntersectWith(kAtLeast2D);
  backend_context_properties.data_type_limits.where_condition.data_types
      .RetainAll(DataTypeConstraint::kUint8);
  return backend_context_properties;
}

}  // namespace webnn
