From 435147ff1168918047496bfec3a976d145586cb6 Mon Sep 17 00:00:00 2001 From: Dmitry Romanenko Date: Tue, 23 Nov 2021 19:41:56 +0000 Subject: [PATCH 1/7] Adjust built-in base_aws methods to avoid Deprecation warnings (#19725) Co-authored-by: Tzu-ping Chung --- airflow/providers/amazon/aws/hooks/base_aws.py | 8 +++++--- airflow/providers/amazon/aws/hooks/glue.py | 3 ++- airflow/providers/amazon/aws/hooks/s3.py | 7 +++++-- airflow/providers/postgres/hooks/postgres.py | 3 ++- 4 files changed, 14 insertions(+), 7 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/base_aws.py b/airflow/providers/amazon/aws/hooks/base_aws.py index 171437cc740c5..90b14733abb39 100644 --- a/airflow/providers/amazon/aws/hooks/base_aws.py +++ b/airflow/providers/amazon/aws/hooks/base_aws.py @@ -491,9 +491,9 @@ def conn(self) -> Union[boto3.client, boto3.resource]: :rtype: Union[boto3.client, boto3.resource] """ if self.client_type: - return self.get_client_type(self.client_type, region_name=self.region_name) + return self.get_client_type(region_name=self.region_name) elif self.resource_type: - return self.get_resource_type(self.resource_type, region_name=self.region_name) + return self.get_resource_type(region_name=self.region_name) else: # Rare possibility - subclasses have not specified a client_type or resource_type raise NotImplementedError('Could not get boto3 connection!') @@ -539,7 +539,9 @@ def expand_role(self, role: str) -> str: if "/" in role: return role else: - return self.get_client_type("iam").get_role(RoleName=role)["Role"]["Arn"] + session, endpoint_url = self._get_credentials() + _client = session.client('iam', endpoint_url=endpoint_url, config=self.config, verify=self.verify) + return _client.get_role(RoleName=role)["Role"]["Arn"] @staticmethod def retry(should_retry: Callable[[Exception], bool]): diff --git a/airflow/providers/amazon/aws/hooks/glue.py b/airflow/providers/amazon/aws/hooks/glue.py index a5e43278dc35a..2cf048000aa4c 100644 --- a/airflow/providers/amazon/aws/hooks/glue.py +++ b/airflow/providers/amazon/aws/hooks/glue.py @@ -85,7 +85,8 @@ def list_jobs(self) -> List: def get_iam_execution_role(self) -> Dict: """:return: iam role for job execution""" - iam_client = self.get_client_type('iam', self.region_name) + session, endpoint_url = self._get_credentials(self.region_name) + iam_client = session.client('iam', endpoint_url=endpoint_url, config=self.config, verify=self.verify) try: glue_execution_role = iam_client.get_role(RoleName=self.role_name) diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index 53b2747104380..569cbade3e18c 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -173,7 +173,8 @@ def get_bucket(self, bucket_name: Optional[str] = None) -> str: :return: the bucket object to the bucket name. :rtype: boto3.S3.Bucket """ - s3_resource = self.get_resource_type('s3') + session, endpoint_url = self._get_credentials() + s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return s3_resource.Bucket(bucket_name) @provide_bucket_name @@ -340,7 +341,9 @@ def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: :return: the key object from the bucket :rtype: boto3.s3.Object """ - obj = self.get_resource_type('s3').Object(bucket_name, key) + session, endpoint_url = self._get_credentials() + s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) + obj = s3_resource.Object(bucket_name, key) obj.load() return obj diff --git a/airflow/providers/postgres/hooks/postgres.py b/airflow/providers/postgres/hooks/postgres.py index 93688d74cd67c..a68046fda03ae 100644 --- a/airflow/providers/postgres/hooks/postgres.py +++ b/airflow/providers/postgres/hooks/postgres.py @@ -184,7 +184,8 @@ def get_iam_token(self, conn: Connection) -> Tuple[str, str, int]: # Pull the custer-identifier from the beginning of the Redshift URL # ex. my-cluster.ccdre4hpd39h.us-east-1.redshift.amazonaws.com returns my-cluster cluster_identifier = conn.extra_dejson.get('cluster-identifier', conn.host.split('.')[0]) - client = aws_hook.get_client_type('redshift') + session, endpoint_url = aws_hook._get_credentials() + client = session.client('redshift', endpoint_url=endpoint_url, config=aws_hook.config, verify=aws_hook.verify) cluster_creds = client.get_cluster_credentials( DbUser=conn.login, DbName=self.schema or conn.schema, From dae144f624eac2624865b7733d5b189148dd1daa Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 24 Nov 2021 03:59:56 +0800 Subject: [PATCH 2/7] Fix missing region_name arguments in AWS provider --- airflow/providers/amazon/aws/hooks/base_aws.py | 2 +- airflow/providers/amazon/aws/hooks/s3.py | 18 ++++++++++++++---- .../amazon/aws/hooks/test_batch_client.py | 2 +- .../amazon/aws/hooks/test_batch_waiters.py | 2 +- .../amazon/aws/operators/test_batch.py | 2 +- 5 files changed, 18 insertions(+), 8 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/base_aws.py b/airflow/providers/amazon/aws/hooks/base_aws.py index 90b14733abb39..1199b43da05a9 100644 --- a/airflow/providers/amazon/aws/hooks/base_aws.py +++ b/airflow/providers/amazon/aws/hooks/base_aws.py @@ -539,7 +539,7 @@ def expand_role(self, role: str) -> str: if "/" in role: return role else: - session, endpoint_url = self._get_credentials() + session, endpoint_url = self._get_credentials(region_name=None) _client = session.client('iam', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return _client.get_role(RoleName=role)["Role"]["Arn"] diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index 569cbade3e18c..e3c58511a26a9 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -164,7 +164,7 @@ def check_for_bucket(self, bucket_name: Optional[str] = None) -> bool: return False @provide_bucket_name - def get_bucket(self, bucket_name: Optional[str] = None) -> str: + def get_bucket(self, bucket_name: Optional[str] = None, *, region_name: Optional[str] = None) -> str: """ Returns a boto3.S3.Bucket object @@ -172,8 +172,10 @@ def get_bucket(self, bucket_name: Optional[str] = None) -> str: :type bucket_name: str :return: the bucket object to the bucket name. :rtype: boto3.S3.Bucket + :param region_name: The name of the aws region in which to get the bucket. + :type region_name: str """ - session, endpoint_url = self._get_credentials() + session, endpoint_url = self._get_credentials(region_name=region_name) s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return s3_resource.Bucket(bucket_name) @@ -330,7 +332,13 @@ def check_for_key(self, key: str, bucket_name: Optional[str] = None) -> bool: @provide_bucket_name @unify_bucket_name_and_key - def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: + def get_key( + self, + key: str, + bucket_name: Optional[str] = None, + *, + region_name: Optional[str] = None, + ) -> S3Transfer: """ Returns a boto3.s3.Object @@ -338,10 +346,12 @@ def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: :type key: str :param bucket_name: the name of the bucket :type bucket_name: str + :param region_name: the name of the bucket's region + :type region_name: str :return: the key object from the bucket :rtype: boto3.s3.Object """ - session, endpoint_url = self._get_credentials() + session, endpoint_url = self._get_credentials(region_name=region_name) s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) obj = s3_resource.Object(bucket_name, key) obj.load() diff --git a/tests/providers/amazon/aws/hooks/test_batch_client.py b/tests/providers/amazon/aws/hooks/test_batch_client.py index dc931426b599f..9cd99c4c95eb2 100644 --- a/tests/providers/amazon/aws/hooks/test_batch_client.py +++ b/tests/providers/amazon/aws/hooks/test_batch_client.py @@ -68,7 +68,7 @@ def test_init(self): assert self.batch_client.aws_conn_id == 'airflow_test' assert self.batch_client.client == self.client_mock - self.get_client_type_mock.assert_called_once_with("batch", region_name=AWS_REGION) + self.get_client_type_mock.assert_called_once_with(region_name=AWS_REGION) def test_wait_for_job_with_success(self): self.client_mock.describe_jobs.return_value = {"jobs": [{"jobId": JOB_ID, "status": "SUCCEEDED"}]} diff --git a/tests/providers/amazon/aws/hooks/test_batch_waiters.py b/tests/providers/amazon/aws/hooks/test_batch_waiters.py index 31c12d3794b42..b5fb4d0626930 100644 --- a/tests/providers/amazon/aws/hooks/test_batch_waiters.py +++ b/tests/providers/amazon/aws/hooks/test_batch_waiters.py @@ -333,7 +333,7 @@ def setUp(self, get_client_type_mock): # init the mock client self.client_mock = self.batch_waiters.client - get_client_type_mock.assert_called_once_with("batch", region_name=self.region_name) + get_client_type_mock.assert_called_once_with(region_name=self.region_name) # don't pause in these unit tests self.mock_delay = mock.Mock(return_value=None) diff --git a/tests/providers/amazon/aws/operators/test_batch.py b/tests/providers/amazon/aws/operators/test_batch.py index af3e0d7a1afd8..9223f517d53d7 100644 --- a/tests/providers/amazon/aws/operators/test_batch.py +++ b/tests/providers/amazon/aws/operators/test_batch.py @@ -95,7 +95,7 @@ def test_init(self): assert self.batch.hook.client == self.client_mock assert self.batch.tags == {} - self.get_client_type_mock.assert_called_once_with("batch", region_name="eu-west-1") + self.get_client_type_mock.assert_called_once_with(region_name="eu-west-1") def test_template_fields_overrides(self): assert self.batch.template_fields == ( From 6e265ebece2b185d1ed52382542d5f0d1e8d34ac Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 24 Nov 2021 04:14:31 +0800 Subject: [PATCH 3/7] Fix Black --- airflow/providers/postgres/hooks/postgres.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/airflow/providers/postgres/hooks/postgres.py b/airflow/providers/postgres/hooks/postgres.py index a68046fda03ae..96c4da10bc0a3 100644 --- a/airflow/providers/postgres/hooks/postgres.py +++ b/airflow/providers/postgres/hooks/postgres.py @@ -185,7 +185,12 @@ def get_iam_token(self, conn: Connection) -> Tuple[str, str, int]: # ex. my-cluster.ccdre4hpd39h.us-east-1.redshift.amazonaws.com returns my-cluster cluster_identifier = conn.extra_dejson.get('cluster-identifier', conn.host.split('.')[0]) session, endpoint_url = aws_hook._get_credentials() - client = session.client('redshift', endpoint_url=endpoint_url, config=aws_hook.config, verify=aws_hook.verify) + client = session.client( + "redshift", + endpoint_url=endpoint_url, + config=aws_hook.config, + verify=aws_hook.verify, + ) cluster_creds = client.get_cluster_credentials( DbUser=conn.login, DbName=self.schema or conn.schema, From 14dcf755e4cf783b25efb741bb05e4d60779551a Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 24 Nov 2021 04:47:52 +0800 Subject: [PATCH 4/7] Revert region_name additions from public API --- airflow/providers/amazon/aws/hooks/s3.py | 18 ++++-------------- 1 file changed, 4 insertions(+), 14 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index e3c58511a26a9..69a5bee63a78a 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -164,7 +164,7 @@ def check_for_bucket(self, bucket_name: Optional[str] = None) -> bool: return False @provide_bucket_name - def get_bucket(self, bucket_name: Optional[str] = None, *, region_name: Optional[str] = None) -> str: + def get_bucket(self, bucket_name: Optional[str] = None) -> str: """ Returns a boto3.S3.Bucket object @@ -172,10 +172,8 @@ def get_bucket(self, bucket_name: Optional[str] = None, *, region_name: Optional :type bucket_name: str :return: the bucket object to the bucket name. :rtype: boto3.S3.Bucket - :param region_name: The name of the aws region in which to get the bucket. - :type region_name: str """ - session, endpoint_url = self._get_credentials(region_name=region_name) + session, endpoint_url = self._get_credentials(region_name=None) s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return s3_resource.Bucket(bucket_name) @@ -332,13 +330,7 @@ def check_for_key(self, key: str, bucket_name: Optional[str] = None) -> bool: @provide_bucket_name @unify_bucket_name_and_key - def get_key( - self, - key: str, - bucket_name: Optional[str] = None, - *, - region_name: Optional[str] = None, - ) -> S3Transfer: + def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: """ Returns a boto3.s3.Object @@ -346,12 +338,10 @@ def get_key( :type key: str :param bucket_name: the name of the bucket :type bucket_name: str - :param region_name: the name of the bucket's region - :type region_name: str :return: the key object from the bucket :rtype: boto3.s3.Object """ - session, endpoint_url = self._get_credentials(region_name=region_name) + session, endpoint_url = self._get_credentials(region_name=None) s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) obj = s3_resource.Object(bucket_name, key) obj.load() From 9fbe186811e698870e074ed064649b54ac9479fc Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 24 Nov 2021 05:35:23 +0800 Subject: [PATCH 5/7] Fix region_name in Postgres hook --- airflow/providers/amazon/aws/hooks/base_aws.py | 7 +++++-- airflow/providers/amazon/aws/hooks/s3.py | 4 ++-- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/base_aws.py b/airflow/providers/amazon/aws/hooks/base_aws.py index 1199b43da05a9..f5042ba13bb0f 100644 --- a/airflow/providers/amazon/aws/hooks/base_aws.py +++ b/airflow/providers/amazon/aws/hooks/base_aws.py @@ -392,7 +392,10 @@ def __init__( if not (self.client_type or self.resource_type): raise AirflowException('Either client_type or resource_type must be provided.') - def _get_credentials(self, region_name: Optional[str]) -> Tuple[boto3.session.Session, Optional[str]]: + def _get_credentials( + self, + region_name: Optional[str] = None, + ) -> Tuple[boto3.session.Session, Optional[str]]: if not self.aws_conn_id: session = boto3.session.Session(region_name=region_name) @@ -539,7 +542,7 @@ def expand_role(self, role: str) -> str: if "/" in role: return role else: - session, endpoint_url = self._get_credentials(region_name=None) + session, endpoint_url = self._get_credentials() _client = session.client('iam', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return _client.get_role(RoleName=role)["Role"]["Arn"] diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index 69a5bee63a78a..569cbade3e18c 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -173,7 +173,7 @@ def get_bucket(self, bucket_name: Optional[str] = None) -> str: :return: the bucket object to the bucket name. :rtype: boto3.S3.Bucket """ - session, endpoint_url = self._get_credentials(region_name=None) + session, endpoint_url = self._get_credentials() s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) return s3_resource.Bucket(bucket_name) @@ -341,7 +341,7 @@ def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: :return: the key object from the bucket :rtype: boto3.s3.Object """ - session, endpoint_url = self._get_credentials(region_name=None) + session, endpoint_url = self._get_credentials() s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) obj = s3_resource.Object(bucket_name, key) obj.load() From d2f59f3155b0a755d37b78f6453c73c5e9c94502 Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 24 Nov 2021 05:37:42 +0800 Subject: [PATCH 6/7] Replace boto3 session.client with session.resource --- airflow/providers/amazon/aws/hooks/s3.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index 569cbade3e18c..6ca4b41f66085 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -174,7 +174,12 @@ def get_bucket(self, bucket_name: Optional[str] = None) -> str: :rtype: boto3.S3.Bucket """ session, endpoint_url = self._get_credentials() - s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) + s3_resource = session.resource( + "s3", + endpoint_url=endpoint_url, + config=self.config, + verify=self.verify, + ) return s3_resource.Bucket(bucket_name) @provide_bucket_name @@ -342,7 +347,12 @@ def get_key(self, key: str, bucket_name: Optional[str] = None) -> S3Transfer: :rtype: boto3.s3.Object """ session, endpoint_url = self._get_credentials() - s3_resource = session.client('s3', endpoint_url=endpoint_url, config=self.config, verify=self.verify) + s3_resource = session.resource( + "s3", + endpoint_url=endpoint_url, + config=self.config, + verify=self.verify, + ) obj = s3_resource.Object(bucket_name, key) obj.load() return obj From 804a89549468c2e21d5cbb9949a51a2432051a6e Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Fri, 26 Nov 2021 16:06:36 +0800 Subject: [PATCH 7/7] Rewrite Postgres Redshift IAM mock for new creds --- .../providers/postgres/hooks/test_postgres.py | 31 ++++++++++++------- 1 file changed, 19 insertions(+), 12 deletions(-) diff --git a/tests/providers/postgres/hooks/test_postgres.py b/tests/providers/postgres/hooks/test_postgres.py index c22a7223cefe1..9d9392da4323b 100644 --- a/tests/providers/postgres/hooks/test_postgres.py +++ b/tests/providers/postgres/hooks/test_postgres.py @@ -113,31 +113,38 @@ def test_get_conn_extra(self, mock_connect): ) @mock.patch('airflow.providers.postgres.hooks.postgres.psycopg2.connect') - @mock.patch('airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook.get_client_type') - def test_get_conn_rds_iam_redshift(self, mock_client, mock_connect): + def test_get_conn_rds_iam_redshift(self, mock_connect): self.connection.extra = '{"iam":true, "redshift":true, "cluster-identifier": "different-identifier"}' self.connection.host = 'cluster-identifier.ccdfre4hpd39h.us-east-1.redshift.amazonaws.com' login = f'IAM:{self.connection.login}' - mock_client.return_value.get_cluster_credentials.return_value = { - 'DbPassword': 'aws_token', - 'DbUser': login, - } - self.db_hook.get_conn() + + mock_session = mock.Mock() + mock_get_cluster_credentials = mock_session.client.return_value.get_cluster_credentials + mock_get_cluster_credentials.return_value = {'DbPassword': 'aws_token', 'DbUser': login} + + aws_get_credentials_patcher = mock.patch( + "airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook._get_credentials", + return_value=(mock_session, None), + ) get_cluster_credentials_call = mock.call( DbUser=self.connection.login, DbName=self.connection.schema, ClusterIdentifier="different-identifier", AutoCreate=False, ) - mock_client.return_value.get_cluster_credentials.assert_has_calls([get_cluster_credentials_call]) + + with aws_get_credentials_patcher: + self.db_hook.get_conn() + assert mock_get_cluster_credentials.mock_calls == [get_cluster_credentials_call] mock_connect.assert_called_once_with( user=login, password='aws_token', host=self.connection.host, dbname='schema', port=5439 ) + # Verify that the connection object has not been mutated. - self.db_hook.get_conn() - mock_client.return_value.get_cluster_credentials.assert_has_calls( - [get_cluster_credentials_call, get_cluster_credentials_call] - ) + mock_get_cluster_credentials.reset_mock() + with aws_get_credentials_patcher: + self.db_hook.get_conn() + assert mock_get_cluster_credentials.mock_calls == [get_cluster_credentials_call] def test_get_uri_from_connection_without_schema_override(self): self.db_hook.get_connection = mock.MagicMock(