# -------------------------------------------------------------------------
#  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 List
from billiard.exceptions import SoftTimeLimitExceeded
from celery import Task

from atk.configs.results_config import TaskResult, RunStatus, ProcessStatus
from atk.db.dao.case_result_dao import CaseResultDAO
from atk.tasks.celery_config import db_pool
from atk.common.log import Logger

DEFAULT_MAX_RETRIES = 1  # 重试次数
DEFAULT_VISIBILITY_TIMEOUT = 3600 * 24  # 控制任务在队列中被取出后的可见性超时。单位是秒

logging = Logger().get_logger()


class BaseCeleryTask(Task):
    celery_clean_output_data_task = None

    def __init__(self):
        super().__init__()
        self.max_retries = DEFAULT_MAX_RETRIES
        self.visibility_timeout = DEFAULT_VISIBILITY_TIMEOUT
        self.case_result_dao = CaseResultDAO(db_pool)

    @staticmethod
    def init_task():
        if not BaseCeleryTask.celery_clean_output_data_task:
            from atk.tasks.celery_tasks import celery_clean_output_data
            BaseCeleryTask.celery_clean_output_data_task = celery_clean_output_data

    def get_task_result(self, args):
        if args is None:
            raise ValueError(f'{self.name} args is None')
        if isinstance(args, List):
            return self.get_task_result(args[0])
        if isinstance(args, dict):
            if 'case_config' in args.keys():
                return TaskResult(**args)
            else:
                raise ValueError(f'{self.name} args have no case_config, args: {args}')
        else:
            raise ValueError(f'get_task_result args type is invalid: {type(args)}')

    def after_return(self, status, retval, task_id, args, kwargs, einfo):
        if 'celery_post_process' not in self.name:
            logging.debug('skip after_return')
            return
        task_result = self.get_task_result(args)
        if task_result is None or task_result.case_config is None or task_result.nodes is None:
            return
        BaseCeleryTask.init_task()
        for node_info in task_result.nodes.get_execute_nodes():
            if not node_info.task:
                continue
            task_result.add_node(node_info)
            task = BaseCeleryTask.celery_clean_output_data_task
            try:
                task(task_result.model_dump())
            except Exception as e:
                logging.warning(f"after_return clean data failed, err {str(e)}")

    def on_failure(self, exc, task_id, args, kwargs, einfo):
        task_result = self.get_task_result(args)
        if task_result is None or task_result.case_config is None:
            return
        msg = (f'{task_result.case_config.name}_{task_result.case_config.id} on {self.name} '
               f'{self.request.delivery_info["routing_key"]} task failed.\n{exc}\n{str(einfo)}')
        logging.error(msg)
        if isinstance(exc, SoftTimeLimitExceeded):
            task_result.run_status = RunStatus.TIMEOUT
        else:
            task_result.run_status = RunStatus.FAILED
        task_result.failed_message += msg
        self.case_result_dao.save_dataset(
            case_id=task_result.case_config.id,
            case_name=task_result.case_config.name,
            run_status=RunStatus.FAILED,
            task_result=task_result,
            save_db_status=ProcessStatus.FINISHED)

    def on_success(self, retval, task_id, args, kwargs):

        if 'celery_post_process' in self.name:
            task_result = TaskResult(**retval)
            logging.debug(f"start on success, task: {self.name}, case id: {task_result.case_config.id}")
            task_result.run_status = RunStatus.SUCCESS
            self.case_result_dao.save_dataset(
                case_id=task_result.case_config.id,
                case_name=task_result.case_config.name,
                run_status=RunStatus.SUCCESS,
                task_result=task_result,
                save_db_status=ProcessStatus.FINISHED)
        else:
            task_result = self.get_task_result(args)
            task_result.run_status = RunStatus.PROCESSING
            self.case_result_dao.save_dataset(
                case_id=task_result.case_config.id,
                case_name=task_result.case_config.name,
                run_status=RunStatus.PROCESSING,
                task_result=task_result,
                save_db_status=ProcessStatus.UNDO)