diff --git a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/sensors/compute.py b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/sensors/compute.py index 1a9c05185dcb6..648b111432d78 100644 --- a/providers/microsoft/azure/src/airflow/providers/microsoft/azure/sensors/compute.py +++ b/providers/microsoft/azure/src/airflow/providers/microsoft/azure/sensors/compute.py @@ -60,10 +60,6 @@ def __init__( deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs, ) -> None: - if target_state not in self.VALID_STATES: - raise ValueError( - f"Invalid target_state: {target_state}. Must be one of {sorted(self.VALID_STATES)}" - ) super().__init__(**kwargs) self.resource_group_name = resource_group_name self.vm_name = vm_name @@ -72,6 +68,11 @@ def __init__( self.deferrable = deferrable def poke(self, context: Context) -> bool: + # target_state is a template field; validate the rendered value here, not in __init__. + if self.target_state not in self.VALID_STATES: + raise ValueError( + f"Invalid target_state: {self.target_state}. Must be one of {sorted(self.VALID_STATES)}" + ) hook = AzureComputeHook(azure_conn_id=self.azure_conn_id) current_state = hook.get_power_state(self.resource_group_name, self.vm_name) self.log.info("VM %s power state: %s", self.vm_name, current_state) diff --git a/providers/microsoft/azure/tests/unit/microsoft/azure/sensors/test_compute.py b/providers/microsoft/azure/tests/unit/microsoft/azure/sensors/test_compute.py index 47e4da37a8d46..85c168ab89bde 100644 --- a/providers/microsoft/azure/tests/unit/microsoft/azure/sensors/test_compute.py +++ b/providers/microsoft/azure/tests/unit/microsoft/azure/sensors/test_compute.py @@ -43,14 +43,24 @@ def test_init(self): assert sensor.target_state == "running" assert sensor.azure_conn_id == CONN_ID - def test_init_invalid_target_state(self): + def test_invalid_target_state_rejected_at_poke(self): + sensor = AzureVirtualMachineStateSensor( + task_id="sense_vm", + resource_group_name=RESOURCE_GROUP, + vm_name=VM_NAME, + target_state="invalid_state", + ) with pytest.raises(ValueError, match="Invalid target_state"): - AzureVirtualMachineStateSensor( - task_id="sense_vm", - resource_group_name=RESOURCE_GROUP, - vm_name=VM_NAME, - target_state="invalid_state", - ) + sensor.poke(context=None) + + def test_templated_target_state_constructs(self): + sensor = AzureVirtualMachineStateSensor( + task_id="sense_vm", + resource_group_name=RESOURCE_GROUP, + vm_name=VM_NAME, + target_state="{{ params.state }}", + ) + assert sensor.target_state == "{{ params.state }}" def test_template_fields(self): sensor = AzureVirtualMachineStateSensor( diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 01fd5bc56dbb7..af9658b934c99 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -73,7 +73,6 @@ providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_bigquery.py providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_gcs.py::GCSToGCSOperator providers/google/src/airflow/providers/google/cloud/transfers/gcs_to_local.py::GCSToLocalFilesystemOperator providers/google/src/airflow/providers/google/marketing_platform/operators/campaign_manager.py::GoogleCampaignManagerDeleteReportOperator -providers/microsoft/azure/src/airflow/providers/microsoft/azure/sensors/compute.py::AzureVirtualMachineStateSensor providers/microsoft/azure/src/airflow/providers/microsoft/azure/transfers/gcs_to_wasb.py::GCSToAzureBlobStorageOperator providers/microsoft/azure/src/airflow/providers/microsoft/azure/transfers/oracle_to_azure_data_lake.py::OracleToAzureDataLakeOperator providers/microsoft/psrp/src/airflow/providers/microsoft/psrp/operators/psrp.py::PsrpOperator