From 6a8927044813edf215b01ca703b7f80e1844f470 Mon Sep 17 00:00:00 2001 From: Dylan Storey Date: Wed, 5 Apr 2023 09:20:19 -0400 Subject: [PATCH 1/5] readding after borked it --- airflow/operators/trigger_dagrun.py | 11 +++--- tests/operators/test_trigger_dagrun.py | 48 ++++++++++++++++++++++---- 2 files changed, 49 insertions(+), 10 deletions(-) diff --git a/airflow/operators/trigger_dagrun.py b/airflow/operators/trigger_dagrun.py index 9a84bfac97dd1..b0115636c3d2c 100644 --- a/airflow/operators/trigger_dagrun.py +++ b/airflow/operators/trigger_dagrun.py @@ -20,7 +20,7 @@ import datetime import json import time -from typing import TYPE_CHECKING, Sequence, cast +from typing import TYPE_CHECKING, Any, Sequence, cast from sqlalchemy.orm.exc import NoResultFound @@ -211,13 +211,16 @@ def execute(self, context: Context): return @provide_session - def execute_complete(self, context: Context, session: Session, **kwargs): - parsed_execution_date = context["execution_date"] + def execute_complete(self, context: Context, session: Session, event: tuple[str, dict[str, Any]]): + # This execution date is parsed from the return trigger event + provided_execution_date = event[1]["execution_dates"][0] try: dag_run = ( session.query(DagRun) - .filter(DagRun.dag_id == self.trigger_dag_id, DagRun.execution_date == parsed_execution_date) + .filter( + DagRun.dag_id == self.trigger_dag_id, DagRun.execution_date == provided_execution_date + ) .one() ) diff --git a/tests/operators/test_trigger_dagrun.py b/tests/operators/test_trigger_dagrun.py index cb2d75e84c8b7..e0423e2e89711 100644 --- a/tests/operators/test_trigger_dagrun.py +++ b/tests/operators/test_trigger_dagrun.py @@ -411,10 +411,21 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true(self): dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 - task.execute_complete(context={"execution_date": execution_date, "logical_date": execution_date}) + task.execute_complete( + context={}, + event=( + "airflow.triggers.external_task.DagStateTrigger", + { + "dag_id": "down_stream", + "execution_dates": [DEFAULT_DATE], + "poll_interval": 20, + "states": ["success", "failed"], + }, + ), + ) def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure(self): - """Test TriggerDagRunOperator with wait_for_completion.""" + """Test TriggerDagRunOperator wait_for_completion dag run in non defined state.""" execution_date = DEFAULT_DATE task = TriggerDagRunOperator( task_id="test_task", @@ -434,10 +445,22 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure(self): assert len(dagruns) == 1 with pytest.raises(AirflowException): - task.execute_complete(context={"execution_date": execution_date, "logical_date": execution_date}) + task.execute_complete( + context={}, + event=( + "airflow.triggers.external_task.DagStateTrigger", + { + "dag_id": "down_stream", + "execution_dates": [DEFAULT_DATE], + "poll_interval": 20, + "states": ["success", "failed"], + }, + ), + ) + assert "which is not in" in str(exception) def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure_2(self): - """Test TriggerDagRunOperator with wait_for_completion.""" + """Test TriggerDagRunOperator wait_for_completion dag run in failed state.""" execution_date = DEFAULT_DATE task = TriggerDagRunOperator( task_id="test_task", @@ -457,5 +480,18 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure_2(self) dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 - with pytest.raises(AirflowException): - task.execute_complete(context={"execution_date": execution_date, "logical_date": execution_date}) + with pytest.raises(AirflowException) as exception: + task.execute_complete( + context={}, + event=( + "airflow.triggers.external_task.DagStateTrigger", + { + "dag_id": "down_stream", + "execution_dates": [DEFAULT_DATE], + "poll_interval": 20, + "states": ["success", "failed"], + }, + ), + ) + + assert "failed with failed state" in str(exception) From 351ec7932f57d70891c52b96f2f8974c90e06fdd Mon Sep 17 00:00:00 2001 From: Dylan Storey Date: Wed, 5 Apr 2023 09:26:55 -0400 Subject: [PATCH 2/5] pre-commit --- tests/operators/test_trigger_dagrun.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/operators/test_trigger_dagrun.py b/tests/operators/test_trigger_dagrun.py index e0423e2e89711..cd5a1ea57eebb 100644 --- a/tests/operators/test_trigger_dagrun.py +++ b/tests/operators/test_trigger_dagrun.py @@ -444,7 +444,7 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure(self): dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 - with pytest.raises(AirflowException): + with pytest.raises(AirflowException) as exception: task.execute_complete( context={}, event=( From a82b8713720ac3bab043cfe2d48299f3ed17655c Mon Sep 17 00:00:00 2001 From: Dylan Storey Date: Tue, 11 Apr 2023 20:42:28 -0400 Subject: [PATCH 3/5] finally fixing after the github issue last week --- airflow/triggers/external_task.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/triggers/external_task.py b/airflow/triggers/external_task.py index 5ed0d3b5e3f92..61e8f0a1e7270 100644 --- a/airflow/triggers/external_task.py +++ b/airflow/triggers/external_task.py @@ -80,7 +80,7 @@ async def run(self) -> typing.AsyncIterator["TriggerEvent"]: while True: num_tasks = await self.count_tasks() if num_tasks == len(self.execution_dates): - yield TriggerEvent(True) + yield TriggerEvent(self.serialize()) await asyncio.sleep(self.poll_interval) @sync_to_async From e51c93659ca4cd76018d969aa168b6d006b3fe6e Mon Sep 17 00:00:00 2001 From: Dylan Storey Date: Tue, 11 Apr 2023 21:03:47 -0400 Subject: [PATCH 4/5] push fix --- airflow/triggers/external_task.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow/triggers/external_task.py b/airflow/triggers/external_task.py index 61e8f0a1e7270..883753401c58f 100644 --- a/airflow/triggers/external_task.py +++ b/airflow/triggers/external_task.py @@ -80,7 +80,7 @@ async def run(self) -> typing.AsyncIterator["TriggerEvent"]: while True: num_tasks = await self.count_tasks() if num_tasks == len(self.execution_dates): - yield TriggerEvent(self.serialize()) + yield TriggerEvent(True) await asyncio.sleep(self.poll_interval) @sync_to_async @@ -144,7 +144,7 @@ async def run(self) -> typing.AsyncIterator["TriggerEvent"]: while True: num_dags = await self.count_dags() if num_dags == len(self.execution_dates): - yield TriggerEvent(True) + yield TriggerEvent(self.serialize()) await asyncio.sleep(self.poll_interval) @sync_to_async From 1ad3ee3d5cc3fd9e4bcf7cfe19495dd6aba09e32 Mon Sep 17 00:00:00 2001 From: Dylan Storey Date: Wed, 12 Apr 2023 07:19:13 -0400 Subject: [PATCH 5/5] feedback from hussein --- tests/operators/test_trigger_dagrun.py | 56 +++++++++++--------------- 1 file changed, 23 insertions(+), 33 deletions(-) diff --git a/tests/operators/test_trigger_dagrun.py b/tests/operators/test_trigger_dagrun.py index cd5a1ea57eebb..3d24315dbc6ec 100644 --- a/tests/operators/test_trigger_dagrun.py +++ b/tests/operators/test_trigger_dagrun.py @@ -28,6 +28,7 @@ from airflow.models import DAG, DagBag, DagModel, DagRun, Log, TaskInstance from airflow.models.serialized_dag import SerializedDagModel from airflow.operators.trigger_dagrun import TriggerDagRunOperator +from airflow.triggers.external_task import DagStateTrigger from airflow.utils import timezone from airflow.utils.session import create_session from airflow.utils.state import State @@ -410,20 +411,15 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true(self): with create_session() as session: dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 - - task.execute_complete( - context={}, - event=( - "airflow.triggers.external_task.DagStateTrigger", - { - "dag_id": "down_stream", - "execution_dates": [DEFAULT_DATE], - "poll_interval": 20, - "states": ["success", "failed"], - }, - ), + trigger = DagStateTrigger( + dag_id="down_stream", + execution_dates=[DEFAULT_DATE], + poll_interval=20, + states=["success", "failed"], ) + task.execute_complete(context={}, event=trigger.serialize()) + def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure(self): """Test TriggerDagRunOperator wait_for_completion dag run in non defined state.""" execution_date = DEFAULT_DATE @@ -444,18 +440,16 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure(self): dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 + trigger = DagStateTrigger( + dag_id="down_stream", + execution_dates=[DEFAULT_DATE], + poll_interval=20, + states=["success", "failed"], + ) with pytest.raises(AirflowException) as exception: task.execute_complete( context={}, - event=( - "airflow.triggers.external_task.DagStateTrigger", - { - "dag_id": "down_stream", - "execution_dates": [DEFAULT_DATE], - "poll_interval": 20, - "states": ["success", "failed"], - }, - ), + event=trigger.serialize(), ) assert "which is not in" in str(exception) @@ -480,18 +474,14 @@ def test_trigger_dagrun_with_wait_for_completion_true_defer_true_failure_2(self) dagruns = session.query(DagRun).filter(DagRun.dag_id == TRIGGERED_DAG_ID).all() assert len(dagruns) == 1 + trigger = DagStateTrigger( + dag_id="down_stream", + execution_dates=[DEFAULT_DATE], + poll_interval=20, + states=["success", "failed"], + ) + with pytest.raises(AirflowException) as exception: - task.execute_complete( - context={}, - event=( - "airflow.triggers.external_task.DagStateTrigger", - { - "dag_id": "down_stream", - "execution_dates": [DEFAULT_DATE], - "poll_interval": 20, - "states": ["success", "failed"], - }, - ), - ) + task.execute_complete(context={}, event=trigger.serialize()) assert "failed with failed state" in str(exception)