diff --git a/providers/databricks/src/airflow/providers/databricks/sensors/databricks.py b/providers/databricks/src/airflow/providers/databricks/sensors/databricks.py index 9497cf8624d35..729d84af78311 100644 --- a/providers/databricks/src/airflow/providers/databricks/sensors/databricks.py +++ b/providers/databricks/src/airflow/providers/databricks/sensors/databricks.py @@ -70,11 +70,6 @@ def __init__( **kwargs, ): # Handle the scenario where either both statement and statement_id are set/not set - if statement and statement_id: - raise AirflowException("Cannot provide both statement and statement_id.") - if not statement and not statement_id: - raise AirflowException("One of either statement or statement_id must be provided.") - if not warehouse_id: raise AirflowException("warehouse_id must be provided.") @@ -112,6 +107,10 @@ def _get_hook(self, caller: str) -> DatabricksHook: ) def execute(self, context: Context): + if self.statement and self.statement_id: + raise AirflowException("Cannot provide both statement and statement_id.") + if not self.statement and not self.statement_id: + raise AirflowException("One of either statement or statement_id must be provided.") if not self.statement_id: # Otherwise, we'll go ahead and "submit" the statement tags = build_query_tags(context, self.query_tags, self.include_airflow_query_tags) diff --git a/providers/databricks/tests/unit/databricks/sensors/test_databricks.py b/providers/databricks/tests/unit/databricks/sensors/test_databricks.py index fe04781f65d1d..08517a17f884a 100644 --- a/providers/databricks/tests/unit/databricks/sensors/test_databricks.py +++ b/providers/databricks/tests/unit/databricks/sensors/test_databricks.py @@ -68,6 +68,18 @@ def test_init_statement_id(self): assert op.statement_id == STATEMENT_ID assert op.warehouse_id == WAREHOUSE_ID + @pytest.mark.parametrize( + ("kwargs", "match"), + [ + ({"statement": STATEMENT, "statement_id": STATEMENT_ID}, "Cannot provide both"), + ({}, "One of either statement or statement_id"), + ], + ) + def test_statement_combination_validated_at_execute(self, kwargs, match): + op = DatabricksSQLStatementsSensor(task_id=TASK_ID, warehouse_id=WAREHOUSE_ID, **kwargs) + with pytest.raises(AirflowException, match=match): + op.execute(None) + @mock.patch("airflow.providers.databricks.sensors.databricks.DatabricksHook") def test_exec_success(self, db_mock_class): """ diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index b78123db0dc66..59535da72df15 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -29,7 +29,6 @@ providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposCreateOperator providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposDeleteOperator providers/databricks/src/airflow/providers/databricks/operators/databricks_repos.py::DatabricksReposUpdateOperator -providers/databricks/src/airflow/providers/databricks/sensors/databricks.py::DatabricksSQLStatementsSensor providers/docker/src/airflow/providers/docker/operators/docker.py::DockerOperator providers/google/src/airflow/providers/google/cloud/operators/bigquery.py::BigQueryInsertJobOperator providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py::CloudBatchSubmitJobOperator