diff --git a/airflow/sensors/external_task.py b/airflow/sensors/external_task.py index e1eec6d192684..8c4b71398df8a 100644 --- a/airflow/sensors/external_task.py +++ b/airflow/sensors/external_task.py @@ -162,8 +162,7 @@ def __init__( f"when `external_task_id` or `external_task_ids` or `external_task_group_id` " f"is not `None`: {State.task_states}" ) - if external_task_ids and len(external_task_ids) > len(set(external_task_ids)): - raise ValueError("Duplicate task_ids passed in external_task_ids parameter") + elif not total_states <= set(State.dag_states): raise ValueError( f"Valid values for `allowed_states` and `failed_states` " @@ -196,6 +195,9 @@ def _get_dttm_filter(self, context): @provide_session def poke(self, context, session=None): + if self.external_task_ids and len(self.external_task_ids) > len(set(self.external_task_ids)): + raise ValueError("Duplicate task_ids passed in external_task_ids parameter") + dttm_filter = self._get_dttm_filter(context) serialized_dttm_filter = ",".join(dt.isoformat() for dt in dttm_filter) diff --git a/tests/sensors/test_external_task_sensor.py b/tests/sensors/test_external_task_sensor.py index b3a22d77f2a11..04e1b4c9431d0 100644 --- a/tests/sensors/test_external_task_sensor.py +++ b/tests/sensors/test_external_task_sensor.py @@ -32,8 +32,10 @@ from airflow.models import DagBag, DagRun, TaskInstance from airflow.models.dag import DAG from airflow.models.serialized_dag import SerializedDagModel +from airflow.models.xcom_arg import XComArg from airflow.operators.bash import BashOperator from airflow.operators.empty import EmptyOperator +from airflow.operators.python import PythonOperator from airflow.sensors.external_task import ExternalTaskMarker, ExternalTaskSensor, ExternalTaskSensorLink from airflow.sensors.time_sensor import TimeSensor from airflow.serialization.serialized_objects import SerializedBaseOperator @@ -44,6 +46,7 @@ from airflow.utils.types import DagRunType from tests.models import TEST_DAGS_FOLDER from tests.test_utils.db import clear_db_runs +from tests.test_utils.mock_operators import MockOperator DEFAULT_DATE = datetime(2015, 1, 1) TEST_DAG_ID = "unit_test_dag" @@ -576,17 +579,70 @@ def test_external_task_sensor_error_task_id_and_task_ids(self): dag=self.dag, ) + def test_external_task_sensor_with_xcom_arg_does_not_fail_on_init(self): + self.add_time_sensor() + op1 = MockOperator(task_id="op1", dag=self.dag) + op2 = ExternalTaskSensor( + task_id="test_external_task_sensor_with_xcom_arg_does_not_fail_on_init", + external_dag_id=TEST_DAG_ID, + external_task_ids=XComArg(op1), + allowed_states=["success"], + dag=self.dag, + ) + assert isinstance(op2.external_task_ids, XComArg) + def test_catch_duplicate_task_ids(self): self.add_time_sensor() - # Test By passing same task_id multiple times + op1 = ExternalTaskSensor( + task_id="test_external_task_duplicate_task_ids", + external_dag_id=TEST_DAG_ID, + external_task_ids=[TEST_TASK_ID, TEST_TASK_ID], + allowed_states=["success"], + dag=self.dag, + ) with pytest.raises(ValueError): - ExternalTaskSensor( - task_id="test_external_task_duplicate_task_ids", - external_dag_id=TEST_DAG_ID, - external_task_ids=[TEST_TASK_ID, TEST_TASK_ID], - allowed_states=["success"], - dag=self.dag, - ) + op1.run(start_date=DEFAULT_DATE, end_date=DEFAULT_DATE) + + def test_catch_duplicate_task_ids_with_xcom_arg(self): + self.add_time_sensor() + op1 = PythonOperator( + python_callable=lambda: ["dupe_value", "dupe_value"], + task_id="op1", + do_xcom_push=True, + dag=self.dag, + ) + + op2 = ExternalTaskSensor( + task_id="test_external_task_duplicate_task_ids_with_xcom_arg", + external_dag_id=TEST_DAG_ID, + external_task_ids=XComArg(op1), + allowed_states=["success"], + dag=self.dag, + ) + with pytest.raises(ValueError): + op1.run(start_date=DEFAULT_DATE, end_date=DEFAULT_DATE) + op2.run(start_date=DEFAULT_DATE, end_date=DEFAULT_DATE) + + def test_catch_duplicate_task_ids_with_multiple_xcom_args(self): + self.add_time_sensor() + + op1 = PythonOperator( + python_callable=lambda: "value", + task_id="op1", + do_xcom_push=True, + dag=self.dag, + ) + + op2 = ExternalTaskSensor( + task_id="test_external_task_duplicate_task_ids_with_xcom_arg", + external_dag_id=TEST_DAG_ID, + external_task_ids=[XComArg(op1), XComArg(op1)], + allowed_states=["success"], + dag=self.dag, + ) + with pytest.raises(ValueError): + op1.run(start_date=DEFAULT_DATE, end_date=DEFAULT_DATE) + op2.run(start_date=DEFAULT_DATE, end_date=DEFAULT_DATE) def test_catch_invalid_allowed_states(self): with pytest.raises(ValueError):