Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading