diff --git a/providers/google/src/airflow/providers/google/cloud/sensors/bigquery.py b/providers/google/src/airflow/providers/google/cloud/sensors/bigquery.py index 29f77ffd27f11..65531b857bb2f 100644 --- a/providers/google/src/airflow/providers/google/cloud/sensors/bigquery.py +++ b/providers/google/src/airflow/providers/google/cloud/sensors/bigquery.py @@ -129,6 +129,7 @@ def execute(self, context: Context) -> None: project_id=self.project_id, poll_interval=self.poke_interval, gcp_conn_id=self.gcp_conn_id, + impersonation_chain=self.impersonation_chain, hook_params={ "impersonation_chain": self.impersonation_chain, }, @@ -301,6 +302,7 @@ def execute(self, context: Context) -> None: partition_id=self.partition_id, poll_interval=self.poke_interval, gcp_conn_id=self.gcp_conn_id, + impersonation_chain=self.impersonation_chain, hook_params={ "impersonation_chain": self.impersonation_chain, }, diff --git a/providers/google/tests/unit/google/cloud/sensors/test_bigquery.py b/providers/google/tests/unit/google/cloud/sensors/test_bigquery.py index 1881c360e5edc..eb953d6af03fc 100644 --- a/providers/google/tests/unit/google/cloud/sensors/test_bigquery.py +++ b/providers/google/tests/unit/google/cloud/sensors/test_bigquery.py @@ -103,6 +103,24 @@ def test_execute_deferred(self, mock_hook): "Trigger is not a BigQueryTableExistenceTrigger" ) + @mock.patch("airflow.providers.google.cloud.sensors.bigquery.BigQueryHook") + def test_deferred_trigger_receives_impersonation_chain(self, mock_hook): + task = BigQueryTableExistenceSensor( + task_id="check_table_exists", + project_id=TEST_PROJECT_ID, + dataset_id=TEST_DATASET_ID, + table_id=TEST_TABLE_ID, + gcp_conn_id=TEST_GCP_CONN_ID, + impersonation_chain=TEST_IMPERSONATION_CHAIN, + deferrable=True, + ) + mock_hook.return_value.table_exists.return_value = False + + with pytest.raises(TaskDeferred) as exc: + task.execute(mock.MagicMock()) + + assert exc.value.trigger.impersonation_chain == TEST_IMPERSONATION_CHAIN + def test_execute_deferred_failure(self): """Tests that an expected exception is raised in case of error event""" task = BigQueryTableExistenceSensor( @@ -206,6 +224,25 @@ def test_execute_with_deferrable_mode(self, mock_hook): "Trigger is not a BigQueryTablePartitionExistenceTrigger" ) + @mock.patch("airflow.providers.google.cloud.sensors.bigquery.BigQueryHook") + def test_deferred_trigger_receives_impersonation_chain(self, mock_hook): + task = BigQueryTablePartitionExistenceSensor( + task_id="test_task_id", + project_id=TEST_PROJECT_ID, + dataset_id=TEST_DATASET_ID, + table_id=TEST_TABLE_ID, + partition_id=TEST_PARTITION_ID, + gcp_conn_id=TEST_GCP_CONN_ID, + impersonation_chain=TEST_IMPERSONATION_CHAIN, + deferrable=True, + ) + mock_hook.return_value.table_partition_exists.return_value = False + + with pytest.raises(TaskDeferred) as exc: + task.execute(context={}) + + assert exc.value.trigger.impersonation_chain == TEST_IMPERSONATION_CHAIN + def test_execute_with_deferrable_mode_execute_failure(self): """Tests that an AirflowException is raised in case of error event""" task = BigQueryTablePartitionExistenceSensor(