This is an automated email from the ASF dual-hosted git repository.
onikolas pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 1ab7ea81a1 uniformize getting hook through cached property in aws
sensors (#29001)
1ab7ea81a1 is described below
commit 1ab7ea81a11073010749103acc97ea92e97dd80a
Author: Raphaƫl Vandon <[email protected]>
AuthorDate: Thu Jan 19 16:03:59 2023 -0800
uniformize getting hook through cached property in aws sensors (#29001)
uniformize getting hook through cached property in aws sensors
---
airflow/providers/amazon/aws/sensors/batch.py | 14 ++++++-----
airflow/providers/amazon/aws/sensors/dms.py | 15 +++++++-----
airflow/providers/amazon/aws/sensors/ec2.py | 8 +++++--
airflow/providers/amazon/aws/sensors/eks.py | 25 ++++++++++++-------
airflow/providers/amazon/aws/sensors/emr.py | 23 +++++++++---------
airflow/providers/amazon/aws/sensors/glacier.py | 8 +++++--
airflow/providers/amazon/aws/sensors/glue.py | 10 +++++---
.../amazon/aws/sensors/glue_catalog_partition.py | 15 +++++++-----
.../providers/amazon/aws/sensors/glue_crawler.py | 18 +++++++-------
airflow/providers/amazon/aws/sensors/quicksight.py | 12 ++++------
airflow/providers/amazon/aws/sensors/rds.py | 9 +++++--
.../amazon/aws/sensors/redshift_cluster.py | 15 +++++++-----
airflow/providers/amazon/aws/sensors/s3.py | 16 +++++++------
airflow/providers/amazon/aws/sensors/sagemaker.py | 28 ++++++++++++----------
airflow/providers/amazon/aws/sensors/sqs.py | 19 +++++++--------
.../providers/amazon/aws/sensors/step_function.py | 15 +++++++-----
.../aws/sensors/test_glue_catalog_partition.py | 2 +-
tests/providers/amazon/aws/sensors/test_sqs.py | 15 ++++++------
18 files changed, 155 insertions(+), 112 deletions(-)
diff --git a/airflow/providers/amazon/aws/sensors/batch.py
b/airflow/providers/amazon/aws/sensors/batch.py
index cbb3791b75..26a5e910a7 100644
--- a/airflow/providers/amazon/aws/sensors/batch.py
+++ b/airflow/providers/amazon/aws/sensors/batch.py
@@ -18,6 +18,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.batch_client import BatchClientHook
@@ -57,10 +59,9 @@ class BatchSensor(BaseSensorOperator):
self.job_id = job_id
self.aws_conn_id = aws_conn_id
self.region_name = region_name
- self.hook: BatchClientHook | None = None
def poke(self, context: Context) -> bool:
- job_description = self.get_hook().get_job_description(self.job_id)
+ job_description = self.hook.get_job_description(self.job_id)
state = job_description["status"]
if state == BatchClientHook.SUCCESS_STATE:
@@ -74,16 +75,17 @@ class BatchSensor(BaseSensorOperator):
raise AirflowException(f"Batch sensor failed. Unknown AWS Batch job
status: {state}")
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> BatchClientHook:
"""Create and return a BatchClientHook"""
- if self.hook:
- return self.hook
+ return self.hook
- self.hook = BatchClientHook(
+ @cached_property
+ def hook(self) -> BatchClientHook:
+ return BatchClientHook(
aws_conn_id=self.aws_conn_id,
region_name=self.region_name,
)
- return self.hook
class BatchComputeEnvironmentSensor(BaseSensorOperator):
diff --git a/airflow/providers/amazon/aws/sensors/dms.py
b/airflow/providers/amazon/aws/sensors/dms.py
index 9a05d77a21..9e2e9ea63c 100644
--- a/airflow/providers/amazon/aws/sensors/dms.py
+++ b/airflow/providers/amazon/aws/sensors/dms.py
@@ -19,6 +19,9 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Iterable, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.dms import DmsHook
from airflow.sensors.base import BaseSensorOperator
@@ -58,18 +61,18 @@ class DmsTaskBaseSensor(BaseSensorOperator):
self.replication_task_arn = replication_task_arn
self.target_statuses: Iterable[str] = target_statuses or []
self.termination_statuses: Iterable[str] = termination_statuses or []
- self.hook: DmsHook | None = None
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> DmsHook:
"""Get DmsHook"""
- if self.hook:
- return self.hook
-
- self.hook = DmsHook(self.aws_conn_id)
return self.hook
+ @cached_property
+ def hook(self) -> DmsHook:
+ return DmsHook(self.aws_conn_id)
+
def poke(self, context: Context):
- status: str | None =
self.get_hook().get_task_status(self.replication_task_arn)
+ status: str | None =
self.hook.get_task_status(self.replication_task_arn)
if not status:
raise AirflowException(
diff --git a/airflow/providers/amazon/aws/sensors/ec2.py
b/airflow/providers/amazon/aws/sensors/ec2.py
index bfdb8fcd41..4377a26444 100644
--- a/airflow/providers/amazon/aws/sensors/ec2.py
+++ b/airflow/providers/amazon/aws/sensors/ec2.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from airflow.compat.functools import cached_property
from airflow.providers.amazon.aws.hooks.ec2 import EC2Hook
from airflow.sensors.base import BaseSensorOperator
@@ -62,8 +63,11 @@ class EC2InstanceStateSensor(BaseSensorOperator):
self.aws_conn_id = aws_conn_id
self.region_name = region_name
+ @cached_property
+ def hook(self):
+ return EC2Hook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
+
def poke(self, context: Context):
- ec2_hook = EC2Hook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
- instance_state =
ec2_hook.get_instance_state(instance_id=self.instance_id)
+ instance_state =
self.hook.get_instance_state(instance_id=self.instance_id)
self.log.info("instance state: %s", instance_state)
return instance_state == self.target_state
diff --git a/airflow/providers/amazon/aws/sensors/eks.py
b/airflow/providers/amazon/aws/sensors/eks.py
index f2ee372151..d275555e3f 100644
--- a/airflow/providers/amazon/aws/sensors/eks.py
+++ b/airflow/providers/amazon/aws/sensors/eks.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.eks import (
ClusterStates,
@@ -98,13 +99,15 @@ class EksClusterStateSensor(BaseSensorOperator):
self.region = region
super().__init__(**kwargs)
- def poke(self, context: Context):
- eks_hook = EksHook(
+ @cached_property
+ def hook(self):
+ return EksHook(
aws_conn_id=self.aws_conn_id,
region_name=self.region,
)
- cluster_state =
eks_hook.get_cluster_state(clusterName=self.cluster_name)
+ def poke(self, context: Context):
+ cluster_state =
self.hook.get_cluster_state(clusterName=self.cluster_name)
self.log.info("Cluster state: %s", cluster_state)
if cluster_state in (CLUSTER_TERMINAL_STATES - {self.target_state}):
# If we reach a terminal state which is not the target state:
@@ -167,13 +170,15 @@ class EksFargateProfileStateSensor(BaseSensorOperator):
self.region = region
super().__init__(**kwargs)
- def poke(self, context: Context):
- eks_hook = EksHook(
+ @cached_property
+ def hook(self):
+ return EksHook(
aws_conn_id=self.aws_conn_id,
region_name=self.region,
)
- fargate_profile_state = eks_hook.get_fargate_profile_state(
+ def poke(self, context: Context):
+ fargate_profile_state = self.hook.get_fargate_profile_state(
clusterName=self.cluster_name,
fargateProfileName=self.fargate_profile_name
)
self.log.info("Fargate profile state: %s", fargate_profile_state)
@@ -238,13 +243,15 @@ class EksNodegroupStateSensor(BaseSensorOperator):
self.region = region
super().__init__(**kwargs)
- def poke(self, context: Context):
- eks_hook = EksHook(
+ @cached_property
+ def hook(self):
+ return EksHook(
aws_conn_id=self.aws_conn_id,
region_name=self.region,
)
- nodegroup_state = eks_hook.get_nodegroup_state(
+ def poke(self, context: Context):
+ nodegroup_state = self.hook.get_nodegroup_state(
clusterName=self.cluster_name, nodegroupName=self.nodegroup_name
)
self.log.info("Nodegroup state: %s", nodegroup_state)
diff --git a/airflow/providers/amazon/aws/sensors/emr.py
b/airflow/providers/amazon/aws/sensors/emr.py
index d1cd0949e0..811846ba45 100644
--- a/airflow/providers/amazon/aws/sensors/emr.py
+++ b/airflow/providers/amazon/aws/sensors/emr.py
@@ -19,6 +19,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Any, Iterable, Sequence
+from deprecated import deprecated
+
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook, EmrHook,
EmrServerlessHook
from airflow.providers.amazon.aws.hooks.s3 import S3Hook
@@ -52,16 +54,15 @@ class EmrBaseSensor(BaseSensorOperator):
self.aws_conn_id = aws_conn_id
self.target_states: Iterable[str] = [] # will be set in subclasses
self.failed_states: Iterable[str] = [] # will be set in subclasses
- self.hook: EmrHook | None = None
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> EmrHook:
- """Get EmrHook"""
- if self.hook:
- return self.hook
-
- self.hook = EmrHook(aws_conn_id=self.aws_conn_id)
return self.hook
+ @cached_property
+ def hook(self) -> EmrHook:
+ return EmrHook(aws_conn_id=self.aws_conn_id)
+
def poke(self, context: Context):
response = self.get_emr_response(context=context)
@@ -332,7 +333,7 @@ class EmrNotebookExecutionSensor(EmrBaseSensor):
self.failed_states = failed_states or self.FAILURE_STATES
def get_emr_response(self, context: Context) -> dict[str, Any]:
- emr_client = self.get_hook().get_conn()
+ emr_client = self.hook.conn
self.log.info("Poking notebook %s", self.notebook_execution_id)
return
emr_client.describe_notebook_execution(NotebookExecutionId=self.notebook_execution_id)
@@ -408,15 +409,15 @@ class EmrJobFlowSensor(EmrBaseSensor):
:return: response
"""
- emr_client = self.get_hook().get_conn()
+ emr_client = self.hook.conn
self.log.info("Poking cluster %s", self.job_flow_id)
response = emr_client.describe_cluster(ClusterId=self.job_flow_id)
log_uri = S3Hook.parse_s3_url(response["Cluster"]["LogUri"])
EmrLogsLink.persist(
context=context,
operator=self,
- region_name=self.get_hook().conn_region_name,
- aws_partition=self.get_hook().conn_partition,
+ region_name=self.hook.conn_region_name,
+ aws_partition=self.hook.conn_partition,
job_flow_id=self.job_flow_id,
log_uri="/".join(log_uri),
)
@@ -497,7 +498,7 @@ class EmrStepSensor(EmrBaseSensor):
:return: response
"""
- emr_client = self.get_hook().get_conn()
+ emr_client = self.hook.conn
self.log.info("Poking step %s on cluster %s", self.step_id,
self.job_flow_id)
return emr_client.describe_step(ClusterId=self.job_flow_id,
StepId=self.step_id)
diff --git a/airflow/providers/amazon/aws/sensors/glacier.py
b/airflow/providers/amazon/aws/sensors/glacier.py
index 857e578327..222027b279 100644
--- a/airflow/providers/amazon/aws/sensors/glacier.py
+++ b/airflow/providers/amazon/aws/sensors/glacier.py
@@ -20,6 +20,7 @@ from __future__ import annotations
from enum import Enum
from typing import TYPE_CHECKING, Any, Sequence
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.glacier import GlacierHook
from airflow.sensors.base import BaseSensorOperator
@@ -81,9 +82,12 @@ class GlacierJobOperationSensor(BaseSensorOperator):
self.poke_interval = poke_interval
self.mode = mode
+ @cached_property
+ def hook(self):
+ return GlacierHook(aws_conn_id=self.aws_conn_id)
+
def poke(self, context: Context) -> bool:
- hook = GlacierHook(aws_conn_id=self.aws_conn_id)
- response = hook.describe_job(vault_name=self.vault_name,
job_id=self.job_id)
+ response = self.hook.describe_job(vault_name=self.vault_name,
job_id=self.job_id)
if response["StatusCode"] == JobStatus.SUCCEEDED.value:
self.log.info("Job status: %s, code status: %s",
response["Action"], response["StatusCode"])
diff --git a/airflow/providers/amazon/aws/sensors/glue.py
b/airflow/providers/amazon/aws/sensors/glue.py
index 87e8f2c249..85db6944d0 100644
--- a/airflow/providers/amazon/aws/sensors/glue.py
+++ b/airflow/providers/amazon/aws/sensors/glue.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.glue import GlueJobHook
from airflow.sensors.base import BaseSensorOperator
@@ -61,10 +62,13 @@ class GlueJobSensor(BaseSensorOperator):
self.errored_states: list[str] = ["FAILED", "STOPPED", "TIMEOUT"]
self.next_log_token: str | None = None
+ @cached_property
+ def hook(self):
+ return GlueJobHook(aws_conn_id=self.aws_conn_id)
+
def poke(self, context: Context):
- hook = GlueJobHook(aws_conn_id=self.aws_conn_id)
self.log.info("Poking for job run status :for Glue Job %s and ID %s",
self.job_name, self.run_id)
- job_state = hook.get_job_state(job_name=self.job_name,
run_id=self.run_id)
+ job_state = self.hook.get_job_state(job_name=self.job_name,
run_id=self.run_id)
job_failed = False
try:
@@ -80,7 +84,7 @@ class GlueJobSensor(BaseSensorOperator):
return False
finally:
if self.verbose:
- self.next_log_token = hook.print_job_logs(
+ self.next_log_token = self.hook.print_job_logs(
job_name=self.job_name,
run_id=self.run_id,
job_failed=job_failed,
diff --git a/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py
b/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py
index 21bf8cb772..d861367466 100644
--- a/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py
+++ b/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py
@@ -19,6 +19,9 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.providers.amazon.aws.hooks.glue_catalog import GlueCatalogHook
from airflow.sensors.base import BaseSensorOperator
@@ -71,7 +74,6 @@ class GlueCatalogPartitionSensor(BaseSensorOperator):
self.table_name = table_name
self.expression = expression
self.database_name = database_name
- self.hook: GlueCatalogHook | None = None
def poke(self, context: Context):
"""Checks for existence of the partition in the AWS Glue Catalog
table"""
@@ -81,12 +83,13 @@ class GlueCatalogPartitionSensor(BaseSensorOperator):
"Poking for table %s. %s, expression %s", self.database_name,
self.table_name, self.expression
)
- return self.get_hook().check_for_partition(self.database_name,
self.table_name, self.expression)
+ return self.hook.check_for_partition(self.database_name,
self.table_name, self.expression)
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> GlueCatalogHook:
"""Gets the GlueCatalogHook"""
- if self.hook:
- return self.hook
-
- self.hook = GlueCatalogHook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
return self.hook
+
+ @cached_property
+ def hook(self) -> GlueCatalogHook:
+ return GlueCatalogHook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
diff --git a/airflow/providers/amazon/aws/sensors/glue_crawler.py
b/airflow/providers/amazon/aws/sensors/glue_crawler.py
index 3032c603c6..6b8b4fcaea 100644
--- a/airflow/providers/amazon/aws/sensors/glue_crawler.py
+++ b/airflow/providers/amazon/aws/sensors/glue_crawler.py
@@ -19,6 +19,9 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.glue_crawler import GlueCrawlerHook
from airflow.sensors.base import BaseSensorOperator
@@ -48,15 +51,13 @@ class GlueCrawlerSensor(BaseSensorOperator):
self.aws_conn_id = aws_conn_id
self.success_statuses = "SUCCEEDED"
self.errored_statuses = ("FAILED", "CANCELLED")
- self.hook: GlueCrawlerHook | None = None
def poke(self, context: Context):
- hook = self.get_hook()
self.log.info("Poking for AWS Glue crawler: %s", self.crawler_name)
- crawler_state = hook.get_crawler(self.crawler_name)["State"]
+ crawler_state = self.hook.get_crawler(self.crawler_name)["State"]
if crawler_state == "READY":
self.log.info("State: %s", crawler_state)
- crawler_status =
hook.get_crawler(self.crawler_name)["LastCrawl"]["Status"]
+ crawler_status =
self.hook.get_crawler(self.crawler_name)["LastCrawl"]["Status"]
if crawler_status == self.success_statuses:
self.log.info("Status: %s", crawler_status)
return True
@@ -65,10 +66,11 @@ class GlueCrawlerSensor(BaseSensorOperator):
else:
return False
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> GlueCrawlerHook:
"""Returns a new or pre-existing GlueCrawlerHook"""
- if self.hook:
- return self.hook
-
- self.hook = GlueCrawlerHook(aws_conn_id=self.aws_conn_id)
return self.hook
+
+ @cached_property
+ def hook(self) -> GlueCrawlerHook:
+ return GlueCrawlerHook(aws_conn_id=self.aws_conn_id)
diff --git a/airflow/providers/amazon/aws/sensors/quicksight.py
b/airflow/providers/amazon/aws/sensors/quicksight.py
index 09cc92cf96..7c71bb24e9 100644
--- a/airflow/providers/amazon/aws/sensors/quicksight.py
+++ b/airflow/providers/amazon/aws/sensors/quicksight.py
@@ -62,8 +62,6 @@ class QuickSightSensor(BaseSensorOperator):
self.aws_conn_id = aws_conn_id
self.success_status = "COMPLETED"
self.errored_statuses = ("FAILED", "CANCELLED")
- self.quicksight_hook: QuickSightHook | None = None
- self.sts_hook: StsHook | None = None
def poke(self, context: Context) -> bool:
"""
@@ -72,11 +70,9 @@ class QuickSightSensor(BaseSensorOperator):
:param context: The task context during execution.
:return: True if it COMPLETED and False if not.
"""
- quicksight_hook = self.get_quicksight_hook
- sts_hook = self.get_sts_hook
self.log.info("Poking for Amazon QuickSight Ingestion ID: %s",
self.ingestion_id)
- aws_account_id = sts_hook.get_account_number()
- quicksight_ingestion_state = quicksight_hook.get_status(
+ aws_account_id = self.sts_hook.get_account_number()
+ quicksight_ingestion_state = self.quicksight_hook.get_status(
aws_account_id, self.data_set_id, self.ingestion_id
)
self.log.info("QuickSight Status: %s", quicksight_ingestion_state)
@@ -85,9 +81,9 @@ class QuickSightSensor(BaseSensorOperator):
return quicksight_ingestion_state == self.success_status
@cached_property
- def get_quicksight_hook(self):
+ def quicksight_hook(self):
return QuickSightHook(aws_conn_id=self.aws_conn_id)
@cached_property
- def get_sts_hook(self):
+ def sts_hook(self):
return StsHook(aws_conn_id=self.aws_conn_id)
diff --git a/airflow/providers/amazon/aws/sensors/rds.py
b/airflow/providers/amazon/aws/sensors/rds.py
index 731c8b5def..50f197ef0c 100644
--- a/airflow/providers/amazon/aws/sensors/rds.py
+++ b/airflow/providers/amazon/aws/sensors/rds.py
@@ -18,6 +18,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowNotFoundException
from airflow.providers.amazon.aws.hooks.rds import RdsHook
from airflow.providers.amazon.aws.utils.rds import RdsDbType
@@ -34,11 +35,15 @@ class RdsBaseSensor(BaseSensorOperator):
ui_fgcolor = "#ffffff"
def __init__(self, *args, aws_conn_id: str = "aws_conn_id", hook_params:
dict | None = None, **kwargs):
- hook_params = hook_params or {}
- self.hook = RdsHook(aws_conn_id=aws_conn_id, **hook_params)
+ self.hook_params = hook_params or {}
+ self.aws_conn_id = aws_conn_id
self.target_statuses: list[str] = []
super().__init__(*args, **kwargs)
+ @cached_property
+ def hook(self):
+ return RdsHook(aws_conn_id=self.aws_conn_id, **self.hook_params)
+
class RdsSnapshotExistenceSensor(RdsBaseSensor):
"""
diff --git a/airflow/providers/amazon/aws/sensors/redshift_cluster.py
b/airflow/providers/amazon/aws/sensors/redshift_cluster.py
index 76f4f90111..93216db13f 100644
--- a/airflow/providers/amazon/aws/sensors/redshift_cluster.py
+++ b/airflow/providers/amazon/aws/sensors/redshift_cluster.py
@@ -18,6 +18,9 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.providers.amazon.aws.hooks.redshift_cluster import RedshiftHook
from airflow.sensors.base import BaseSensorOperator
@@ -51,16 +54,16 @@ class RedshiftClusterSensor(BaseSensorOperator):
self.cluster_identifier = cluster_identifier
self.target_status = target_status
self.aws_conn_id = aws_conn_id
- self.hook: RedshiftHook | None = None
def poke(self, context: Context):
self.log.info("Poking for status : %s\nfor cluster %s",
self.target_status, self.cluster_identifier)
- return self.get_hook().cluster_status(self.cluster_identifier) ==
self.target_status
+ return self.hook.cluster_status(self.cluster_identifier) ==
self.target_status
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> RedshiftHook:
"""Create and return a RedshiftHook"""
- if self.hook:
- return self.hook
-
- self.hook = RedshiftHook(aws_conn_id=self.aws_conn_id)
return self.hook
+
+ @cached_property
+ def hook(self) -> RedshiftHook:
+ return RedshiftHook(aws_conn_id=self.aws_conn_id)
diff --git a/airflow/providers/amazon/aws/sensors/s3.py
b/airflow/providers/amazon/aws/sensors/s3.py
index 57c9393c0c..407a054184 100644
--- a/airflow/providers/amazon/aws/sensors/s3.py
+++ b/airflow/providers/amazon/aws/sensors/s3.py
@@ -23,6 +23,8 @@ import re
from datetime import datetime
from typing import TYPE_CHECKING, Callable, Sequence
+from deprecated import deprecated
+
if TYPE_CHECKING:
from airflow.utils.context import Context
@@ -91,7 +93,6 @@ class S3KeySensor(BaseSensorOperator):
self.check_fn = check_fn
self.aws_conn_id = aws_conn_id
self.verify = verify
- self.hook: S3Hook | None = None
def _check_key(self, key):
bucket_name, key = S3Hook.get_s3_bucket_key(self.bucket_name, key,
"bucket_name", "bucket_key")
@@ -106,7 +107,7 @@ class S3KeySensor(BaseSensorOperator):
"""
if self.wildcard_match:
prefix = re.split(r"[\[\*\?]", key, 1)[0]
- keys = self.get_hook().get_file_metadata(prefix, bucket_name)
+ keys = self.hook.get_file_metadata(prefix, bucket_name)
key_matches = [k for k in keys if fnmatch.fnmatch(k["Key"], key)]
if len(key_matches) == 0:
return False
@@ -114,7 +115,7 @@ class S3KeySensor(BaseSensorOperator):
# Reduce the set of metadata to size only
files = list(map(lambda f: {"Size": f["Size"]}, key_matches))
else:
- obj = self.get_hook().head_object(key, bucket_name)
+ obj = self.hook.head_object(key, bucket_name)
if obj is None:
return False
files = [{"Size": obj["ContentLength"]}]
@@ -130,14 +131,15 @@ class S3KeySensor(BaseSensorOperator):
else:
return all(self._check_key(key) for key in self.bucket_key)
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> S3Hook:
"""Create and return an S3Hook"""
- if self.hook:
- return self.hook
-
- self.hook = S3Hook(aws_conn_id=self.aws_conn_id, verify=self.verify)
return self.hook
+ @cached_property
+ def hook(self) -> S3Hook:
+ return S3Hook(aws_conn_id=self.aws_conn_id, verify=self.verify)
+
@poke_mode_only
class S3KeysUnchangedSensor(BaseSensorOperator):
diff --git a/airflow/providers/amazon/aws/sensors/sagemaker.py
b/airflow/providers/amazon/aws/sensors/sagemaker.py
index f8527fbb2c..b02ea8902b 100644
--- a/airflow/providers/amazon/aws/sensors/sagemaker.py
+++ b/airflow/providers/amazon/aws/sensors/sagemaker.py
@@ -19,6 +19,9 @@ from __future__ import annotations
import time
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.sagemaker import LogState,
SageMakerHook
from airflow.sensors.base import BaseSensorOperator
@@ -41,15 +44,16 @@ class SageMakerBaseSensor(BaseSensorOperator):
super().__init__(**kwargs)
self.aws_conn_id = aws_conn_id
self.resource_type = resource_type # only used for logs, to say what
kind of resource we are sensing
- self.hook: SageMakerHook | None = None
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> SageMakerHook:
"""Get SageMakerHook."""
- if self.hook:
- return self.hook
- self.hook = SageMakerHook(aws_conn_id=self.aws_conn_id)
return self.hook
+ @cached_property
+ def hook(self) -> SageMakerHook:
+ return SageMakerHook(aws_conn_id=self.aws_conn_id)
+
def poke(self, context: Context):
response = self.get_sagemaker_response()
if response["ResponseMetadata"]["HTTPStatusCode"] != 200:
@@ -114,7 +118,7 @@ class SageMakerEndpointSensor(SageMakerBaseSensor):
def get_sagemaker_response(self):
self.log.info("Poking Sagemaker Endpoint %s", self.endpoint_name)
- return self.get_hook().describe_endpoint(self.endpoint_name)
+ return self.hook.describe_endpoint(self.endpoint_name)
def get_failed_reason_from_response(self, response):
return response["FailureReason"]
@@ -150,7 +154,7 @@ class SageMakerTransformSensor(SageMakerBaseSensor):
def get_sagemaker_response(self):
self.log.info("Poking Sagemaker Transform Job %s", self.job_name)
- return self.get_hook().describe_transform_job(self.job_name)
+ return self.hook.describe_transform_job(self.job_name)
def get_failed_reason_from_response(self, response):
return response["FailureReason"]
@@ -186,7 +190,7 @@ class SageMakerTuningSensor(SageMakerBaseSensor):
def get_sagemaker_response(self):
self.log.info("Poking Sagemaker Tuning Job %s", self.job_name)
- return self.get_hook().describe_tuning_job(self.job_name)
+ return self.hook.describe_tuning_job(self.job_name)
def get_failed_reason_from_response(self, response):
return response["FailureReason"]
@@ -243,12 +247,12 @@ class SageMakerTrainingSensor(SageMakerBaseSensor):
def get_sagemaker_response(self):
if self.print_log:
if not self.log_resource_inited:
- self.init_log_resource(self.get_hook())
+ self.init_log_resource(self.hook)
(
self.state,
self.last_description,
self.last_describe_job_call,
- ) = self.get_hook().describe_training_job_with_log(
+ ) = self.hook.describe_training_job_with_log(
self.job_name,
self.positions,
self.stream_names,
@@ -258,7 +262,7 @@ class SageMakerTrainingSensor(SageMakerBaseSensor):
self.last_describe_job_call,
)
else:
- self.last_description =
self.get_hook().describe_training_job(self.job_name)
+ self.last_description =
self.hook.describe_training_job(self.job_name)
status = self.state_from_response(self.last_description)
if (status not in self.non_terminal_states()) and (status not in
self.failed_states()):
billable_time = (
@@ -303,7 +307,7 @@ class SageMakerPipelineSensor(SageMakerBaseSensor):
def get_sagemaker_response(self) -> dict:
self.log.info("Poking Sagemaker Pipeline Execution %s",
self.pipeline_exec_arn)
- return self.get_hook().describe_pipeline_exec(self.pipeline_exec_arn,
self.verbose)
+ return self.hook.describe_pipeline_exec(self.pipeline_exec_arn,
self.verbose)
def state_from_response(self, response: dict) -> str:
return response["PipelineExecutionStatus"]
@@ -335,7 +339,7 @@ class SageMakerAutoMLSensor(SageMakerBaseSensor):
def get_sagemaker_response(self) -> dict:
self.log.info("Poking Sagemaker AutoML Execution %s", self.job_name)
- return self.get_hook()._describe_auto_ml_job(self.job_name)
+ return self.hook._describe_auto_ml_job(self.job_name)
def state_from_response(self, response: dict) -> str:
return response["AutoMLJobStatus"]
diff --git a/airflow/providers/amazon/aws/sensors/sqs.py
b/airflow/providers/amazon/aws/sensors/sqs.py
index 8b09203d2e..6dc032c3fe 100644
--- a/airflow/providers/amazon/aws/sensors/sqs.py
+++ b/airflow/providers/amazon/aws/sensors/sqs.py
@@ -21,9 +21,11 @@ from __future__ import annotations
import json
from typing import TYPE_CHECKING, Any, Collection, Sequence
+from deprecated import deprecated
from jsonpath_ng import parse
from typing_extensions import Literal
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.base_aws import BaseAwsConnection
from airflow.providers.amazon.aws.hooks.sqs import SqsHook
@@ -111,8 +113,6 @@ class SqsSensor(BaseSensorOperator):
self.message_filtering_config = message_filtering_config
- self.hook: SqsHook | None = None
-
def poll_sqs(self, sqs_conn: BaseAwsConnection) -> Collection:
"""
Poll SQS queue to retrieve messages.
@@ -152,13 +152,11 @@ class SqsSensor(BaseSensorOperator):
:param context: the context object
:return: ``True`` if message is available or ``False``
"""
- sqs_conn = self.get_hook().get_conn()
-
message_batch: list[Any] = []
# perform multiple SQS call to retrieve messages in series
for _ in range(self.num_batches):
- messages = self.poll_sqs(sqs_conn=sqs_conn)
+ messages = self.poll_sqs(sqs_conn=self.hook.conn)
if not len(messages):
continue
@@ -173,7 +171,7 @@ class SqsSensor(BaseSensorOperator):
{"Id": message["MessageId"], "ReceiptHandle":
message["ReceiptHandle"]}
for message in messages
]
- response =
sqs_conn.delete_message_batch(QueueUrl=self.sqs_queue, Entries=entries)
+ response =
self.hook.conn.delete_message_batch(QueueUrl=self.sqs_queue, Entries=entries)
if "Successful" not in response:
raise AirflowException(
@@ -185,14 +183,15 @@ class SqsSensor(BaseSensorOperator):
context["ti"].xcom_push(key="messages", value=message_batch)
return True
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> SqsHook:
"""Create and return an SqsHook"""
- if self.hook:
- return self.hook
-
- self.hook = SqsHook(aws_conn_id=self.aws_conn_id)
return self.hook
+ @cached_property
+ def hook(self) -> SqsHook:
+ return SqsHook(aws_conn_id=self.aws_conn_id)
+
def filter_messages(self, messages):
if self.message_filtering == "literal":
return self.filter_messages_literal(messages)
diff --git a/airflow/providers/amazon/aws/sensors/step_function.py
b/airflow/providers/amazon/aws/sensors/step_function.py
index fda6f932d8..2a0c8b10db 100644
--- a/airflow/providers/amazon/aws/sensors/step_function.py
+++ b/airflow/providers/amazon/aws/sensors/step_function.py
@@ -19,6 +19,9 @@ from __future__ import annotations
import json
from typing import TYPE_CHECKING, Sequence
+from deprecated import deprecated
+
+from airflow.compat.functools import cached_property
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.step_function import StepFunctionHook
from airflow.sensors.base import BaseSensorOperator
@@ -68,10 +71,9 @@ class StepFunctionExecutionSensor(BaseSensorOperator):
self.execution_arn = execution_arn
self.aws_conn_id = aws_conn_id
self.region_name = region_name
- self.hook: StepFunctionHook | None = None
def poke(self, context: Context):
- execution_status =
self.get_hook().describe_execution(self.execution_arn)
+ execution_status = self.hook.describe_execution(self.execution_arn)
state = execution_status["status"]
output = json.loads(execution_status["output"]) if "output" in
execution_status else None
@@ -85,10 +87,11 @@ class StepFunctionExecutionSensor(BaseSensorOperator):
self.xcom_push(context, "output", output)
return True
+ @deprecated(reason="use `hook` property instead.")
def get_hook(self) -> StepFunctionHook:
"""Create and return a StepFunctionHook"""
- if self.hook:
- return self.hook
-
- self.hook = StepFunctionHook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
return self.hook
+
+ @cached_property
+ def hook(self) -> StepFunctionHook:
+ return StepFunctionHook(aws_conn_id=self.aws_conn_id,
region_name=self.region_name)
diff --git a/tests/providers/amazon/aws/sensors/test_glue_catalog_partition.py
b/tests/providers/amazon/aws/sensors/test_glue_catalog_partition.py
index ab89b09ecc..e5726c58a2 100644
--- a/tests/providers/amazon/aws/sensors/test_glue_catalog_partition.py
+++ b/tests/providers/amazon/aws/sensors/test_glue_catalog_partition.py
@@ -76,7 +76,7 @@ class TestGlueCatalogPartitionSensor:
)
# We're mocking all actual AWS calls and don't need a connection. This
# avoids an Airflow warning about connection cannot be found.
- op.get_hook().get_connection = lambda _: None
+ op.hook.get_connection = lambda _: None
op.poke({})
assert op.hook.region_name == region_name
diff --git a/tests/providers/amazon/aws/sensors/test_sqs.py
b/tests/providers/amazon/aws/sensors/test_sqs.py
index 73a3ce8278..e6175c61d3 100644
--- a/tests/providers/amazon/aws/sensors/test_sqs.py
+++ b/tests/providers/amazon/aws/sensors/test_sqs.py
@@ -19,6 +19,7 @@ from __future__ import annotations
import json
from unittest import mock
+from unittest.mock import patch
import pytest
from moto import mock_sqs
@@ -70,7 +71,7 @@ class TestSqsSensor:
assert self.mock_context["ti"].method_calls == context_calls, "context
call should be same"
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_delete_raise_airflow_exception(self, mock_conn):
message = {
"Messages": [
@@ -103,7 +104,7 @@ class TestSqsSensor:
assert "Delete SQS Messages failed" in ctx.value.args[0]
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_receive_raise_exception(self, mock_conn):
mock_conn.return_value.receive_message.side_effect = Exception("test
exception")
with pytest.raises(Exception) as ctx:
@@ -111,7 +112,7 @@ class TestSqsSensor:
assert "test exception" in ctx.value.args[0]
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_visibility_timeout(self, mock_conn):
# Check without visibility_timeout parameter
self.sqs_hook.create_queue(QUEUE_NAME)
@@ -155,7 +156,7 @@ class TestSqsSensor:
sensor.poke(self.mock_context)
assert "Override this method to define custom filters" in
ctx.value.args[0]
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_message_filtering_literal_values(self, mock_conn):
self.sqs_hook.create_queue(QUEUE_NAME)
matching = [{"id": 11, "body": "a matching message"}]
@@ -194,7 +195,7 @@ class TestSqsSensor:
]
mock_conn.assert_has_calls(calls_delete_message_batch)
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_message_filtering_jsonpath(self, mock_conn):
self.sqs_hook.create_queue(QUEUE_NAME)
matching = [
@@ -240,7 +241,7 @@ class TestSqsSensor:
]
mock_conn.assert_has_calls(calls_delete_message_batch)
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_message_filtering_jsonpath_values(self, mock_conn):
self.sqs_hook.create_queue(QUEUE_NAME)
matching = [
@@ -288,7 +289,7 @@ class TestSqsSensor:
]
mock_conn.assert_has_calls(calls_delete_message_batch)
- @mock.patch.object(SqsHook, "get_conn")
+ @patch("airflow.providers.amazon.aws.hooks.sqs.SqsHook.conn",
new_callable=mock.PropertyMock)
def test_poke_do_not_delete_message_on_received(self, mock_conn):
self.sqs_hook.create_queue(QUEUE_NAME)