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)


Reply via email to