From 6e50d2ad1abb1152e0f1978cb36d8a848adb4317 Mon Sep 17 00:00:00 2001 From: SameerMesiah97 <75502260+SameerMesiah97@users.noreply.github.com> Date: Mon, 27 Jul 2026 16:21:43 +0100 Subject: [PATCH] Rename StepFunctionStartExecutionOperator's input field to match its constructor argument --- .../src/airflow/jobs/scheduler_job_runner.py | 149 ++++++++++-------- .../amazon/aws/operators/step_function.py | 6 +- .../aws/operators/test_step_function.py | 2 +- .../validate_operators_init_exemptions.txt | 1 - 4 files changed, 89 insertions(+), 69 deletions(-) diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py b/airflow-core/src/airflow/jobs/scheduler_job_runner.py index a9efa1c03d62c..daef05b2b7c07 100644 --- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py +++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py @@ -543,6 +543,80 @@ def _debug_dump(self, signum: int, frame: FrameType | None) -> None: self.log.info("\n\t".join(map(repr, callstack))) self.log.info("-" * 80) + def _task_concurrency_allows_execution( + self, + *, + task_instance: TI, + concurrency_map: ConcurrencyMap, + session: Session, + starved_tasks: set[tuple[str, str]], + starved_tasks_task_dagrun_concurrency: set[tuple[str, str, str]], + ) -> bool: + """Evaluate task-level concurrency constraints for a task instance.""" + dag_id = task_instance.dag_id + task_id = task_instance.task_id + run_id = task_instance.run_id + + serialized_dag = self.scheduler_dag_bag.get_dag_for_run( + dag_run=task_instance.dag_run, + session=session, + ) + + # If the DAG is missing, fail all scheduled TIs for this DAG. + if not serialized_dag: + self.log.error( + "DAG '%s' for task instance %s not found in serialized_dag table", + dag_id, + task_instance, + ) + + session.execute( + update(TI) + .where(TI.dag_id == dag_id, TI.state == TaskInstanceState.SCHEDULED) + .values(state=TaskInstanceState.FAILED) + .execution_options(synchronize_session="fetch") + ) + + return False + + if not serialized_dag.has_task(task_id): + return True + + task = serialized_dag.get_task(task_id) + + task_concurrency_limit = task.max_active_tis_per_dag + + if task_concurrency_limit is not None: + current_task_concurrency = concurrency_map.task_concurrency_map[(dag_id, task_id)] + + if current_task_concurrency >= task_concurrency_limit: + self.log.info( + "Not executing %s since the task concurrency for this task has been reached.", + task_instance, + ) + + starved_tasks.add((dag_id, task_id)) + return False + + task_dagrun_concurrency_limit = task.max_active_tis_per_dagrun + + if task_dagrun_concurrency_limit is not None: + current_task_dagrun_concurrency = concurrency_map.task_dagrun_concurrency_map[ + (dag_id, run_id, task_id) + ] + + if current_task_dagrun_concurrency >= task_dagrun_concurrency_limit: + self.log.info( + "Not executing %s since the task concurrency per DAG run for this task has been reached.", + task_instance, + ) + + starved_tasks_task_dagrun_concurrency.add((dag_id, run_id, task_id)) + + return False + + return True + def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -> list[TI]: """ Find TIs that are ready for execution based on conditions. @@ -868,71 +942,18 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) - starved_dags.add(dag_id) continue - if task_instance.dag_model.has_task_concurrency_limits: - # Many dags don't have a task_concurrency, so where we can avoid loading the full - # serialized DAG the better. - serialized_dag = self.scheduler_dag_bag.get_dag_for_run( - dag_run=task_instance.dag_run, session=session + # Many DAGs do not define task concurrency limits, so avoid + # loading the serialized DAG unless required. + if task_instance.dag_model.has_task_concurrency_limits and not ( + self._task_concurrency_allows_execution( + task_instance=task_instance, + concurrency_map=concurrency_map, + session=session, + starved_tasks=starved_tasks, + starved_tasks_task_dagrun_concurrency=(starved_tasks_task_dagrun_concurrency), ) - # If the dag is missing, fail the task and continue to the next task. - if not serialized_dag: - self.log.error( - "DAG '%s' for task instance %s not found in serialized_dag table", - dag_id, - task_instance, - ) - session.execute( - update(TI) - .where(TI.dag_id == dag_id, TI.state == TaskInstanceState.SCHEDULED) - .values(state=TaskInstanceState.FAILED) - .execution_options(synchronize_session="fetch") - ) - continue - - task_concurrency_limit: int | None = None - if serialized_dag.has_task(task_instance.task_id): - task_concurrency_limit = serialized_dag.get_task( - task_instance.task_id - ).max_active_tis_per_dag - - if task_concurrency_limit is not None: - current_task_concurrency = concurrency_map.task_concurrency_map[ - (task_instance.dag_id, task_instance.task_id) - ] - - if current_task_concurrency >= task_concurrency_limit: - self.log.info( - "Not executing %s since the task concurrency for this task has been reached.", - task_instance, - ) - starved_tasks.add((task_instance.dag_id, task_instance.task_id)) - continue - - task_dagrun_concurrency_limit: int | None = None - if serialized_dag.has_task(task_instance.task_id): - task_dagrun_concurrency_limit = serialized_dag.get_task( - task_instance.task_id - ).max_active_tis_per_dagrun - - if task_dagrun_concurrency_limit is not None: - current_task_dagrun_concurrency = concurrency_map.task_dagrun_concurrency_map[ - (task_instance.dag_id, task_instance.run_id, task_instance.task_id) - ] - - if current_task_dagrun_concurrency >= task_dagrun_concurrency_limit: - self.log.info( - "Not executing %s since the task concurrency per DAG run for" - " this task has been reached.", - task_instance, - ) - starved_tasks_task_dagrun_concurrency.add( - ( - task_instance.dag_id, - task_instance.run_id, - task_instance.task_id, - ) - ) - continue + ): + continue if executor_obj := self._try_to_load_executor( task_instance, session, team_name=dag_id_to_team_name.get(task_instance.dag_id, NOTSET) diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py b/providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py index 7c5c6a31cb91c..4ae4367329046 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py @@ -76,7 +76,7 @@ class StepFunctionStartExecutionOperator(AwsBaseOperator[StepFunctionHook]): aws_hook_class = StepFunctionHook template_fields: Sequence[str] = aws_template_fields( - "state_machine_arn", "name", "input", "is_redrive_execution" + "state_machine_arn", "name", "state_machine_input", "is_redrive_execution" ) ui_color = "#f9c915" operator_extra_links = (StateMachineDetailsLink(), StateMachineExecutionsDetailsLink()) @@ -97,7 +97,7 @@ def __init__( self.state_machine_arn = state_machine_arn self.name = name self.is_redrive_execution = is_redrive_execution - self.input = state_machine_input + self.state_machine_input = state_machine_input self.waiter_delay = waiter_delay self.waiter_max_attempts = waiter_max_attempts self.deferrable = deferrable @@ -113,7 +113,7 @@ def execute(self, context: Context): if not ( execution_arn := self.hook.start_execution( - self.state_machine_arn, self.name, self.input, self.is_redrive_execution + self.state_machine_arn, self.name, self.state_machine_input, self.is_redrive_execution ) ): raise AirflowException(f"Failed to start State Machine execution for: {self.state_machine_arn}") diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_step_function.py b/providers/amazon/tests/unit/amazon/aws/operators/test_step_function.py index 55a0913405f9f..9d58968a70279 100644 --- a/providers/amazon/tests/unit/amazon/aws/operators/test_step_function.py +++ b/providers/amazon/tests/unit/amazon/aws/operators/test_step_function.py @@ -144,7 +144,7 @@ def test_init(self): assert op.state_machine_arn == STATE_MACHINE_ARN assert op.state_machine_arn == STATE_MACHINE_ARN assert op.name == NAME - assert op.input == INPUT + assert op.state_machine_input == INPUT assert op.hook.aws_conn_id == AWS_CONN_ID assert op.hook._region_name == REGION_NAME assert op.hook._verify is False diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 86ae9cd475947..818dc48cb84bb 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -12,7 +12,6 @@ providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneS providers/amazon/src/airflow/providers/amazon/aws/operators/s3.py::S3DeleteObjectsOperator providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py::SageMakerCreateNotebookOperator providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py::SageMakerProcessingOperator -providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py::StepFunctionStartExecutionOperator providers/amazon/src/airflow/providers/amazon/aws/transfers/gcs_to_s3.py::GCSToS3Operator providers/amazon/src/airflow/providers/amazon/aws/transfers/s3_to_redshift.py::S3ToRedshiftOperator providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator