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)