Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions airflow/sensors/external_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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` "
Expand Down Expand Up @@ -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)

Expand Down
72 changes: 64 additions & 8 deletions tests/sensors/test_external_task_sensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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"
Expand Down Expand Up @@ -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):
Expand Down