From f00f8065362a1c0859a10303b15dcf8541890952 Mon Sep 17 00:00:00 2001 From: Jens Scheffler Date: Mon, 4 Mar 2024 21:50:25 +0100 Subject: [PATCH] Pass task output as outlet to dataset trigger params --- airflow/example_dags/example_datasets.py | 11 +++++++++-- airflow/jobs/scheduler_job_runner.py | 8 ++++++++ airflow/models/taskinstance.py | 13 +++++++++---- 3 files changed, 26 insertions(+), 6 deletions(-) diff --git a/airflow/example_dags/example_datasets.py b/airflow/example_dags/example_datasets.py index ac7cc2b3c1702..15586063edc4e 100644 --- a/airflow/example_dags/example_datasets.py +++ b/airflow/example_dags/example_datasets.py @@ -51,18 +51,21 @@ """ from __future__ import annotations +import random + import pendulum from airflow.datasets import Dataset from airflow.models.dag import DAG from airflow.operators.bash import BashOperator +from airflow.operators.python import PythonOperator from airflow.timetables.datasets import DatasetOrTimeSchedule from airflow.timetables.trigger import CronTriggerTimetable # [START dataset_def] dag1_dataset = Dataset("s3://dag1/output_1.txt", extra={"hi": "bye"}) # [END dataset_def] -dag2_dataset = Dataset("s3://dag2/output_1.txt", extra={"hi": "bye"}) +dag2_dataset = Dataset("s3://dag2/output_1.txt") dag3_dataset = Dataset("s3://dag3/output_3.txt", extra={"hi": "bye"}) with DAG( @@ -83,7 +86,11 @@ schedule=None, tags=["produces", "dataset-scheduled"], ) as dag2: - BashOperator(outlets=[dag2_dataset], task_id="producing_task_2", bash_command="sleep 5") + + def some_python_callable(): + return {"some_context": "dynamic data 123", "random_number": random.randint(1, 100)} + + PythonOperator(python_callable=some_python_callable, outlets=[dag2_dataset], task_id="producing_task_2") # [START dag_dep] with DAG( diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index 9a5ba78b6f65f..702efb0e6fd0a 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -1283,9 +1283,17 @@ def _create_dag_runs_dataset_triggered( events=dataset_events, ) + run_conf = {} + for item in dataset_events: + event: DatasetEvent = item + extra: dict | None = event.extra + if extra: + run_conf.update(extra) + dag_run = dag.create_dagrun( run_id=run_id, run_type=DagRunType.DATASET_TRIGGERED, + conf=run_conf, execution_date=exec_date, data_interval=data_interval, state=DagRunState.QUEUED, diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index 57c9483cd4ee7..da9772964584b 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -2374,7 +2374,7 @@ def _run_raw_task( try: if not mark_success: - self._execute_task_with_callbacks(context, test_mode, session=session) + result = self._execute_task_with_callbacks(context, test_mode, session=session) if not test_mode: self.refresh_from_db(lock_for_update=True, session=session) self.state = TaskInstanceState.SUCCESS @@ -2462,7 +2462,7 @@ def _run_raw_task( session.add(Log(self.state, self)) session.merge(self).task = self.task if self.state == TaskInstanceState.SUCCESS: - self._register_dataset_changes(session=session) + self._register_dataset_changes(result, session=session) session.commit() if self.state == TaskInstanceState.SUCCESS: @@ -2472,7 +2472,7 @@ def _run_raw_task( return None - def _register_dataset_changes(self, *, session: Session) -> None: + def _register_dataset_changes(self, result: Any, *, session: Session) -> None: for obj in self.task.outlets or []: self.log.debug("outlet obj %s", obj) # Lineage can have other types of objects besides datasets @@ -2480,10 +2480,13 @@ def _register_dataset_changes(self, *, session: Session) -> None: dataset_manager.register_dataset_change( task_instance=self, dataset=obj, + extra=obj.extra or result if isinstance(result, dict) else {self.task_id: result}, session=session, ) - def _execute_task_with_callbacks(self, context: Context, test_mode: bool = False, *, session: Session): + def _execute_task_with_callbacks( + self, context: Context, test_mode: bool = False, *, session: Session + ) -> Any: """Prepare Task for Execution.""" from airflow.models.renderedtifields import RenderedTaskInstanceFields @@ -2571,6 +2574,8 @@ def signal_handler(signum, frame): Stats.incr("operator_successes", tags={**self.stats_tags, "task_type": self.task.task_type}) Stats.incr("ti_successes", tags=self.stats_tags) + return result + def _execute_task(self, context, task_orig): """ Execute Task (optionally with a Timeout) and push Xcom results.