diff --git a/providers/dbt/cloud/src/airflow/providers/dbt/cloud/operators/dbt.py b/providers/dbt/cloud/src/airflow/providers/dbt/cloud/operators/dbt.py index 0eb178d4d5ec2..0420ef1208c9e 100644 --- a/providers/dbt/cloud/src/airflow/providers/dbt/cloud/operators/dbt.py +++ b/providers/dbt/cloud/src/airflow/providers/dbt/cloud/operators/dbt.py @@ -289,6 +289,7 @@ def execute(self, context: Context): execution_deadline=execution_deadline, account_id=self.account_id, poll_interval=self.check_interval, + hook_params=self.hook_params, ), method_name="execute_complete", ) diff --git a/providers/dbt/cloud/tests/unit/dbt/cloud/operators/test_dbt.py b/providers/dbt/cloud/tests/unit/dbt/cloud/operators/test_dbt.py index be1ebec812e2b..91c0d3863cf10 100644 --- a/providers/dbt/cloud/tests/unit/dbt/cloud/operators/test_dbt.py +++ b/providers/dbt/cloud/tests/unit/dbt/cloud/operators/test_dbt.py @@ -236,8 +236,44 @@ def test_execute_deferrable_does_not_pass_execution_timeout_to_defer( execution_deadline=ANY, account_id=None, poll_interval=1, + hook_params={}, ) + @patch( + "airflow.providers.dbt.cloud.hooks.dbt.DbtCloudHook.get_job_run_status", + return_value=DbtCloudJobRunStatus.QUEUED.value, + ) + @patch("airflow.providers.dbt.cloud.operators.dbt.DbtCloudRunJobOperator.defer") + @patch("airflow.providers.dbt.cloud.operators.dbt.DbtCloudRunJobTrigger") + @patch("airflow.providers.dbt.cloud.hooks.dbt.DbtCloudHook.get_connection") + @patch( + "airflow.providers.dbt.cloud.hooks.dbt.DbtCloudHook.trigger_job_run", + return_value=mock_response_json(DEFAULT_ACCOUNT_JOB_RUN_RESPONSE), + ) + def test_execute_deferrable_hands_hook_params_to_the_trigger( + self, + mock_trigger_job_run, + mock_dbt_hook, + mock_dbt_trigger, + mock_defer, + mock_job_run_status, + ): + hook_params = {"retry_limit": 3, "retry_delay": 2.0} + dbt_op = DbtCloudRunJobOperator( + dbt_cloud_conn_id=ACCOUNT_ID_CONN, + task_id=TASK_ID, + job_id=JOB_ID, + check_interval=1, + timeout=3, + dag=self.dag, + deferrable=True, + hook_params=hook_params, + ) + + dbt_op.execute(MagicMock()) + + assert mock_dbt_trigger.call_args.kwargs["hook_params"] == hook_params + @pytest.mark.parametrize( "status", (