# -------------------------------------------------------------------------
#  This file is part of the MindStudio project.
# Copyright (c) 2025 Huawei Technologies Co.,Ltd.
#
# MindStudio is licensed under Mulan PSL v2.
# You can use this software according to the terms and conditions of the Mulan PSL v2.
# You may obtain a copy of Mulan PSL v2 at:
#
#          http://license.coscl.org.cn/MulanPSL2
#
# 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 FIT FOR A PARTICULAR PURPOSE.
# See the Mulan PSL v2 for more details.
# -------------------------------------------------------------------------

from typing import Optional, List, Any
import os
import sys
import torch
import numpy as np
from pydantic import BaseModel
from packaging import version

from atk.common.log import Logger
from atk.common.file_check import safe_file_open
from atk.configs.results_config import TaskResult

logging = Logger().get_logger()

SUPPORT_GRAD_TYPE = [
    torch.float,
    torch.float16,
    torch.float32,
    torch.float64,
    torch.double,
    torch.bfloat16,
    torch.half,
    torch.complex64,
    torch.complex128,
]
if version.parse(torch.__version__) >= version.parse("2.3"):
    UINT_TO_INT = {
        torch.uint64: torch.int64,
        torch.uint32: torch.int32,
        torch.uint16: torch.int16,
    }
if version.parse(torch.__version__) >= version.parse("1.12.0"):
    if hasattr(torch, "complex32"):
        SUPPORT_GRAD_TYPE.append(torch.complex32)


class InputDataset(BaseModel):
    args: Optional[List[object]] = []
    kwargs: Optional[dict] = {}
    method_args: Optional[List[object]] = []
    method_kwargs: Optional[dict] = {}
    tensor_args: Optional[object] = None
    require_grad: Optional[bool] = False
    backward_args: Optional[List[bool]] = []
    backward_kwargs: Optional[dict] = {}
    backward_tensor_args: Optional[bool] = None

    @staticmethod
    def to_screen(data: Any, item: str = None, name: str = None):
        logging.info(f"--- print {item} for backend: {name} ---")
        InputDataset._recursive_txt_writer(sys.stdout, data)

    @staticmethod
    def to_txt_file(txt_file_path: str, data: Any):
        try:
            with safe_file_open(txt_file_path, 'w', encoding='utf-8') as f:
                f.write(f"--- [Export:txt] Data Structure Map for {os.path.basename(txt_file_path)} ---\n\n")
                InputDataset._recursive_txt_writer(f, data)
            logging.debug(f"Successfully exported data to TXT: {txt_file_path}")
        except Exception as e:
            logging.error(f"Failed to export data to TXT file {txt_file_path}: {e}")

    @staticmethod
    def _get_upper_type(data_type, is_bm_task):
        """
        获取升精度后的精度类型
        """
        if data_type in [torch.float16, torch.bfloat16]:
            return torch.float32
        elif data_type == torch.float32:
            return torch.float64 if not is_bm_task else torch.float32
        elif hasattr(torch, 'complex32') and data_type == torch.complex32:
            return torch.complex64
        elif hasattr(torch, 'complex64') and data_type == torch.complex64:
            return torch.complex128 if not is_bm_task else torch.complex64
        logging.debug(f"data_type: {data_type} not in [fp16, bf16, fp32], skip it.")
        return None

    @staticmethod
    def _recursive_txt_writer(f, data, indent_level: int = 0, path: str = "root"):
        """
        完整打印Tensor数据,顶层容器不打印自身信息。
        """
        indent = "    " * indent_level

        if isinstance(data, torch.Tensor):
            f.write(f"{indent}[TENSOR] Path: {path}\n")
            f.write(f"{indent}-Dtype: {data.dtype}\n")
            f.write(f"{indent}-Shape: {list(data.shape)}\n")
            f.write(f"{indent}-Device: {data.device}\n")
            f.write(f"{indent}-Data:\n")

            numpy_data = data.detach().cpu().numpy()
            with np.printoptions(threshold=np.inf):
                tensor_str = np.array2string(numpy_data, separator=', ')
                for line in tensor_str.split('\n'):
                    f.write(f"{indent} {line}\n")
            f.write("\n")

        elif isinstance(data, (list, tuple)):
            # 只有在非顶层时,才打印容器自身的信息
            if indent_level > 0:
                f.write(f"{indent}[STRUCT] Path: {path} | Type: {type(data).__name__} | Size: {len(data)}\n")

            for i, item in enumerate(data):
                InputDataset._recursive_txt_writer(f, item, indent_level + 1, path=f"{path}[{i}]")

            # 如果容器为空,且不是顶层容器,可以加个提示
            if not data and indent_level > 0:
                f.write(f"{indent}  (empty)\n\n")

        elif isinstance(data, dict):
            # 只有在非顶层时,才打印容器自身的信息
            if indent_level > 0:
                f.write(f"{indent}[STRUCT] Path: {path} | Type: {type(data).__name__} | Size: {len(data)}\n")

            for key, value in data.items():
                InputDataset._recursive_txt_writer(f, value, indent_level + 1, path=f"{path}['{key}']")

            if not data and indent_level > 0:
                f.write(f"{indent}  (empty)\n\n")

        # 基本类型的处理逻辑保持不变
        else:
            f.write(f"{indent}[PRIMITIVE] Path: {path} | Type: {type(data).__name__} | Value: {data}\n\n")

    @staticmethod
    def _save_torch_with_check(data, file_path: str, export_formats: Optional[List[str]] = None,
                               data_item_key: str = None, backend_name: str = None):
        """
        保存包含Tensor的Python对象,并根据export_format执行额外操作。
        平台默认会生成二进制文件,此处的逻辑是决定是否要生成额外文件或打印。

        Args:
            data (Any): 需要保存的数据。
            file_path (str): 保存二进制文件的路径 (例如 'input.bin')。
            export_formats (Optional[str]): 导出格式 。
        """
        if os.path.exists(file_path):
            logging.warning(f"File already exists, skipping save for: {file_path}")
        else:
            try:
                # torch.save不支持uint64类型,因此需要将其转换为int64
                if version.parse(torch.__version__) >= version.parse("2.3"):
                    InputDataset._data_check_uint(data)
                torch.save(data, file_path, pickle_protocol=4)
            except Exception as e:
                logging.error(f"Failed to save binary file {file_path}: {e}")
                return

        try:
            if 'txt' in export_formats:
                base_name, _ = os.path.splitext(file_path)
                InputDataset.to_txt_file(base_name + ".txt", data)

        except Exception as e:
            logging.error(f"Failed during export operation (format: {export_formats}) for {file_path}: {e}")

    @staticmethod
    def _data_check_uint(data):
        if isinstance(data, list):
            for i, d in enumerate(data):
                if isinstance(d, torch.Tensor) and d.dtype in UINT_TO_INT:
                    data[i] = d.to(UINT_TO_INT[d.dtype])
        elif isinstance(data, dict):
            for k, item in data.items():
                if isinstance(item, torch.Tensor) and item.dtype in UINT_TO_INT:
                    data[k] = item.to(UINT_TO_INT[item.dtype])

    def save_input_data(self, file_path: str, other_path: str = None,
                        task_result: TaskResult = None, item: str = "input"):
        """
        保存输入数据到文件。
        """
        export_formats = []
        if task_result.export_config:
            export_formats = task_result.export_config.get(item, [])
        api_type = task_result.case_config.api_type
        # 将获取到的格式传递给保存函数
        name = task_result.get_backend_name()
        if self.args:
            self._save_torch_with_check(self.args, file_path, export_formats, item, name)
        if self.kwargs:
            self._save_torch_with_check(self.kwargs, file_path, export_formats, item, name)
        if "tensor" in api_type:
            self._save_torch_with_check(self.tensor_args, other_path, export_formats, item, name)
        elif "method" in api_type:
            self._save_torch_with_check(self.method_args, other_path, export_formats, item, name)

        logging.debug(f"Saved primary input to: {file_path}")
        if other_path:
            logging.debug(f"Saved other input to: {other_path}")

    def to_device(self, device: str):
        if not self.args and not self.kwargs:
            logging.debug("input data is None")
        else:
            self.args = self._set_input_device(self.args, device)
            self.kwargs = self._set_input_device(self.kwargs, device)
        if self.method_args or self.method_kwargs:
            self.method_args = self._set_input_device(self.method_args, device)
            self.method_kwargs = self._set_input_device(self.method_kwargs, device)
        if self.tensor_args is not None:
            self.tensor_args = self._set_input_device(self.tensor_args, device)

    def set_grad(self):
        """
        检查是否有输入需要设置require_grad=True,
        是则针对各个输入的backward设置require_grad, 否则报错
        """
        self.require_grad = False
        # RuntimeError: a leaf Variable that requires grad is being used in
        # an in-place operation.
        if self.tensor_args is not None:
            self._set_grad_data(self.tensor_args, self.backward_tensor_args)
        self._set_grad_data(self.args, self.backward_args)
        self._set_grad_data(self.kwargs, self.backward_kwargs)
        if not self.require_grad:
            logging.warning("not input support set require grad, so op backward can't execute")
            raise ValueError('set input grad error.')

    def get_grad(self):
        grad_data = []
        self._get_grad_data(self.tensor_args, grad_data)
        self._get_grad_data(self.args, grad_data)
        self._get_grad_data(self.kwargs, grad_data)
        return grad_data

    def clear_grad(self):
        self._clear_grad_data(self.tensor_args)
        self._clear_grad_data(self.args)
        self._clear_grad_data(self.kwargs)

    def get_benchmark_data(self, is_bm_task=False):
        """
        将所有数据转换为更高精度
        """
        if not self.args and not self.kwargs:
            logging.warning("input data is None")
        else:
            self.args = self._set_input_upper(self.args, is_bm_task)
            self.kwargs = self._set_input_upper(self.kwargs, is_bm_task)
        if self.method_args or self.method_kwargs:
            self.method_args = self._set_input_upper(self.method_args, is_bm_task)
            self.method_kwargs = self._set_input_upper(self.method_kwargs, is_bm_task)
        if self.tensor_args is not None:
            self.tensor_args = self._set_input_upper(self.tensor_args, is_bm_task)

    def _set_input_data_by_list_and_tuple(self, input_data, device):
        device_data_list = []
        for data in input_data:
            device_data = self._set_input_device(data, device)
            device_data_list.append(device_data)
        if isinstance(input_data, list):
            return device_data_list
        else:
            return tuple(device_data_list)

    def _set_input_data_by_dict(self, input_data, device):
        for name, data in input_data.items():
            input_data[name] = self._set_input_device(data, device)
        return input_data

    def _set_input_device(self, data, device):

        def _set_input_data_by_tensor(input_data, device):
            if not device:
                raise ValueError('InputDataset not set device.')
            return input_data.to(device)

        if isinstance(data, list) or isinstance(data, tuple):
            return self._set_input_data_by_list_and_tuple(data, device)
        if isinstance(data, dict):
            return self._set_input_data_by_dict(data, device)
        if isinstance(data, torch.Tensor):
            return _set_input_data_by_tensor(data, device)
        return data

    def _set_grad_data(self, input_data, backward):
        """
        :input_data: 输入数据 Union[Tensor,List[Tensor],Dict[str, Tensor]]
        :backward: 输入数据的backward信息 Union[bool,List[bool],Dict[str, bool]]
        return: 设置grad后的输入数据 Union[Tensor,List[Tensor],Dict[str, Tensor]]
        """
        if isinstance(input_data, (list, tuple)):
            if isinstance(backward, (list, tuple)) and len(input_data) == len(backward):
                for i, data in enumerate(input_data):
                    self._set_grad_data(data, backward[i])
            elif isinstance(backward, bool):
                for _, data in enumerate(input_data):
                    self._set_grad_data(data, backward)
            else:
                raise ValueError(f"input and backward mismatch: backward: {backward}, inputdata: {input_data}")
        elif isinstance(input_data, dict):
            if isinstance(backward, dict) and set(input_data.keys()) == set(backward.keys()):
                for key, value in input_data.items():
                    self._set_grad_data(value, backward.get(key))
            elif isinstance(backward, bool):
                for _, value in input_data.items():
                    self._set_grad_data(value, backward)
            else:
                raise ValueError(f"input and backward mismatch: backward: {backward}, inputdata: {input_data}")
        elif isinstance(input_data, torch.Tensor):
            if input_data.dtype not in SUPPORT_GRAD_TYPE:
                logging.warning(f"The dtype of tensor is {input_data.dtype}, don't support require gradients,"
                                f"only Tensors of floating point and complex dtype can require gradients")
            else:
                input_data.requires_grad = backward
                self.require_grad = True
        else:
            logging.debug(f'input_data type: {type(input_data)} is not support, skip it')

    def _get_grad_data(self, input_data, grad_data):
        if isinstance(input_data, list) or isinstance(input_data, tuple):
            for data in input_data:
                self._get_grad_data(data, grad_data)
        elif isinstance(input_data, dict):
            for value in input_data.values():
                self._get_grad_data(value, grad_data)
        elif isinstance(input_data, torch.Tensor):
            if input_data.requires_grad and input_data.grad is not None:
                grad_data.append(input_data.grad)

    def _clear_grad_data(self, input_data):
        """
        清理输入数据中所有张量的梯度。

        参数:
        input_data: 可以是列表、元组、字典或单个的torch.Tensor。
        """
        if isinstance(input_data, list) or isinstance(input_data, tuple):
            for data in input_data:
                self._clear_grad_data(data)
        elif isinstance(input_data, dict):
            for value in input_data.values():
                self._clear_grad_data(value)
        elif isinstance(input_data, torch.Tensor):
            if input_data.requires_grad and input_data.grad is not None:
                logging.debug("clear grad success!")
                input_data.grad.zero_()

    def _set_input_upper(self, data, is_bm_task):
        """
        对不同类型的数据递归处理
        """
        if isinstance(data, (list, tuple)):
            data_list = []
            for value in data:
                data_list.append(self._set_input_upper(value, is_bm_task))
            return tuple(data_list) if isinstance(data, tuple) else data_list
        if isinstance(data, dict):
            data_dict = {}
            for name, value in data.items():
                data_dict[name] = self._set_input_upper(value, is_bm_task)
            return data_dict
        if isinstance(data, torch.Tensor):
            data_type = InputDataset._get_upper_type(data.dtype, is_bm_task)
            if data_type is not None:
                data = data.to(data_type)
        return data