// 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 SERVICES_WEBNN_COREML_TENSOR_IMPL_COREML_H_
#define SERVICES_WEBNN_COREML_TENSOR_IMPL_COREML_H_

#include "base/sequence_checker.h"
#include "base/thread_annotations.h"
#include "base/types/pass_key.h"
#include "gpu/command_buffer/service/shared_image/shared_image_representation.h"
#include "mojo/public/cpp/bindings/pending_associated_receiver.h"
#include "services/webnn/public/mojom/webnn_tensor.mojom.h"
#include "services/webnn/queueable_resource_state.h"
#include "services/webnn/webnn_tensor_impl.h"

namespace webnn {

class WebNNContextImpl;

namespace coreml {

class BufferContent;

class TensorImplCoreml final : public WebNNTensorImpl {
 public:
  static base::expected<scoped_refptr<WebNNTensorImpl>, mojom::ErrorPtr> Create(
      mojo::PendingAssociatedReceiver<mojom::WebNNTensor> receiver,
      WebNNContextImpl& context,
      mojom::TensorInfoPtr tensor_info);

  static base::expected<scoped_refptr<WebNNTensorImpl>, mojom::ErrorPtr> Create(
      mojo::PendingAssociatedReceiver<mojom::WebNNTensor> receiver,
      WebNNContextImpl& context,
      mojom::TensorInfoPtr tensor_info,
      RepresentationPtr representation);

  TensorImplCoreml(
      mojo::PendingAssociatedReceiver<mojom::WebNNTensor> receiver,
      WebNNContextImpl& context,
      mojom::TensorInfoPtr tensor_info,
      scoped_refptr<QueueableResourceState<BufferContent>> buffer_state,
      RepresentationPtr representation,
      base::PassKey<TensorImplCoreml> pass_key);

  TensorImplCoreml(const TensorImplCoreml&) = delete;
  TensorImplCoreml& operator=(const TensorImplCoreml&) = delete;

  // WebNNTensorImpl:
  void ReadTensorImpl(mojom::WebNNTensor::ReadTensorCallback callback) override;
  void WriteTensorImpl(mojo_base::BigBuffer src_buffer) override;
  bool ImportTensorImpl(ScopedAccessPtr access) override;
  void ExportTensorImpl(ScopedAccessPtr access) override;

  const scoped_refptr<QueueableResourceState<BufferContent>>& GetBufferState()
      const;

  // mojom::WebNNTensor
  void ExportTensor(uint64_t flow_id, uint64_t release_count) override;
  void ExportTensorSync(uint64_t flow_id,
                        uint64_t release_count,
                        ExportTensorSyncCallback callback) override;

 private:
  ~TensorImplCoreml() override;

  scoped_refptr<QueueableResourceState<BufferContent>> buffer_state_
      GUARDED_BY_CONTEXT(sequence_checker_);
};

}  // namespace coreml

}  // namespace webnn

#endif  // SERVICES_WEBNN_COREML_TENSOR_IMPL_COREML_H_
