diff --git a/airflow/executors/base_executor.py b/airflow/executors/base_executor.py index 73f7f57edd1f6..7fcbd0642e14d 100644 --- a/airflow/executors/base_executor.py +++ b/airflow/executors/base_executor.py @@ -49,6 +49,9 @@ # Tuple of: state, info EventBufferValueType = Tuple[Optional[str], Any] +# Task tuple to send to be executed +TaskTuple = Tuple[TaskInstanceKey, CommandType, Optional[str], Optional[Any]] + class BaseExecutor(LoggingMixin): """ @@ -186,6 +189,7 @@ def trigger_tasks(self, open_slots: int) -> None: :param open_slots: Number of open slots """ sorted_queue = self.order_queued_tasks_by_priority() + task_tuples = [] for _ in range(min((open_slots, len(self.queued_tasks)))): key, (command, _, queue, ti) = sorted_queue.pop(0) @@ -212,9 +216,16 @@ def trigger_tasks(self, open_slots: int) -> None: del self.attempts[key] del self.queued_tasks[key] else: - del self.queued_tasks[key] - self.running.add(key) - self.execute_async(key=key, command=command, queue=queue, executor_config=ti.executor_config) + task_tuples.append((key, command, queue, ti.executor_config)) + + if task_tuples: + self._process_tasks(task_tuples) + + def _process_tasks(self, task_tuples: List[TaskTuple]) -> None: + for key, command, queue, executor_config in task_tuples: + del self.queued_tasks[key] + self.execute_async(key=key, command=command, queue=queue, executor_config=executor_config) + self.running.add(key) def change_state(self, key: TaskInstanceKey, state: str, info=None) -> None: """ diff --git a/airflow/executors/celery_executor.py b/airflow/executors/celery_executor.py index 6099042a53704..752244d5528c6 100644 --- a/airflow/executors/celery_executor.py +++ b/airflow/executors/celery_executor.py @@ -29,7 +29,7 @@ import subprocess import time import traceback -from collections import OrderedDict +from collections import Counter, OrderedDict from concurrent.futures import ProcessPoolExecutor from multiprocessing import cpu_count from typing import Any, Dict, List, Mapping, MutableMapping, Optional, Set, Tuple, Union @@ -45,7 +45,7 @@ from airflow.config_templates.default_celery import DEFAULT_CELERY_CONFIG from airflow.configuration import conf from airflow.exceptions import AirflowException, AirflowTaskTimeout -from airflow.executors.base_executor import BaseExecutor, CommandType, EventBufferValueType +from airflow.executors.base_executor import BaseExecutor, CommandType, EventBufferValueType, TaskTuple from airflow.models.taskinstance import TaskInstance, TaskInstanceKey from airflow.stats import Stats from airflow.utils.log.logging_mixin import LoggingMixin @@ -235,7 +235,7 @@ def __init__(self): self.task_adoption_timeout = datetime.timedelta( seconds=conf.getint('celery', 'task_adoption_timeout', fallback=600) ) - self.task_publish_retries: Dict[TaskInstanceKey, int] = OrderedDict() + self.task_publish_retries: Counter[TaskInstanceKey] = Counter() self.task_publish_max_retries = conf.getint('celery', 'task_publish_max_retries', fallback=3) def start(self) -> None: @@ -250,28 +250,8 @@ def _num_tasks_per_send_process(self, to_send_count: int) -> int: """ return max(1, int(math.ceil(1.0 * to_send_count / self._sync_parallelism))) - def trigger_tasks(self, open_slots: int) -> None: - """ - Overwrite trigger_tasks function from BaseExecutor - - :param open_slots: Number of open slots - :return: - """ - sorted_queue = self.order_queued_tasks_by_priority() - - task_tuples_to_send: List[TaskInstanceInCelery] = [] - - for _ in range(min(open_slots, len(self.queued_tasks))): - key, (command, _, queue, _) = sorted_queue.pop(0) - task_tuple = (key, command, queue, execute_command) - task_tuples_to_send.append(task_tuple) - if key not in self.task_publish_retries: - self.task_publish_retries[key] = 1 - - if task_tuples_to_send: - self._process_tasks(task_tuples_to_send) - - def _process_tasks(self, task_tuples_to_send: List[TaskInstanceInCelery]) -> None: + def _process_tasks(self, task_tuples: List[TaskTuple]) -> None: + task_tuples_to_send = [task_tuple[:3] + (execute_command,) for task_tuple in task_tuples] first_task = next(t[3] for t in task_tuples_to_send) # Celery state queries will stuck if we do not use one same backend @@ -285,20 +265,19 @@ def _process_tasks(self, task_tuples_to_send: List[TaskInstanceInCelery]) -> Non if isinstance(result, ExceptionWithTraceback) and isinstance( result.exception, AirflowTaskTimeout ): - if key in self.task_publish_retries and ( - self.task_publish_retries.get(key) <= self.task_publish_max_retries - ): + retries = self.task_publish_retries[key] + if retries < self.task_publish_max_retries: Stats.incr("celery.task_timeout_error") self.log.info( "[Try %s of %s] Task Timeout Error for Task: (%s).", - self.task_publish_retries[key], + self.task_publish_retries[key] + 1, self.task_publish_max_retries, key, ) - self.task_publish_retries[key] += 1 + self.task_publish_retries[key] = retries + 1 continue self.queued_tasks.pop(key) - self.task_publish_retries.pop(key) + self.task_publish_retries.pop(key, None) if isinstance(result, ExceptionWithTraceback): self.log.error(CELERY_SEND_ERR_MSG_HEADER + ": %s\n%s\n", result.exception, result.traceback) self.event_buffer[key] = (State.FAILED, None) @@ -430,16 +409,6 @@ def end(self, synchronous: bool = False) -> None: time.sleep(5) self.sync() - def execute_async( - self, - key: TaskInstanceKey, - command: CommandType, - queue: Optional[str] = None, - executor_config: Optional[Any] = None, - ): - """Do not allow async execution for Celery executor.""" - raise AirflowException("No Async execution for Celery executor.") - def terminate(self): pass diff --git a/tests/executors/test_celery_executor.py b/tests/executors/test_celery_executor.py index 78a55ef79817b..d5bd3dbeae7a3 100644 --- a/tests/executors/test_celery_executor.py +++ b/tests/executors/test_celery_executor.py @@ -225,19 +225,19 @@ def test_retry_on_error_sending_task(self, caplog): # Test that when heartbeat is called again, task is published again to Celery Queue executor.heartbeat() - assert dict(executor.task_publish_retries) == {key: 2} + assert dict(executor.task_publish_retries) == {key: 1} assert 1 == len(executor.queued_tasks), "Task should remain in queue" assert executor.event_buffer == {} assert f"[Try 1 of 3] Task Timeout Error for Task: ({key})." in caplog.text executor.heartbeat() - assert dict(executor.task_publish_retries) == {key: 3} + assert dict(executor.task_publish_retries) == {key: 2} assert 1 == len(executor.queued_tasks), "Task should remain in queue" assert executor.event_buffer == {} assert f"[Try 2 of 3] Task Timeout Error for Task: ({key})." in caplog.text executor.heartbeat() - assert dict(executor.task_publish_retries) == {key: 4} + assert dict(executor.task_publish_retries) == {key: 3} assert 1 == len(executor.queued_tasks), "Task should remain in queue" assert executor.event_buffer == {} assert f"[Try 3 of 3] Task Timeout Error for Task: ({key})." in caplog.text diff --git a/tests/jobs/test_local_task_job.py b/tests/jobs/test_local_task_job.py index 34cfc263c872c..7eb044ec8bac6 100644 --- a/tests/jobs/test_local_task_job.py +++ b/tests/jobs/test_local_task_job.py @@ -755,7 +755,7 @@ def test_fast_follow( scheduler_job.processor_agent.end() @conf_vars({('scheduler', 'schedule_after_task_execution'): 'True'}) - def test_mini_scheduler_works_with_wait_for_upstream(self, caplog, dag_maker): + def test_mini_scheduler_works_with_wait_for_downstream(self, caplog, dag_maker): session = settings.Session() with dag_maker(default_args={'wait_for_downstream': True}, catchup=False) as dag: task_a = PythonOperator(task_id='A', python_callable=lambda: True) @@ -786,13 +786,17 @@ def test_mini_scheduler_works_with_wait_for_upstream(self, caplog, dag_maker): job1 = LocalTaskJob(task_instance=ti2_a, ignore_ti_state=True, executor=SequentialExecutor()) job1.task_runner = StandardTaskRunner(job1) + t = time.time() job1.run() + d = time.time() - t ti2_a.refresh_from_db(session) ti2_b.refresh_from_db(session) assert ti2_a.state == State.SUCCESS assert ti2_b.state == State.NONE - assert "0 downstream tasks scheduled from follow-on schedule" in caplog.text + assert ( + "0 downstream tasks scheduled from follow-on schedule" in caplog.text + ), f"Failed after {d.total_seconds()}: {caplog.text}" failed_deps = list(ti2_b.get_failed_dep_statuses(session=session)) assert len(failed_deps) == 1