/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */

#include "sampleModel.h"
#include "utils.h"

aclmdlDataset *AclModelOutput::GetDataSet() {
  return dataset_;
}

aclmdlDataset *AclModelInput::GetDataSet() {
  return dataset_;
}

aclmdlDataset *AclLlmModelInput::GetDataSet() {
  return dataset_;
}

size_t AclModelWork::GetModelWorkSize() {
  return modelWorkSize_;
}

void *AclModelWork::GetModelWorkPtr() {
  return modelWorkPtr_;
}

size_t AclModelWeight::GetModelWeightSize() {
  return modelWeightSize_;
}

void *AclModelWeight::GetModelWeightPrt() {
  return modelWeightPtr_;
}

aclmdlDesc *AclModelDesc::GetModelDesc() {
  return modelDesc_;
}

AclModelInput::AclModelInput(void *inputDataBuffer, size_t bufferSize, aclmdlDesc *modelDesc) {
  CHECK_NOT_NULL(modelDesc);
  uint32_t dataNum = aclmdlGetNumInputs(modelDesc);
  dataset_ = aclmdlCreateDataset();
  CHECK_NOT_NULL(dataset_);
  aclDataBuffer *inputData = aclCreateDataBuffer(inputDataBuffer, bufferSize);
  CHECK_NOT_NULL(inputData);
  CHECK(aclmdlAddDatasetBuffer(dataset_, inputData));

  size_t dynamicIdx = 0;
  auto ret = aclmdlGetInputIndexByName(modelDesc, ACL_DYNAMIC_TENSOR_NAME, &dynamicIdx);
  if ((ret == ACL_SUCCESS) && (dynamicIdx == (dataNum - 1))) {
    size_t dataLen = aclmdlGetInputSizeByIndex(modelDesc, dynamicIdx);
    void *data = nullptr;
    CHECK(aclrtMalloc(&data, dataLen, ACL_MEM_MALLOC_HUGE_FIRST));
    aclDataBuffer *dataBuf = aclCreateDataBuffer(data, dataLen);
    CHECK_NOT_NULL(dataBuf);
    CHECK(aclmdlAddDatasetBuffer(dataset_, dataBuf));
  }
}

AclModelInput::~AclModelInput() {
  CHECK_NOT_NULL(dataset_);

  for (size_t i = 0; i < aclmdlGetDatasetNumBuffers(dataset_); ++i) {
    aclDataBuffer *dataBuffer = aclmdlGetDatasetBuffer(dataset_, i);
    CHECK(aclDestroyDataBuffer(dataBuffer));
  }
  CHECK(aclmdlDestroyDataset(dataset_));
  dataset_ = nullptr;
}

AclLlmModelInput::AclLlmModelInput(const std::vector<InputTensor> &inputs, aclmdlDesc *modelDesc) {
  dataset_ = aclmdlCreateDataset();
  CHECK_NOT_NULL(dataset_);

  size_t num_inputs = aclmdlGetNumInputs(modelDesc);
  for (size_t i = 0; i < num_inputs; ++i) {
    size_t expected_size = aclmdlGetInputSizeByIndex(modelDesc, i);
    void *device_ptr = nullptr;
    CHECK(aclrtMalloc(&device_ptr, expected_size, ACL_MEM_MALLOC_HUGE_FIRST));
    CHECK(aclrtMemset(device_ptr, expected_size, 0, expected_size));
    if (i < inputs.size()) {
      const auto &input = inputs[i];
      if (input.byte_size > expected_size) {
        ERROR_LOG("Input %zu: host data size (%zu) > model expected size (%zu)", i, input.byte_size, expected_size);
        aclrtFree(device_ptr);
      }
      CHECK(aclrtMemcpy(device_ptr, expected_size, input.data, input.byte_size, ACL_MEMCPY_HOST_TO_DEVICE));
    }

    aclDataBuffer *buffer = aclCreateDataBuffer(device_ptr, expected_size);
    CHECK(aclmdlAddDatasetBuffer(dataset_, buffer));
  }
}

AclLlmModelInput::~AclLlmModelInput() {
  CHECK_NOT_NULL(dataset_);

  for (size_t i = 0; i < aclmdlGetDatasetNumBuffers(dataset_); ++i) {
    aclDataBuffer *dataBuffer = aclmdlGetDatasetBuffer(dataset_, i);
    void *data = aclGetDataBufferAddr(dataBuffer);
    CHECK(aclrtFree(data));
    CHECK(aclDestroyDataBuffer(dataBuffer));
  }
  CHECK(aclmdlDestroyDataset(dataset_));
  dataset_ = nullptr;
}

AclModelOutput::AclModelOutput(aclmdlDesc *modelDesc) {
  CHECK_NOT_NULL(modelDesc);

  dataset_ = aclmdlCreateDataset();
  CHECK_NOT_NULL(dataset_);

  size_t outputSize = aclmdlGetNumOutputs(modelDesc);
  for (size_t i = 0; i < outputSize; ++i) {
    size_t modelOutputSize = aclmdlGetOutputSizeByIndex(modelDesc, i);
    void *outputBuffer = nullptr;
    CHECK(aclrtMalloc(&outputBuffer, modelOutputSize, ACL_MEM_MALLOC_HUGE_FIRST));
    aclDataBuffer *outputData = aclCreateDataBuffer(outputBuffer, modelOutputSize);
    CHECK_NOT_NULL(outputData);
    CHECK(aclmdlAddDatasetBuffer(dataset_, outputData));
  }
}

AclModelOutput::~AclModelOutput() {
  if (dataset_ == nullptr) {
    return;
  }
  for (size_t i = 0; i < aclmdlGetDatasetNumBuffers(dataset_); ++i) {
    aclDataBuffer *dataBuffer = aclmdlGetDatasetBuffer(dataset_, i);
    void *data = aclGetDataBufferAddr(dataBuffer);
    CHECK(aclrtFree(data));
    CHECK(aclDestroyDataBuffer(dataBuffer));
  }
  CHECK(aclmdlDestroyDataset(dataset_));
  dataset_ = nullptr;
}