From 3002452a3afaed19e8a8e62edd7c1e0da486eb23 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 13:29:33 +0100 Subject: [PATCH 1/6] adding template_fields tests in operators --- .../amazon/aws/operators/test_athena.py | 9 ++ .../amazon/aws/operators/test_bedrock.py | 58 ++++++++++++ .../aws/operators/test_cloud_formation.py | 39 ++++++++ .../amazon/aws/operators/test_comprehend.py | 20 ++++ .../amazon/aws/operators/test_datasync.py | 10 ++ .../amazon/aws/operators/test_dms.py | 87 ++++++++++++++++++ .../amazon/aws/operators/test_ec2.py | 92 +++++++++++++++++++ .../amazon/aws/operators/test_ecs.py | 56 +++++++++++ .../amazon/aws/operators/test_eks.py | 90 ++++++++++++++++++ 9 files changed, 461 insertions(+) diff --git a/tests/providers/amazon/aws/operators/test_athena.py b/tests/providers/amazon/aws/operators/test_athena.py index c132e6456f1d8..90713284bc96c 100644 --- a/tests/providers/amazon/aws/operators/test_athena.py +++ b/tests/providers/amazon/aws/operators/test_athena.py @@ -397,3 +397,12 @@ def mock_get_table_metadata(CatalogName, DatabaseName, TableName): run_facets={"externalQuery": ExternalQueryRunFacet(externalQueryId="12345", source="awsathena")}, ) assert op.get_openlineage_facets_on_complete(None) == expected_lineage + + def test_template_fields(self): + template_fields = list(self.athena.template_fields) + list(self.athena.template_fields_renderers.keys()) + + class_fields = self.athena.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_bedrock.py b/tests/providers/amazon/aws/operators/test_bedrock.py index b49d09b52a5c0..d9e303f3ca6e0 100644 --- a/tests/providers/amazon/aws/operators/test_bedrock.py +++ b/tests/providers/amazon/aws/operators/test_bedrock.py @@ -176,6 +176,15 @@ def test_ensure_unique_job_name(self, _, side_effect, ensure_unique_name, mock_c bedrock_hook.get_waiter.assert_not_called() self.operator.defer.assert_not_called() + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestBedrockCreateProvisionedModelThroughputOperator: MODEL_ARN = "testProvisionedModelArn" @@ -222,7 +231,14 @@ def test_provisioned_model_wait_combinations( assert bedrock_hook.get_waiter.call_count == wait_for_completion assert self.operator.defer.call_count == deferrable + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" class TestBedrockCreateKnowledgeBaseOperator: KNOWLEDGE_BASE_ID = "knowledge_base_id" @@ -288,7 +304,14 @@ def test_returns_id(self, mock_conn): assert result == self.KNOWLEDGE_BASE_ID + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" class TestBedrockCreateDataSourceOperator: DATA_SOURCE_ID = "data_source_id" @@ -317,6 +340,15 @@ def test_id_returned(self, mock_conn): assert result == self.DATA_SOURCE_ID + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestBedrockIngestDataOperator: INGESTION_JOB_ID = "ingestion_job_id" @@ -348,6 +380,15 @@ def test_id_returned(self, mock_conn): assert result == self.INGESTION_JOB_ID + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestBedrockRaGOperator: VECTOR_SEARCH_CONFIG = {"filter": {"equals": {"key": "some key", "value": "some value"}}} @@ -520,3 +561,20 @@ def test_external_sources_build_rag_config(self, prompt_template): **expected_config_without_template, **expected_config_template, } + + def test_template_fields(self): + op = BedrockRaGOperator( + task_id="test_rag", + input="some text prompt", + source_type="EXTERNAL_SOURCES", + model_arn=self.MODEL_ARN, + knowledge_base_id=self.KNOWLEDGE_BASE_ID, + vector_search_config=self.VECTOR_SEARCH_CONFIG, + ) + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_cloud_formation.py b/tests/providers/amazon/aws/operators/test_cloud_formation.py index 5de02c3622cfb..b7eb06f520827 100644 --- a/tests/providers/amazon/aws/operators/test_cloud_formation.py +++ b/tests/providers/amazon/aws/operators/test_cloud_formation.py @@ -87,6 +87,26 @@ def test_create_stack(self, mocked_hook_client): StackName=stack_name, TemplateBody=template_body, TimeoutInMinutes=timeout ) + def test_template_fields(self): + op = CloudFormationCreateStackOperator( + task_id="cf_create_stack_init", + stack_name="fake-stack", + cloudformation_parameters={}, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="eu-west-1", + verify=True, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestCloudFormationDeleteStackOperator: def test_init(self): @@ -125,3 +145,22 @@ def test_delete_stack(self, mocked_hook_client): operator.execute(MagicMock()) mocked_hook_client.delete_stack.assert_any_call(StackName=stack_name) + + def test_template_fields(self): + op = CloudFormationDeleteStackOperator( + task_id="cf_delete_stack_init", + stack_name="fake-stack", + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="us-east-1", + verify=False, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_comprehend.py b/tests/providers/amazon/aws/operators/test_comprehend.py index 60f0fca219111..e2d674fb421cf 100644 --- a/tests/providers/amazon/aws/operators/test_comprehend.py +++ b/tests/providers/amazon/aws/operators/test_comprehend.py @@ -163,6 +163,16 @@ def test_start_pii_entities_detection_job_wait_combinations( assert comprehend_hook.get_waiter.call_count == wait_for_completion assert self.operator.defer.call_count == deferrable + def test_template_fields(self): + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestComprehendCreateDocumentClassifierOperator: CLASSIFIER_ARN = ( @@ -259,3 +269,13 @@ def test_create_document_classifier_wait_combinations( assert response == self.CLASSIFIER_ARN assert comprehend_hook.get_waiter.call_count == wait_for_completion assert self.operator.defer.call_count == deferrable + + def test_template_fields(self): + + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_datasync.py b/tests/providers/amazon/aws/operators/test_datasync.py index e1a44ce99e28c..39dd279af8a93 100644 --- a/tests/providers/amazon/aws/operators/test_datasync.py +++ b/tests/providers/amazon/aws/operators/test_datasync.py @@ -363,6 +363,16 @@ def test_return_value(self, mock_get_conn, session, clean_dags_and_dagruns): # ### Check mocks: mock_get_conn.assert_called() + def test_template_fields(self, mock_get_conn): + self.set_up_operator() + template_fields = list(self.datasync.template_fields) + list(self.datasync.template_fields_renderers.keys()) + + class_fields = self.datasync.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + @mock_aws @mock.patch.object(DataSyncHook, "get_conn") diff --git a/tests/providers/amazon/aws/operators/test_dms.py b/tests/providers/amazon/aws/operators/test_dms.py index fba14a6370dd7..7ff6e1843fd10 100644 --- a/tests/providers/amazon/aws/operators/test_dms.py +++ b/tests/providers/amazon/aws/operators/test_dms.py @@ -121,6 +121,25 @@ def test_create_task_with_migration_type( assert dms_hook.get_task_status(TASK_ARN) == "ready" + def test_template_fields(self): + op = DmsCreateTaskOperator( + task_id="create_task", + **self.TASK_DATA, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="ca-west-1", + verify=True, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestDmsDeleteTaskOperator: TASK_DATA = { @@ -174,6 +193,25 @@ def test_delete_task( assert dms_hook.get_task_status(TASK_ARN) == "deleting" + def test_template_fields(self): + op = DmsDeleteTaskOperator( + task_id="delete_task", + replication_task_arn=TASK_ARN, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="us-east-1", + verify=False, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestDmsDescribeTasksOperator: FILTER = {"Name": "replication-task-arn", "Values": [TASK_ARN]} @@ -267,6 +305,17 @@ def test_describe_tasks_return_value(self, mock_conn, mock_describe_replication_ assert marker is None assert response == self.MOCK_RESPONSE + def test_template_fields(self): + op = DmsDescribeTasksOperator( + task_id="describe_tasks", + describe_tasks_kwargs={"Filters": [self.FILTER]}, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="eu-west-2", + verify="/foo/bar/spam.egg", + botocore_config={"read_timeout": 42}, + ) + class TestDmsStartTaskOperator: TASK_DATA = { @@ -324,6 +373,25 @@ def test_start_task( assert dms_hook.get_task_status(TASK_ARN) == "starting" + def test_template_fields(self): + op = DmsStartTaskOperator( + task_id="start_task", + replication_task_arn=TASK_ARN, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="us-west-1", + verify=False, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestDmsStopTaskOperator: TASK_DATA = { @@ -376,3 +444,22 @@ def test_stop_task( mock_stop_replication_task.assert_called_once_with(replication_task_arn=TASK_ARN) assert dms_hook.get_task_status(TASK_ARN) == "stopping" + + def test_template_fields(self): + op = DmsStopTaskOperator( + task_id="stop_task", + replication_task_arn=TASK_ARN, + # Generic hooks parameters + aws_conn_id="fake-conn-id", + region_name="eu-west-1", + verify=True, + botocore_config={"read_timeout": 42}, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_ec2.py b/tests/providers/amazon/aws/operators/test_ec2.py index 8f8a755a84357..d0dc0e299aff7 100644 --- a/tests/providers/amazon/aws/operators/test_ec2.py +++ b/tests/providers/amazon/aws/operators/test_ec2.py @@ -87,6 +87,20 @@ def test_create_multiple_instances(self): for id in instance_ids: assert ec2_hook.get_instance_state(instance_id=id) == "running" + def test_template_fields(self): + ec2_operator = EC2CreateInstanceOperator( + task_id="test_create_instance", + image_id="test_image_id", + ) + + template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEC2TerminateInstanceOperator(BaseEc2TestClass): def test_init(self): @@ -140,6 +154,20 @@ def test_terminate_multiple_instances(self): for id in instance_ids: assert ec2_hook.get_instance_state(instance_id=id) == "terminated" + def test_template_fields(self): + ec2_operator = EC2TerminateInstanceOperator( + task_id="test_terminate_instance", + instance_ids="test_image_id", + ) + + template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEC2StartInstanceOperator(BaseEc2TestClass): def test_init(self): @@ -156,6 +184,7 @@ def test_init(self): assert ec2_operator.region_name == "region-test" assert ec2_operator.check_interval == 3 + @mock_aws def test_start_instance(self): # create instance @@ -175,6 +204,24 @@ def test_start_instance(self): # assert instance state is running assert ec2_hook.get_instance_state(instance_id=instance_id[0]) == "running" + def test_template_fields(self): + ec2_operator = EC2StartInstanceOperator( + task_id="task_test", + instance_id="i-123abc", + aws_conn_id="aws_conn_test", + region_name="region-test", + check_interval=3, + ) + + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEC2StopInstanceOperator(BaseEc2TestClass): def test_init(self): @@ -210,6 +257,23 @@ def test_stop_instance(self): # assert instance state is running assert ec2_hook.get_instance_state(instance_id=instance_id[0]) == "stopped" + def test_template_fields(self): + ec2_operator = EC2StopInstanceOperator( + task_id="task_test", + instance_id="i-123abc", + aws_conn_id="aws_conn_test", + region_name="region-test", + check_interval=3, + ) + + template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEC2HibernateInstanceOperator(BaseEc2TestClass): def test_init(self): @@ -322,6 +386,20 @@ def test_cannot_hibernate_some_instances(self): for id in instance_ids: assert ec2_hook.get_instance_state(instance_id=id) == "running" + def test_template_fields(self): + ec2_operator = EC2HibernateInstanceOperator( + task_id="task_test", + instance_ids="i-123abc", + ) + + template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEC2RebootInstanceOperator(BaseEc2TestClass): def test_init(self): @@ -372,3 +450,17 @@ def test_reboot_multiple_instances(self): terminate_instance.execute(None) for id in instance_ids: assert ec2_hook.get_instance_state(instance_id=id) == "running" + + def test_template_fields(self): + ec2_operator = EC2RebootInstanceOperator( + task_id="task_test", + instance_ids="i-123abc", + ) + + template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + + class_fields = ec2_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_ecs.py b/tests/providers/amazon/aws/operators/test_ecs.py index a6915214a0764..ba388515dd605 100644 --- a/tests/providers/amazon/aws/operators/test_ecs.py +++ b/tests/providers/amazon/aws/operators/test_ecs.py @@ -793,6 +793,23 @@ def test_execute_without_waiter(self, patch_hook_waiters): patch_hook_waiters.assert_not_called() assert result is not None + def test_template_fields(self): + op = EcsCreateClusterOperator( + task_id="task", + cluster_name=CLUSTER_NAME, + deferrable=True, + waiter_delay=12, + waiter_max_attempts=34, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEcsDeleteClusterOperator(EcsBaseTestCase): @pytest.mark.parametrize("waiter_delay, waiter_max_attempts", WAITERS_TEST_CASES) @@ -857,6 +874,23 @@ def test_execute_without_waiter(self, patch_hook_waiters): patch_hook_waiters.assert_not_called() assert result is not None + def test_template_fields(self): + op = EcsDeleteClusterOperator( + task_id="task", + cluster_name=CLUSTER_NAME, + deferrable=True, + waiter_delay=12, + waiter_max_attempts=34, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEcsDeregisterTaskDefinitionOperator(EcsBaseTestCase): warn_message = "'wait_for_completion' and waiter related params have no effect" @@ -913,6 +947,17 @@ def test_partial_deprecation_waiters_params( assert not hasattr(ti.task, "waiter_delay") assert not hasattr(ti.task, "waiter_max_attempts") + def test_template_fields(self): + op = EcsDeregisterTaskDefinitionOperator(task_id="task", task_definition=TASK_DEFINITION_NAME) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEcsRegisterTaskDefinitionOperator(EcsBaseTestCase): warn_message = "'wait_for_completion' and waiter related params have no effect" @@ -990,3 +1035,14 @@ def test_partial_deprecation_waiters_params( assert not hasattr(ti.task, "wait_for_completion") assert not hasattr(ti.task, "waiter_delay") assert not hasattr(ti.task, "waiter_max_attempts") + + def test_template_fields(self): + op = EcsRegisterTaskDefinitionOperator(task_id="task", **TASK_DEFINITION_CONFIG) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_eks.py b/tests/providers/amazon/aws/operators/test_eks.py index 9571ca0962005..20889257f2f76 100644 --- a/tests/providers/amazon/aws/operators/test_eks.py +++ b/tests/providers/amazon/aws/operators/test_eks.py @@ -365,6 +365,21 @@ def test_eks_create_cluster_with_deferrable(self, mock_create_cluster, caplog): eks_create_cluster_operator.execute({}) assert "Waiting for EKS Cluster to provision. This will take some time." in caplog.messages + def test_template_fields(self): + op = EksCreateClusterOperator( + task_id=TASK_ID, + **self.create_cluster_params, + compute="fargate", + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksCreateFargateProfileOperator: def setup_method(self) -> None: @@ -445,6 +460,18 @@ def test_create_fargate_profile_deferrable(self, _): exc.value.trigger, EksCreateFargateProfileTrigger ), "Trigger is not a EksCreateFargateProfileTrigger" + def test_template_fields(self): + + op = EksCreateFargateProfileOperator(task_id=TASK_ID, **self.create_fargate_profile_params) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksCreateNodegroupOperator: def setup_method(self) -> None: @@ -536,6 +563,18 @@ def test_create_nodegroup_deferrable_versus_wait_for_completion(self): ) assert operator.wait_for_completion is True + def test_template_fields(self): + op_kwargs = {**self.create_nodegroup_params} + op = EksCreateNodegroupOperator(task_id=TASK_ID, **op_kwargs) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksDeleteClusterOperator: def setup_method(self) -> None: @@ -575,6 +614,16 @@ def test_eks_delete_cluster_operator_with_deferrable(self): with pytest.raises(TaskDeferred): self.delete_cluster_operator.execute({}) + def test_template_fields(self): + template_fields = list(self.delete_cluster_operator.template_fields) + list( + self.delete_cluster_operator.template_fields_renderers.keys()) + + class_fields = self.delete_cluster_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksDeleteNodegroupOperator: def setup_method(self) -> None: @@ -608,6 +657,16 @@ def test_existing_nodegroup_with_wait(self, mock_delete_nodegroup, mock_waiter): mock_waiter.assert_called_with(mock.ANY, clusterName=CLUSTER_NAME, nodegroupName=NODEGROUP_NAME) assert_expected_waiter_type(mock_waiter, "NodegroupDeleted") + def test_template_fields(self): + template_fields = list(self.delete_nodegroup_operator.template_fields) + list( + self.delete_nodegroup_operator.template_fields_renderers.keys()) + + class_fields = self.delete_nodegroup_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksDeleteFargateProfileOperator: def setup_method(self) -> None: @@ -656,6 +715,16 @@ def test_delete_fargate_profile_deferrable(self, _): exc.value.trigger, EksDeleteFargateProfileTrigger ), "Trigger is not a EksDeleteFargateProfileTrigger" + def test_template_fields(self): + template_fields = list(self.delete_fargate_profile_operator.template_fields) + list( + self.delete_fargate_profile_operator.template_fields_renderers.keys()) + + class_fields = self.delete_fargate_profile_operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEksPodOperator: @mock.patch("airflow.providers.cncf.kubernetes.operators.pod.KubernetesPodOperator.execute") @@ -767,3 +836,24 @@ def test_on_finish_action_handler( ) for expected_attr in expected_attributes: assert op.__getattribute__(expected_attr) == expected_attributes[expected_attr] + + def test_template_fields(self): + op = EksPodOperator( + task_id="run_pod", + pod_name="run_pod", + cluster_name=CLUSTER_NAME, + image="amazon/aws-cli:latest", + cmds=["sh", "-c", "ls"], + labels={"demo": "hello_world"}, + get_logs=True, + on_finish_action="delete_pod", + ) + + template_fields = list(op.template_fields) + list( + op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" From 02d006024697e162939476e6c6d7156c13c5b674 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 13:34:14 +0100 Subject: [PATCH 2/6] fix static checks --- .../amazon/aws/operators/test_athena.py | 4 +++- .../amazon/aws/operators/test_bedrock.py | 24 +++++++++++++++---- .../amazon/aws/operators/test_comprehend.py | 8 ++++--- .../amazon/aws/operators/test_datasync.py | 4 +++- .../amazon/aws/operators/test_ec2.py | 24 +++++++++++++------ .../amazon/aws/operators/test_eks.py | 13 +++++----- 6 files changed, 54 insertions(+), 23 deletions(-) diff --git a/tests/providers/amazon/aws/operators/test_athena.py b/tests/providers/amazon/aws/operators/test_athena.py index 90713284bc96c..99cff9a573efc 100644 --- a/tests/providers/amazon/aws/operators/test_athena.py +++ b/tests/providers/amazon/aws/operators/test_athena.py @@ -399,7 +399,9 @@ def mock_get_table_metadata(CatalogName, DatabaseName, TableName): assert op.get_openlineage_facets_on_complete(None) == expected_lineage def test_template_fields(self): - template_fields = list(self.athena.template_fields) + list(self.athena.template_fields_renderers.keys()) + template_fields = list(self.athena.template_fields) + list( + self.athena.template_fields_renderers.keys() + ) class_fields = self.athena.__dict__ diff --git a/tests/providers/amazon/aws/operators/test_bedrock.py b/tests/providers/amazon/aws/operators/test_bedrock.py index d9e303f3ca6e0..01a42f4443ec1 100644 --- a/tests/providers/amazon/aws/operators/test_bedrock.py +++ b/tests/providers/amazon/aws/operators/test_bedrock.py @@ -177,7 +177,9 @@ def test_ensure_unique_job_name(self, _, side_effect, ensure_unique_name, mock_c self.operator.defer.assert_not_called() def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ @@ -232,13 +234,17 @@ def test_provisioned_model_wait_combinations( assert self.operator.defer.call_count == deferrable def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ missing_fields = [field for field in template_fields if field not in class_fields] assert not missing_fields, f"Templated fields are not available {missing_fields}" + + class TestBedrockCreateKnowledgeBaseOperator: KNOWLEDGE_BASE_ID = "knowledge_base_id" @@ -305,13 +311,17 @@ def test_returns_id(self, mock_conn): assert result == self.KNOWLEDGE_BASE_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ missing_fields = [field for field in template_fields if field not in class_fields] assert not missing_fields, f"Templated fields are not available {missing_fields}" + + class TestBedrockCreateDataSourceOperator: DATA_SOURCE_ID = "data_source_id" @@ -341,7 +351,9 @@ def test_id_returned(self, mock_conn): assert result == self.DATA_SOURCE_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ @@ -381,7 +393,9 @@ def test_id_returned(self, mock_conn): assert result == self.INGESTION_JOB_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ diff --git a/tests/providers/amazon/aws/operators/test_comprehend.py b/tests/providers/amazon/aws/operators/test_comprehend.py index e2d674fb421cf..4e72010ff2ea0 100644 --- a/tests/providers/amazon/aws/operators/test_comprehend.py +++ b/tests/providers/amazon/aws/operators/test_comprehend.py @@ -165,7 +165,8 @@ def test_start_pii_entities_detection_job_wait_combinations( def test_template_fields(self): template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys()) + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ @@ -271,8 +272,9 @@ def test_create_document_classifier_wait_combinations( assert self.operator.defer.call_count == deferrable def test_template_fields(self): - - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + template_fields = list(self.operator.template_fields) + list( + self.operator.template_fields_renderers.keys() + ) class_fields = self.operator.__dict__ diff --git a/tests/providers/amazon/aws/operators/test_datasync.py b/tests/providers/amazon/aws/operators/test_datasync.py index 39dd279af8a93..ee670e181c439 100644 --- a/tests/providers/amazon/aws/operators/test_datasync.py +++ b/tests/providers/amazon/aws/operators/test_datasync.py @@ -365,7 +365,9 @@ def test_return_value(self, mock_get_conn, session, clean_dags_and_dagruns): def test_template_fields(self, mock_get_conn): self.set_up_operator() - template_fields = list(self.datasync.template_fields) + list(self.datasync.template_fields_renderers.keys()) + template_fields = list(self.datasync.template_fields) + list( + self.datasync.template_fields_renderers.keys() + ) class_fields = self.datasync.__dict__ diff --git a/tests/providers/amazon/aws/operators/test_ec2.py b/tests/providers/amazon/aws/operators/test_ec2.py index d0dc0e299aff7..397397daad423 100644 --- a/tests/providers/amazon/aws/operators/test_ec2.py +++ b/tests/providers/amazon/aws/operators/test_ec2.py @@ -93,7 +93,9 @@ def test_template_fields(self): image_id="test_image_id", ) - template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ @@ -160,7 +162,9 @@ def test_template_fields(self): instance_ids="test_image_id", ) - template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ @@ -184,7 +188,6 @@ def test_init(self): assert ec2_operator.region_name == "region-test" assert ec2_operator.check_interval == 3 - @mock_aws def test_start_instance(self): # create instance @@ -214,7 +217,8 @@ def test_template_fields(self): ) template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys()) + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ @@ -266,7 +270,9 @@ def test_template_fields(self): check_interval=3, ) - template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ @@ -392,7 +398,9 @@ def test_template_fields(self): instance_ids="i-123abc", ) - template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ @@ -457,7 +465,9 @@ def test_template_fields(self): instance_ids="i-123abc", ) - template_fields = list(ec2_operator.template_fields) + list(ec2_operator.template_fields_renderers.keys()) + template_fields = list(ec2_operator.template_fields) + list( + ec2_operator.template_fields_renderers.keys() + ) class_fields = ec2_operator.__dict__ diff --git a/tests/providers/amazon/aws/operators/test_eks.py b/tests/providers/amazon/aws/operators/test_eks.py index 20889257f2f76..dd75d92a053c6 100644 --- a/tests/providers/amazon/aws/operators/test_eks.py +++ b/tests/providers/amazon/aws/operators/test_eks.py @@ -461,7 +461,6 @@ def test_create_fargate_profile_deferrable(self, _): ), "Trigger is not a EksCreateFargateProfileTrigger" def test_template_fields(self): - op = EksCreateFargateProfileOperator(task_id=TASK_ID, **self.create_fargate_profile_params) template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) @@ -616,7 +615,8 @@ def test_eks_delete_cluster_operator_with_deferrable(self): def test_template_fields(self): template_fields = list(self.delete_cluster_operator.template_fields) + list( - self.delete_cluster_operator.template_fields_renderers.keys()) + self.delete_cluster_operator.template_fields_renderers.keys() + ) class_fields = self.delete_cluster_operator.__dict__ @@ -659,7 +659,8 @@ def test_existing_nodegroup_with_wait(self, mock_delete_nodegroup, mock_waiter): def test_template_fields(self): template_fields = list(self.delete_nodegroup_operator.template_fields) + list( - self.delete_nodegroup_operator.template_fields_renderers.keys()) + self.delete_nodegroup_operator.template_fields_renderers.keys() + ) class_fields = self.delete_nodegroup_operator.__dict__ @@ -717,7 +718,8 @@ def test_delete_fargate_profile_deferrable(self, _): def test_template_fields(self): template_fields = list(self.delete_fargate_profile_operator.template_fields) + list( - self.delete_fargate_profile_operator.template_fields_renderers.keys()) + self.delete_fargate_profile_operator.template_fields_renderers.keys() + ) class_fields = self.delete_fargate_profile_operator.__dict__ @@ -849,8 +851,7 @@ def test_template_fields(self): on_finish_action="delete_pod", ) - template_fields = list(op.template_fields) + list( - op.template_fields_renderers.keys()) + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) class_fields = op.__dict__ From be82a3d72f23241c6bb68ad54d6b38c51f13394f Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 15:03:24 +0100 Subject: [PATCH 3/6] adding template_fields test to emr operator --- .../aws/operators/test_emr_add_steps.py | 16 +++++ .../aws/operators/test_emr_containers.py | 10 +++ .../aws/operators/test_emr_create_job_flow.py | 10 +++ .../aws/operators/test_emr_modify_cluster.py | 10 +++ .../operators/test_emr_notebook_execution.py | 24 +++++++ .../aws/operators/test_emr_serverless.py | 67 +++++++++++++++++++ .../operators/test_emr_terminate_job_flow.py | 17 +++++ 7 files changed, 154 insertions(+) diff --git a/tests/providers/amazon/aws/operators/test_emr_add_steps.py b/tests/providers/amazon/aws/operators/test_emr_add_steps.py index 9ee99864e00e3..218fec3e3b861 100644 --- a/tests/providers/amazon/aws/operators/test_emr_add_steps.py +++ b/tests/providers/amazon/aws/operators/test_emr_add_steps.py @@ -274,3 +274,19 @@ def test_emr_add_steps_deferrable(self, mock_add_job_flow_steps, mock_get_log_ur operator.execute(MagicMock()) assert isinstance(exc.value.trigger, EmrAddStepsTrigger), "Trigger is not a EmrAddStepsTrigger" + + def test_template_fields(self): + op = EmrAddStepsOperator( + task_id="test_task", + job_flow_id="j-8989898989", + aws_conn_id="aws_default", + steps=self._config, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_containers.py b/tests/providers/amazon/aws/operators/test_emr_containers.py index feeec1278e155..4b0142e5ddc94 100644 --- a/tests/providers/amazon/aws/operators/test_emr_containers.py +++ b/tests/providers/amazon/aws/operators/test_emr_containers.py @@ -194,3 +194,13 @@ def test_emr_on_eks_execute_with_failure(self, mock_create_emr_on_eks_cluster): with pytest.raises(AirflowException) as ctx: self.emr_container.execute(None) assert expected_exception_msg in str(ctx.value) + + def test_template_fields(self): + + template_fields = list(self.emr_container.template_fields) + list(self.emr_container.template_fields_renderers.keys()) + + class_fields = self.emr_container.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py b/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py index 204d292c67b46..633f93fb9d9ff 100644 --- a/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py +++ b/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py @@ -203,3 +203,13 @@ def test_create_job_flow_deferrable(self, mocked_hook_client): assert isinstance( exc.value.trigger, EmrCreateJobFlowTrigger ), "Trigger is not a EmrCreateJobFlowTrigger" + + def test_template_fields(self): + + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py b/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py index 6dada442ff79f..af902da9e9c19 100644 --- a/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py +++ b/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py @@ -65,3 +65,13 @@ def test_execute_returns_error(self, mocked_hook_client): with pytest.raises(AirflowException, match="Modify cluster failed"): self.operator.execute(self.mock_context) + + def test_template_fields(self): + + template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) + + class_fields = self.operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py b/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py index ef6cb7ebc70ec..6aa52df6e5c4a 100644 --- a/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py +++ b/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py @@ -303,3 +303,27 @@ def test_stop_notebook_execution_waiter_config(self, mock_conn, mock_waiter, _): WaiterConfig={"Delay": delay, "MaxAttempts": waiter_max_attempts}, ) assert_expected_waiter_type(mock_waiter, "notebook_stopped") + + def test_template_fields(self): + + op = EmrStartNotebookExecutionOperator( + task_id="test-id", + editor_id=PARAMS["EditorId"], + relative_path=PARAMS["RelativePath"], + cluster_id=PARAMS["ExecutionEngine"]["Id"], + service_role=PARAMS["ServiceRole"], + notebook_execution_name=PARAMS["NotebookExecutionName"], + notebook_params=PARAMS["NotebookParams"], + notebook_instance_security_group_id=PARAMS["NotebookInstanceSecurityGroupId"], + master_instance_security_group_id=PARAMS["ExecutionEngine"]["MasterInstanceSecurityGroupId"], + tags=PARAMS["Tags"], + wait_for_completion=True, + ) + + template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) + + class_fields = op.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_serverless.py b/tests/providers/amazon/aws/operators/test_emr_serverless.py index 12c5cc938018e..a05b8deebc462 100644 --- a/tests/providers/amazon/aws/operators/test_emr_serverless.py +++ b/tests/providers/amazon/aws/operators/test_emr_serverless.py @@ -393,6 +393,26 @@ def test_create_application_deferrable(self, mock_conn): with pytest.raises(TaskDeferred): operator.execute(None) + def test_template_fields(self): + + operator = EmrServerlessCreateApplicationOperator( + task_id=task_id, + release_label=release_label, + job_type=job_type, + client_request_token=client_request_token, + config=config, + waiter_max_attempts=3, + waiter_delay=0, + ) + + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEmrServerlessStartJobOperator: def setup_method(self): @@ -1163,6 +1183,25 @@ def test_links_spark_without_applicationui_enabled( job_run_id=job_run_id, ) + def test_template_fields(self): + + operator = EmrServerlessStartJobOperator( + task_id=task_id, + client_request_token=client_request_token, + application_id=application_id, + execution_role_arn=execution_role_arn, + job_driver=job_driver, + configuration_overrides=configuration_overrides, + ) + + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEmrServerlessDeleteOperator: @mock.patch.object(EmrServerlessHook, "get_waiter") @@ -1277,6 +1316,20 @@ def test_delete_application_deferrable(self, mock_conn): with pytest.raises(TaskDeferred): operator.execute(None) + def test_template_fields(self): + + operator = EmrServerlessDeleteApplicationOperator( + task_id=task_id, application_id=application_id_delete_operator + ) + + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + class TestEmrServerlessStopOperator: @mock.patch.object(EmrServerlessHook, "get_waiter") @@ -1344,3 +1397,17 @@ def test_stop_application_deferrable_without_force_stop( operator.execute({}) assert "no running jobs found with application ID test" in caplog.messages + + def test_template_fields(self): + + operator = EmrServerlessStopApplicationOperator( + task_id=task_id, application_id="test", deferrable=True, force_stop=True + ) + + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py b/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py index 2c27c146d2d34..8d7d85e5af914 100644 --- a/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py +++ b/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py @@ -57,3 +57,20 @@ def test_create_job_flow_deferrable(self, mocked_hook_client): assert isinstance( exc.value.trigger, EmrTerminateJobFlowTrigger ), "Trigger is not a EmrTerminateJobFlowTrigger" + + def test_template_fields(self): + + operator = EmrTerminateJobFlowOperator( + task_id="test_task", + job_flow_id="j-8989898989", + aws_conn_id="aws_default", + deferrable=True, + ) + + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" From edabacb05b0026ba29c8521bac7c75d37a33d57c Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 22:26:03 +0100 Subject: [PATCH 4/6] refactor asserts for template fields --- .pre-commit-config.yaml | 2 +- scripts/cov/core_coverage.py | 2 +- .../amazon/aws/operators/test_athena.py | 11 +--- .../amazon/aws/operators/test_bedrock.py | 59 ++--------------- .../aws/operators/test_cloud_formation.py | 17 +---- .../amazon/aws/operators/test_comprehend.py | 21 +----- .../amazon/aws/operators/test_datasync.py | 11 +--- .../amazon/aws/operators/test_dms.py | 35 ++-------- .../amazon/aws/operators/test_ec2.py | 65 ++----------------- .../amazon/aws/operators/test_ecs.py | 33 ++-------- .../amazon/aws/operators/test_eks.py | 62 +++--------------- .../aws/operators/test_emr_add_steps.py | 10 +-- .../aws/operators/test_emr_containers.py | 9 +-- .../aws/operators/test_emr_create_job_flow.py | 10 +-- .../aws/operators/test_emr_modify_cluster.py | 10 +-- .../operators/test_emr_notebook_execution.py | 10 +-- .../aws/operators/test_emr_serverless.py | 13 +--- .../operators/test_emr_terminate_job_flow.py | 10 +-- .../amazon/aws/utils/test_template_fields.py | 30 +++++++++ 19 files changed, 89 insertions(+), 331 deletions(-) create mode 100644 tests/providers/amazon/aws/utils/test_template_fields.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 9e91b09613926..7d2875c1f2551 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1171,7 +1171,7 @@ repos: entry: "^\\s*from re\\s|^\\s*import re\\s" pass_filenames: true files: \.py$ - exclude: ^airflow/providers|^dev/.*\.py$|^scripts/.*\.py$|^tests/|^\w+_tests/|^docs/.*\.py$|^airflow/utils/helpers.py$|^hatch_build.py$ + exclude: ^airflow/providers|^dev/.*\.py$|^scripts/.*\.py$|^tests/|^\w+_tests/|^docs/.*\.py$|^airflow/utils/test_template_fields.py$|^hatch_build.py$ - id: check-provider-docs-valid name: Validate provider doc files entry: ./scripts/ci/pre_commit/check_provider_docs.py diff --git a/scripts/cov/core_coverage.py b/scripts/cov/core_coverage.py index 0facd4bb1c5d7..edc1fd00294dd 100644 --- a/scripts/cov/core_coverage.py +++ b/scripts/cov/core_coverage.py @@ -118,7 +118,7 @@ "airflow/utils/entry_points.py", "airflow/utils/file.py", "airflow/utils/hashlib_wrapper.py", - "airflow/utils/helpers.py", + "airflow/utils/test_template_fields.py", "airflow/utils/json.py", "airflow/utils/log/action_logger.py", "airflow/utils/log/colored_log.py", diff --git a/tests/providers/amazon/aws/operators/test_athena.py b/tests/providers/amazon/aws/operators/test_athena.py index 99cff9a573efc..102d1fe31e5c1 100644 --- a/tests/providers/amazon/aws/operators/test_athena.py +++ b/tests/providers/amazon/aws/operators/test_athena.py @@ -39,6 +39,7 @@ from airflow.utils import timezone from airflow.utils.timezone import datetime from airflow.utils.types import DagRunType +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields TEST_DAG_ID = "unit_tests" DEFAULT_DATE = datetime(2018, 1, 1) @@ -399,12 +400,4 @@ def mock_get_table_metadata(CatalogName, DatabaseName, TableName): assert op.get_openlineage_facets_on_complete(None) == expected_lineage def test_template_fields(self): - template_fields = list(self.athena.template_fields) + list( - self.athena.template_fields_renderers.keys() - ) - - class_fields = self.athena.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.athena) diff --git a/tests/providers/amazon/aws/operators/test_bedrock.py b/tests/providers/amazon/aws/operators/test_bedrock.py index 01a42f4443ec1..8cbb67d6f50df 100644 --- a/tests/providers/amazon/aws/operators/test_bedrock.py +++ b/tests/providers/amazon/aws/operators/test_bedrock.py @@ -35,6 +35,7 @@ BedrockInvokeModelOperator, BedrockRaGOperator, ) +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields if TYPE_CHECKING: from airflow.providers.amazon.aws.hooks.base_aws import BaseAwsConnection @@ -177,15 +178,7 @@ def test_ensure_unique_job_name(self, _, side_effect, ensure_unique_name, mock_c self.operator.defer.assert_not_called() def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestBedrockCreateProvisionedModelThroughputOperator: @@ -234,15 +227,7 @@ def test_provisioned_model_wait_combinations( assert self.operator.defer.call_count == deferrable def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestBedrockCreateKnowledgeBaseOperator: @@ -311,15 +296,7 @@ def test_returns_id(self, mock_conn): assert result == self.KNOWLEDGE_BASE_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestBedrockCreateDataSourceOperator: @@ -351,15 +328,7 @@ def test_id_returned(self, mock_conn): assert result == self.DATA_SOURCE_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestBedrockIngestDataOperator: @@ -393,15 +362,7 @@ def test_id_returned(self, mock_conn): assert result == self.INGESTION_JOB_ID def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestBedrockRaGOperator: @@ -585,10 +546,4 @@ def test_template_fields(self): knowledge_base_id=self.KNOWLEDGE_BASE_ID, vector_search_config=self.VECTOR_SEARCH_CONFIG, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_cloud_formation.py b/tests/providers/amazon/aws/operators/test_cloud_formation.py index b7eb06f520827..4d8fb4d12bd3c 100644 --- a/tests/providers/amazon/aws/operators/test_cloud_formation.py +++ b/tests/providers/amazon/aws/operators/test_cloud_formation.py @@ -28,6 +28,7 @@ CloudFormationDeleteStackOperator, ) from airflow.utils import timezone +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields DEFAULT_DATE = timezone.datetime(2019, 1, 1) DEFAULT_ARGS = {"owner": "airflow", "start_date": DEFAULT_DATE} @@ -99,13 +100,7 @@ def test_template_fields(self): botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestCloudFormationDeleteStackOperator: @@ -157,10 +152,4 @@ def test_template_fields(self): botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_comprehend.py b/tests/providers/amazon/aws/operators/test_comprehend.py index 4e72010ff2ea0..a86b779b1d502 100644 --- a/tests/providers/amazon/aws/operators/test_comprehend.py +++ b/tests/providers/amazon/aws/operators/test_comprehend.py @@ -29,6 +29,7 @@ ComprehendStartPiiEntitiesDetectionJobOperator, ) from airflow.utils.types import NOTSET +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields if TYPE_CHECKING: from airflow.providers.amazon.aws.hooks.base_aws import BaseAwsConnection @@ -164,15 +165,7 @@ def test_start_pii_entities_detection_job_wait_combinations( assert self.operator.defer.call_count == deferrable def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) class TestComprehendCreateDocumentClassifierOperator: @@ -272,12 +265,4 @@ def test_create_document_classifier_wait_combinations( assert self.operator.defer.call_count == deferrable def test_template_fields(self): - template_fields = list(self.operator.template_fields) + list( - self.operator.template_fields_renderers.keys() - ) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) diff --git a/tests/providers/amazon/aws/operators/test_datasync.py b/tests/providers/amazon/aws/operators/test_datasync.py index ee670e181c439..18b0e86103c0b 100644 --- a/tests/providers/amazon/aws/operators/test_datasync.py +++ b/tests/providers/amazon/aws/operators/test_datasync.py @@ -29,6 +29,7 @@ from airflow.utils import timezone from airflow.utils.timezone import datetime from airflow.utils.types import DagRunType +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields TEST_DAG_ID = "unit_tests" DEFAULT_DATE = datetime(2018, 1, 1) @@ -365,15 +366,7 @@ def test_return_value(self, mock_get_conn, session, clean_dags_and_dagruns): def test_template_fields(self, mock_get_conn): self.set_up_operator() - template_fields = list(self.datasync.template_fields) + list( - self.datasync.template_fields_renderers.keys() - ) - - class_fields = self.datasync.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.datasync) @mock_aws diff --git a/tests/providers/amazon/aws/operators/test_dms.py b/tests/providers/amazon/aws/operators/test_dms.py index 7ff6e1843fd10..2528edaef9e0a 100644 --- a/tests/providers/amazon/aws/operators/test_dms.py +++ b/tests/providers/amazon/aws/operators/test_dms.py @@ -34,6 +34,7 @@ ) from airflow.utils import timezone from airflow.utils.types import DagRunType +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields TASK_ARN = "test_arn" @@ -125,20 +126,13 @@ def test_template_fields(self): op = DmsCreateTaskOperator( task_id="create_task", **self.TASK_DATA, - # Generic hooks parameters aws_conn_id="fake-conn-id", region_name="ca-west-1", verify=True, botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestDmsDeleteTaskOperator: @@ -204,13 +198,7 @@ def test_template_fields(self): botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestDmsDescribeTasksOperator: @@ -315,6 +303,7 @@ def test_template_fields(self): verify="/foo/bar/spam.egg", botocore_config={"read_timeout": 42}, ) + validate_template_fields(op) class TestDmsStartTaskOperator: @@ -384,13 +373,7 @@ def test_template_fields(self): botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestDmsStopTaskOperator: @@ -456,10 +439,4 @@ def test_template_fields(self): botocore_config={"read_timeout": 42}, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_ec2.py b/tests/providers/amazon/aws/operators/test_ec2.py index 397397daad423..a5ea81ff6ae87 100644 --- a/tests/providers/amazon/aws/operators/test_ec2.py +++ b/tests/providers/amazon/aws/operators/test_ec2.py @@ -30,6 +30,7 @@ EC2StopInstanceOperator, EC2TerminateInstanceOperator, ) +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields class BaseEc2TestClass: @@ -92,16 +93,7 @@ def test_template_fields(self): task_id="test_create_instance", image_id="test_image_id", ) - - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) class TestEC2TerminateInstanceOperator(BaseEc2TestClass): @@ -161,16 +153,7 @@ def test_template_fields(self): task_id="test_terminate_instance", instance_ids="test_image_id", ) - - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) class TestEC2StartInstanceOperator(BaseEc2TestClass): @@ -216,15 +199,7 @@ def test_template_fields(self): check_interval=3, ) - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) class TestEC2StopInstanceOperator(BaseEc2TestClass): @@ -270,15 +245,7 @@ def test_template_fields(self): check_interval=3, ) - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) class TestEC2HibernateInstanceOperator(BaseEc2TestClass): @@ -397,16 +364,7 @@ def test_template_fields(self): task_id="task_test", instance_ids="i-123abc", ) - - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) class TestEC2RebootInstanceOperator(BaseEc2TestClass): @@ -464,13 +422,4 @@ def test_template_fields(self): task_id="task_test", instance_ids="i-123abc", ) - - template_fields = list(ec2_operator.template_fields) + list( - ec2_operator.template_fields_renderers.keys() - ) - - class_fields = ec2_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(ec2_operator) diff --git a/tests/providers/amazon/aws/operators/test_ecs.py b/tests/providers/amazon/aws/operators/test_ecs.py index ba388515dd605..be06a8802e449 100644 --- a/tests/providers/amazon/aws/operators/test_ecs.py +++ b/tests/providers/amazon/aws/operators/test_ecs.py @@ -39,6 +39,7 @@ from airflow.providers.amazon.aws.utils.task_log_fetcher import AwsTaskLogFetcher from airflow.utils.task_instance_session import set_current_task_instance_session from airflow.utils.types import NOTSET +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields CLUSTER_NAME = "test_cluster" CONTAINER_NAME = "e1ed7aac-d9b2-4315-8726-d2432bf11868" @@ -802,13 +803,7 @@ def test_template_fields(self): waiter_max_attempts=34, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEcsDeleteClusterOperator(EcsBaseTestCase): @@ -883,13 +878,7 @@ def test_template_fields(self): waiter_max_attempts=34, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEcsDeregisterTaskDefinitionOperator(EcsBaseTestCase): @@ -950,13 +939,7 @@ def test_partial_deprecation_waiters_params( def test_template_fields(self): op = EcsDeregisterTaskDefinitionOperator(task_id="task", task_definition=TASK_DEFINITION_NAME) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEcsRegisterTaskDefinitionOperator(EcsBaseTestCase): @@ -1039,10 +1022,4 @@ def test_partial_deprecation_waiters_params( def test_template_fields(self): op = EcsRegisterTaskDefinitionOperator(task_id="task", **TASK_DEFINITION_CONFIG) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_eks.py b/tests/providers/amazon/aws/operators/test_eks.py index dd75d92a053c6..fe8601d1a6ea0 100644 --- a/tests/providers/amazon/aws/operators/test_eks.py +++ b/tests/providers/amazon/aws/operators/test_eks.py @@ -51,6 +51,7 @@ TASK_ID, ) from tests.providers.amazon.aws.utils.eks_test_utils import convert_keys +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields from tests.providers.amazon.aws.utils.test_waiter import assert_expected_waiter_type CLUSTER_NAME = "cluster1" @@ -372,13 +373,7 @@ def test_template_fields(self): compute="fargate", ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEksCreateFargateProfileOperator: @@ -463,13 +458,7 @@ def test_create_fargate_profile_deferrable(self, _): def test_template_fields(self): op = EksCreateFargateProfileOperator(task_id=TASK_ID, **self.create_fargate_profile_params) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEksCreateNodegroupOperator: @@ -566,13 +555,7 @@ def test_template_fields(self): op_kwargs = {**self.create_nodegroup_params} op = EksCreateNodegroupOperator(task_id=TASK_ID, **op_kwargs) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) class TestEksDeleteClusterOperator: @@ -614,15 +597,8 @@ def test_eks_delete_cluster_operator_with_deferrable(self): self.delete_cluster_operator.execute({}) def test_template_fields(self): - template_fields = list(self.delete_cluster_operator.template_fields) + list( - self.delete_cluster_operator.template_fields_renderers.keys() - ) - - class_fields = self.delete_cluster_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.delete_cluster_operator) class TestEksDeleteNodegroupOperator: @@ -658,15 +634,7 @@ def test_existing_nodegroup_with_wait(self, mock_delete_nodegroup, mock_waiter): assert_expected_waiter_type(mock_waiter, "NodegroupDeleted") def test_template_fields(self): - template_fields = list(self.delete_nodegroup_operator.template_fields) + list( - self.delete_nodegroup_operator.template_fields_renderers.keys() - ) - - class_fields = self.delete_nodegroup_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.delete_nodegroup_operator) class TestEksDeleteFargateProfileOperator: @@ -717,15 +685,7 @@ def test_delete_fargate_profile_deferrable(self, _): ), "Trigger is not a EksDeleteFargateProfileTrigger" def test_template_fields(self): - template_fields = list(self.delete_fargate_profile_operator.template_fields) + list( - self.delete_fargate_profile_operator.template_fields_renderers.keys() - ) - - class_fields = self.delete_fargate_profile_operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.delete_fargate_profile_operator) class TestEksPodOperator: @@ -851,10 +811,4 @@ def test_template_fields(self): on_finish_action="delete_pod", ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_emr_add_steps.py b/tests/providers/amazon/aws/operators/test_emr_add_steps.py index 218fec3e3b861..d5a999349aa53 100644 --- a/tests/providers/amazon/aws/operators/test_emr_add_steps.py +++ b/tests/providers/amazon/aws/operators/test_emr_add_steps.py @@ -31,6 +31,7 @@ from airflow.providers.amazon.aws.triggers.emr import EmrAddStepsTrigger from airflow.utils import timezone from airflow.utils.types import DagRunType +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields from tests.test_utils import AIRFLOW_MAIN_FOLDER DEFAULT_DATE = timezone.datetime(2017, 1, 1) @@ -282,11 +283,4 @@ def test_template_fields(self): aws_conn_id="aws_default", steps=self._config, ) - - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_emr_containers.py b/tests/providers/amazon/aws/operators/test_emr_containers.py index 4b0142e5ddc94..4f6c7d8fa7156 100644 --- a/tests/providers/amazon/aws/operators/test_emr_containers.py +++ b/tests/providers/amazon/aws/operators/test_emr_containers.py @@ -25,6 +25,7 @@ from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook from airflow.providers.amazon.aws.operators.emr import EmrContainerOperator, EmrEksCreateClusterOperator from airflow.providers.amazon.aws.triggers.emr import EmrContainerTrigger +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields SUBMIT_JOB_SUCCESS_RETURN = { "ResponseMetadata": {"HTTPStatusCode": 200}, @@ -196,11 +197,5 @@ def test_emr_on_eks_execute_with_failure(self, mock_create_emr_on_eks_cluster): assert expected_exception_msg in str(ctx.value) def test_template_fields(self): + validate_template_fields(self.emr_container) - template_fields = list(self.emr_container.template_fields) + list(self.emr_container.template_fields_renderers.keys()) - - class_fields = self.emr_container.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" diff --git a/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py b/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py index 633f93fb9d9ff..860df8c7219ac 100644 --- a/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py +++ b/tests/providers/amazon/aws/operators/test_emr_create_job_flow.py @@ -32,6 +32,7 @@ from airflow.providers.amazon.aws.triggers.emr import EmrCreateJobFlowTrigger from airflow.utils import timezone from airflow.utils.types import DagRunType +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields from tests.providers.amazon.aws.utils.test_waiter import assert_expected_waiter_type from tests.test_utils import AIRFLOW_MAIN_FOLDER @@ -205,11 +206,4 @@ def test_create_job_flow_deferrable(self, mocked_hook_client): ), "Trigger is not a EmrCreateJobFlowTrigger" def test_template_fields(self): - - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) diff --git a/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py b/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py index af902da9e9c19..6f257288760c3 100644 --- a/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py +++ b/tests/providers/amazon/aws/operators/test_emr_modify_cluster.py @@ -25,6 +25,7 @@ from airflow.models.dag import DAG from airflow.providers.amazon.aws.operators.emr import EmrModifyClusterOperator from airflow.utils import timezone +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields DEFAULT_DATE = timezone.datetime(2017, 1, 1) MODIFY_CLUSTER_SUCCESS_RETURN = {"ResponseMetadata": {"HTTPStatusCode": 200}, "StepConcurrencyLevel": 1} @@ -67,11 +68,4 @@ def test_execute_returns_error(self, mocked_hook_client): self.operator.execute(self.mock_context) def test_template_fields(self): - - template_fields = list(self.operator.template_fields) + list(self.operator.template_fields_renderers.keys()) - - class_fields = self.operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(self.operator) diff --git a/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py b/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py index 6aa52df6e5c4a..6fcd4eeb74629 100644 --- a/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py +++ b/tests/providers/amazon/aws/operators/test_emr_notebook_execution.py @@ -28,6 +28,7 @@ EmrStartNotebookExecutionOperator, EmrStopNotebookExecutionOperator, ) +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields from tests.providers.amazon.aws.utils.test_waiter import assert_expected_waiter_type PARAMS = { @@ -305,7 +306,6 @@ def test_stop_notebook_execution_waiter_config(self, mock_conn, mock_waiter, _): assert_expected_waiter_type(mock_waiter, "notebook_stopped") def test_template_fields(self): - op = EmrStartNotebookExecutionOperator( task_id="test-id", editor_id=PARAMS["EditorId"], @@ -320,10 +320,4 @@ def test_template_fields(self): wait_for_completion=True, ) - template_fields = list(op.template_fields) + list(op.template_fields_renderers.keys()) - - class_fields = op.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(op) diff --git a/tests/providers/amazon/aws/operators/test_emr_serverless.py b/tests/providers/amazon/aws/operators/test_emr_serverless.py index a05b8deebc462..e7a43cf079f0b 100644 --- a/tests/providers/amazon/aws/operators/test_emr_serverless.py +++ b/tests/providers/amazon/aws/operators/test_emr_serverless.py @@ -32,6 +32,7 @@ EmrServerlessStopApplicationOperator, ) from airflow.utils.types import NOTSET +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields if TYPE_CHECKING: from unittest.mock import MagicMock @@ -394,7 +395,6 @@ def test_create_application_deferrable(self, mock_conn): operator.execute(None) def test_template_fields(self): - operator = EmrServerlessCreateApplicationOperator( task_id=task_id, release_label=release_label, @@ -1184,7 +1184,6 @@ def test_links_spark_without_applicationui_enabled( ) def test_template_fields(self): - operator = EmrServerlessStartJobOperator( task_id=task_id, client_request_token=client_request_token, @@ -1317,7 +1316,6 @@ def test_delete_application_deferrable(self, mock_conn): operator.execute(None) def test_template_fields(self): - operator = EmrServerlessDeleteApplicationOperator( task_id=task_id, application_id=application_id_delete_operator ) @@ -1399,15 +1397,8 @@ def test_stop_application_deferrable_without_force_stop( assert "no running jobs found with application ID test" in caplog.messages def test_template_fields(self): - operator = EmrServerlessStopApplicationOperator( task_id=task_id, application_id="test", deferrable=True, force_stop=True ) - template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) - - class_fields = operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(operator) diff --git a/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py b/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py index 8d7d85e5af914..06ab35e4510ba 100644 --- a/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py +++ b/tests/providers/amazon/aws/operators/test_emr_terminate_job_flow.py @@ -24,6 +24,7 @@ from airflow.exceptions import TaskDeferred from airflow.providers.amazon.aws.operators.emr import EmrTerminateJobFlowOperator from airflow.providers.amazon.aws.triggers.emr import EmrTerminateJobFlowTrigger +from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields TERMINATE_SUCCESS_RETURN = {"ResponseMetadata": {"HTTPStatusCode": 200}} @@ -59,7 +60,6 @@ def test_create_job_flow_deferrable(self, mocked_hook_client): ), "Trigger is not a EmrTerminateJobFlowTrigger" def test_template_fields(self): - operator = EmrTerminateJobFlowOperator( task_id="test_task", job_flow_id="j-8989898989", @@ -67,10 +67,4 @@ def test_template_fields(self): deferrable=True, ) - template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) - - class_fields = operator.__dict__ - - missing_fields = [field for field in template_fields if field not in class_fields] - - assert not missing_fields, f"Templated fields are not available {missing_fields}" + validate_template_fields(operator) diff --git a/tests/providers/amazon/aws/utils/test_template_fields.py b/tests/providers/amazon/aws/utils/test_template_fields.py new file mode 100644 index 0000000000000..8ec0d87a59413 --- /dev/null +++ b/tests/providers/amazon/aws/utils/test_template_fields.py @@ -0,0 +1,30 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +def validate_template_fields(operator): + template_fields = list(operator.template_fields) + list( + operator.template_fields_renderers.keys() + ) + + class_fields = operator.__dict__ + + missing_fields = [field for field in template_fields if field not in class_fields] + + assert not missing_fields, f"Templated fields are not available {missing_fields}" + + From dc471a35fc2342285d342f21c894ecad65ddf921 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 22:27:38 +0100 Subject: [PATCH 5/6] fix static checks --- tests/providers/amazon/aws/operators/test_eks.py | 1 - .../providers/amazon/aws/operators/test_emr_containers.py | 1 - tests/providers/amazon/aws/utils/test_template_fields.py | 8 +++----- 3 files changed, 3 insertions(+), 7 deletions(-) diff --git a/tests/providers/amazon/aws/operators/test_eks.py b/tests/providers/amazon/aws/operators/test_eks.py index fe8601d1a6ea0..399c8e40823ae 100644 --- a/tests/providers/amazon/aws/operators/test_eks.py +++ b/tests/providers/amazon/aws/operators/test_eks.py @@ -597,7 +597,6 @@ def test_eks_delete_cluster_operator_with_deferrable(self): self.delete_cluster_operator.execute({}) def test_template_fields(self): - validate_template_fields(self.delete_cluster_operator) diff --git a/tests/providers/amazon/aws/operators/test_emr_containers.py b/tests/providers/amazon/aws/operators/test_emr_containers.py index 4f6c7d8fa7156..52306864f3597 100644 --- a/tests/providers/amazon/aws/operators/test_emr_containers.py +++ b/tests/providers/amazon/aws/operators/test_emr_containers.py @@ -198,4 +198,3 @@ def test_emr_on_eks_execute_with_failure(self, mock_create_emr_on_eks_cluster): def test_template_fields(self): validate_template_fields(self.emr_container) - diff --git a/tests/providers/amazon/aws/utils/test_template_fields.py b/tests/providers/amazon/aws/utils/test_template_fields.py index 8ec0d87a59413..689977de9bcc5 100644 --- a/tests/providers/amazon/aws/utils/test_template_fields.py +++ b/tests/providers/amazon/aws/utils/test_template_fields.py @@ -15,16 +15,14 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +from __future__ import annotations + def validate_template_fields(operator): - template_fields = list(operator.template_fields) + list( - operator.template_fields_renderers.keys() - ) + template_fields = list(operator.template_fields) + list(operator.template_fields_renderers.keys()) class_fields = operator.__dict__ missing_fields = [field for field in template_fields if field not in class_fields] assert not missing_fields, f"Templated fields are not available {missing_fields}" - - From fc5091f05703eef48952838fb199ea30063ee6de Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Wed, 11 Sep 2024 22:39:34 +0100 Subject: [PATCH 6/6] revert changes in pre-commit and coverage --- .pre-commit-config.yaml | 2 +- scripts/cov/core_coverage.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 7d2875c1f2551..9e91b09613926 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1171,7 +1171,7 @@ repos: entry: "^\\s*from re\\s|^\\s*import re\\s" pass_filenames: true files: \.py$ - exclude: ^airflow/providers|^dev/.*\.py$|^scripts/.*\.py$|^tests/|^\w+_tests/|^docs/.*\.py$|^airflow/utils/test_template_fields.py$|^hatch_build.py$ + exclude: ^airflow/providers|^dev/.*\.py$|^scripts/.*\.py$|^tests/|^\w+_tests/|^docs/.*\.py$|^airflow/utils/helpers.py$|^hatch_build.py$ - id: check-provider-docs-valid name: Validate provider doc files entry: ./scripts/ci/pre_commit/check_provider_docs.py diff --git a/scripts/cov/core_coverage.py b/scripts/cov/core_coverage.py index edc1fd00294dd..0facd4bb1c5d7 100644 --- a/scripts/cov/core_coverage.py +++ b/scripts/cov/core_coverage.py @@ -118,7 +118,7 @@ "airflow/utils/entry_points.py", "airflow/utils/file.py", "airflow/utils/hashlib_wrapper.py", - "airflow/utils/test_template_fields.py", + "airflow/utils/helpers.py", "airflow/utils/json.py", "airflow/utils/log/action_logger.py", "airflow/utils/log/colored_log.py",