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 648b111432d78..e40ca5f4fcfd9 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 @@ -68,11 +68,6 @@ 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) @@ -85,6 +80,10 @@ def execute(self, context: Context) -> None: In deferrable mode, the polling is deferred to the triggerer. Otherwise the sensor waits synchronously. """ + 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)}" + ) if not self.deferrable: super().execute(context=context) else: 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 85c168ab89bde..5d3c93b030548 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,15 +43,19 @@ def test_init(self): assert sensor.target_state == "running" assert sensor.azure_conn_id == CONN_ID - def test_invalid_target_state_rejected_at_poke(self): + @pytest.mark.parametrize("deferrable", [False, True]) + def test_invalid_target_state_rejected_at_execute(self, deferrable): sensor = AzureVirtualMachineStateSensor( task_id="sense_vm", resource_group_name=RESOURCE_GROUP, vm_name=VM_NAME, target_state="invalid_state", + deferrable=deferrable, + silent_fail=True, + timeout=0, ) with pytest.raises(ValueError, match="Invalid target_state"): - sensor.poke(context=None) + sensor.execute(context=None) def test_templated_target_state_constructs(self): sensor = AzureVirtualMachineStateSensor(