diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py b/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py index e926dbdfe3b28..a666710849198 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py @@ -194,8 +194,6 @@ def __init__( **kwargs, ): super().__init__(aws_conn_id=aws_conn_id, **kwargs) - if cdc_start_time and cdc_start_position: - raise ValueError("Only one of cdc_start_time or cdc_start_position can be provided.") self.replication_task_arn = replication_task_arn self.table_mappings = table_mappings self.migration_type = migration_type @@ -216,6 +214,9 @@ def _wait_for_modification_completion(self) -> None: ) def execute(self, context: Context) -> dict: + if self.cdc_start_time and self.cdc_start_position: + raise ValueError("Only one of cdc_start_time or cdc_start_position can be provided.") + tasks = self.hook.find_replication_tasks_by_arn( replication_task_arn=self.replication_task_arn, without_settings=True ) @@ -799,10 +800,10 @@ def __init__( self.waiter_max_attempts = waiter_max_attempts self.wait_for_completion = wait_for_completion + def execute(self, context: Context): if self.cdc_start_time and self.cdc_start_pos: raise AirflowException("Only one of cdc_start_time or cdc_start_pos should be provided.") - def execute(self, context: Context): result = self.hook.describe_replications( filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}] ) diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py b/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py index b33270360fcbc..e70149a8c18af 100644 --- a/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py +++ b/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py @@ -207,15 +207,16 @@ def _running_task(self): def _modifying_task(self): return [{"ReplicationTaskArn": self.TASK_ARN, "Status": "modifying"}] - def test_init_raises_if_both_cdc_start_params_provided(self): + def test_execute_raises_if_both_cdc_start_params_provided(self): + op = DmsModifyTaskOperator( + task_id="modify_task", + replication_task_arn=self.TASK_ARN, + cdc_start_time=datetime(2024, 1, 1), + cdc_start_position="mysql-bin.000001:4", + ) with pytest.raises(ValueError, match="Only one of"): - DmsModifyTaskOperator( - task_id="modify_task", - replication_task_arn=self.TASK_ARN, - cdc_start_time=datetime(2024, 1, 1), - cdc_start_position="mysql-bin.000001:4", - ) + op.execute(None) @pytest.mark.parametrize("status", ["stopped", "ready", "failed"]) @mock.patch.object(DmsHook, "find_replication_tasks_by_arn") @@ -1225,29 +1226,18 @@ def mock_replication_response(self, status: str): } } - def test_arg_validation(self): - with pytest.raises(AirflowException): - DmsStartReplicationOperator( - task_id="start_replication", - replication_config_arn="XXXXXXXXXXXXXXX", - replication_start_type="cdc", - cdc_start_pos=1, - cdc_start_time="2024-01-01 00:00:00", - ) - DmsStartReplicationOperator( + def test_execute_raises_if_both_cdc_start_params_provided(self): + op = DmsStartReplicationOperator( task_id="start_replication", replication_config_arn="XXXXXXXXXXXXXXX", replication_start_type="cdc", cdc_start_pos=1, - ) - - DmsStartReplicationOperator( - task_id="start_replication", - replication_config_arn="XXXXXXXXXXXXXXX", - replication_start_type="cdc", cdc_start_time="2024-01-01 00:00:00", ) + with pytest.raises(AirflowException, match="Only one of"): + op.execute({}) + @mock.patch.object(DmsHook, "describe_replications") @mock.patch.object(DmsHook, "start_replication") def test_already_running(self, mock_replication, mock_describe): diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 29ee3dd263922..ffc69836d916f 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -7,8 +7,6 @@ # execute()) MUST remove its entry in the same PR — the hook fails on stale entries. # Burn-down tracked at https://github.com/apache/airflow/issues/70296 providers/amazon/src/airflow/providers/amazon/aws/operators/appflow.py::AppflowBaseOperator -providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py::DmsModifyTaskOperator -providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py::DmsStartReplicationOperator providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py::EcsRunTaskOperator providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py::EmrAddStepsOperator providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStartDbClusterOperator