From 0083fc390f0f2622061b7360091777cb0da0220d Mon Sep 17 00:00:00 2001 From: Alex Begg Date: Wed, 10 Nov 2021 12:32:14 -0800 Subject: [PATCH 1/5] Do not require all extras for SalesforceHook and allow extras without prefix --- .../providers/salesforce/hooks/salesforce.py | 39 ++- .../salesforce/hooks/test_salesforce.py | 273 +++++++++++++----- 2 files changed, 223 insertions(+), 89 deletions(-) diff --git a/airflow/providers/salesforce/hooks/salesforce.py b/airflow/providers/salesforce/hooks/salesforce.py index 36426e8522fbd..275c79a22b651 100644 --- a/airflow/providers/salesforce/hooks/salesforce.py +++ b/airflow/providers/salesforce/hooks/salesforce.py @@ -131,22 +131,39 @@ def get_conn(self) -> api.Salesforce: if not self.conn: connection = self.get_connection(self.conn_id) extras = connection.extra_dejson + # all extras below (besides the version one) are explicitly defaulted to None + # because simple-salesforce has a built-in authentication-choosing method that + # relies on which arguments are None and without "or None" setting this connection + # in the UI will result in the blank extras being empty strings instead of None, + # which would break the connection if "get" was used on its own. self.conn = Salesforce( username=connection.login, password=connection.password, - security_token=extras["extra__salesforce__security_token"] or None, - domain=extras["extra__salesforce__domain"] or None, + security_token=extras.get('security_token') + or extras.get('extra__salesforce__security_token') + or None, + domain=extras.get('domain') or extras.get('extra__salesforce__domain') or None, session_id=self.session_id, - instance=extras["extra__salesforce__instance"] or None, - instance_url=extras["extra__salesforce__instance_url"] or None, - organizationId=extras["extra__salesforce__organization_id"] or None, - version=extras["extra__salesforce__version"] or api.DEFAULT_API_VERSION, - proxies=extras["extra__salesforce__proxies"] or None, + instance=extras.get('instance') or extras.get('extra__salesforce__instance') or None, + instance_url=extras.get('instance_url') + or extras.get('extra__salesforce__instance_url') + or None, + organizationId=extras.get('organization_id') + or extras.get('extra__salesforce__organization_id') + or None, + version=extras.get('version') + or extras.get('extra__salesforce__version') + or api.DEFAULT_API_VERSION, + proxies=extras.get('proxies') or extras.get('extra__salesforce__proxies') or None, session=self.session, - client_id=extras["extra__salesforce__client_id"] or None, - consumer_key=extras["extra__salesforce__consumer_key"] or None, - privatekey_file=extras["extra__salesforce__private_key_file_path"] or None, - privatekey=extras["extra__salesforce__private_key"] or None, + client_id=extras.get('client_id') or extras.get('extra__salesforce__client_id') or None, + consumer_key=extras.get('consumer_key') + or extras.get('extra__salesforce__consumer_key') + or None, + privatekey_file=extras.get('private_key_file_path') + or extras.get('extra__salesforce__private_key_file_path') + or None, + privatekey=extras.get('private_key') or extras.get('extra__salesforce__private_key') or None, ) return self.conn diff --git a/tests/providers/salesforce/hooks/test_salesforce.py b/tests/providers/salesforce/hooks/test_salesforce.py index 766fa898178b4..03c71e76731c7 100644 --- a/tests/providers/salesforce/hooks/test_salesforce.py +++ b/tests/providers/salesforce/hooks/test_salesforce.py @@ -23,6 +23,7 @@ import pandas as pd import pytest from numpy import nan +from parameterized import parameterized from requests import Session as request_session from simple_salesforce import Salesforce, api @@ -48,8 +49,52 @@ def test_get_conn_exists(self): assert self.salesforce_hook.conn.return_value is not None + @parameterized.expand( + [ + ( + "all extras", + ''' + { + "extra__salesforce__client_id": "my_client", + "extra__salesforce__consumer_key": "", + "extra__salesforce__domain": "test", + "extra__salesforce__instance": "", + "extra__salesforce__instance_url": "", + "extra__salesforce__organization_id": "", + "extra__salesforce__private_key": "", + "extra__salesforce__private_key_file_path": "", + "extra__salesforce__proxies": "", + "extra__salesforce__security_token": "token", + "extra__salesforce__version": "42.0" + } + ''', + ), + ( + "required extras", + ''' + { + "extra__salesforce__client_id": "my_client", + "extra__salesforce__domain": "test", + "extra__salesforce__security_token": "token", + "extra__salesforce__version": "42.0" + } + ''', + ), + ( + "required extras no prefix", + ''' + { + "client_id": "my_client", + "domain": "test", + "security_token": "token", + "version": "42.0" + } + ''', + ), + ] + ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_password_auth(self, mock_salesforce): + def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): """ Testing mock password authentication to Salesforce. Users should provide a username, password, and security token in the Connection. Providing a client ID, Salesforce API version, proxy mapping, and @@ -61,21 +106,7 @@ def test_get_conn_password_auth(self, mock_salesforce): conn_type="salesforce", login=None, password=None, - extra=''' - { - "extra__salesforce__client_id": "my_client", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "test", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "token", - "extra__salesforce__version": "42.0" - } - ''', + extra=extras_string, ) TestSalesforceHook._insert_conn_db_entry(password_auth_conn.conn_id, password_auth_conn) @@ -86,23 +117,67 @@ def test_get_conn_password_auth(self, mock_salesforce): mock_salesforce.assert_called_once_with( username=password_auth_conn.login, password=password_auth_conn.password, - security_token=extras["extra__salesforce__security_token"], - domain=extras["extra__salesforce__domain"], + security_token=extras.get('extra__salesforce__security_token') or extras.get('security_token'), + domain=extras.get('extra__salesforce__domain') or extras.get('domain'), session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras["extra__salesforce__version"], + version=extras.get('extra__salesforce__version') or extras.get('version'), proxies=None, session=None, - client_id=extras["extra__salesforce__client_id"], + client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), consumer_key=None, privatekey_file=None, privatekey=None, ) + @parameterized.expand( + [ + ( + "all extras", + ''' + { + "extra__salesforce__client_id": "my_client2", + "extra__salesforce__consumer_key": "", + "extra__salesforce__domain": "test", + "extra__salesforce__instance": "", + "extra__salesforce__instance_url": "https://my.salesforce.com", + "extra__salesforce__organization_id": "", + "extra__salesforce__private_key": "", + "extra__salesforce__private_key_file_path": "", + "extra__salesforce__proxies": "", + "extra__salesforce__security_token": "", + "extra__salesforce__version": "29.0" + } + ''', + ), + ( + "required extras", + ''' + { + "extra__salesforce__client_id": "my_client2", + "extra__salesforce__domain": "test", + "extra__salesforce__instance_url": "https://my.salesforce.com", + "extra__salesforce__version": "29.0" + } + ''', + ), + ( + "required extras no prefix", + ''' + { + "client_id": "my_client2", + "domain": "test", + "instance_url": "https://my.salesforce.com", + "version": "29.0" + } + ''', + ), + ] + ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_direct_session_access(self, mock_salesforce): + def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce): """ Testing mock direct session access to Salesforce. Users should provide an instance (or instance URL) in the Connection and set a `session_id` value when calling `SalesforceHook`. @@ -115,21 +190,7 @@ def test_get_conn_direct_session_access(self, mock_salesforce): conn_type="salesforce", login=None, password=None, - extra=''' - { - "extra__salesforce__client_id": "my_client2", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "test", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "https://my.salesforce.com", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "29.0" - } - ''', + extra=extras_string, ) TestSalesforceHook._insert_conn_db_entry(direct_access_conn.conn_id, direct_access_conn) @@ -145,22 +206,68 @@ def test_get_conn_direct_session_access(self, mock_salesforce): username=direct_access_conn.login, password=direct_access_conn.password, security_token=None, - domain=extras["extra__salesforce__domain"], + domain=extras.get('extra__salesforce__domain') or extras.get('domain'), session_id=self.salesforce_hook.session_id, instance=None, - instance_url=extras["extra__salesforce__instance_url"], + instance_url=extras.get('extra__salesforce__instance_url') or extras.get('instance_url'), organizationId=None, - version=extras["extra__salesforce__version"], + version=extras.get('extra__salesforce__version') or extras.get('version'), proxies=None, session=self.salesforce_hook.session, - client_id=extras["extra__salesforce__client_id"], + client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), consumer_key=None, privatekey_file=None, privatekey=None, ) + @parameterized.expand( + [ + ( + "all extras", + ''' + { + "extra__salesforce__client_id": "my_client3", + "extra__salesforce__consumer_key": "consumer_key", + "extra__salesforce__domain": "login", + "extra__salesforce__instance": "", + "extra__salesforce__instance_url": "", + "extra__salesforce__organization_id": "", + "extra__salesforce__private_key": "private_key", + "extra__salesforce__private_key_file_path": "", + "extra__salesforce__proxies": "", + "extra__salesforce__security_token": "", + "extra__salesforce__version": "34.0" + } + ''', + ), + ( + "required extras", + ''' + { + "extra__salesforce__client_id": "my_client3", + "extra__salesforce__consumer_key": "consumer_key", + "extra__salesforce__domain": "login", + "extra__salesforce__private_key": "private_key", + "extra__salesforce__version": "34.0" + } + ''', + ), + ( + "required extras no prefix", + ''' + { + "client_id": "my_client3", + "consumer_key": "consumer_key", + "domain": "login", + "private_key": "private_key", + "version": "34.0" + } + ''', + ), + ] + ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_jwt_auth(self, mock_salesforce): + def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): """ Testing mock JWT bearer authentication to Salesforce. Users should provide consumer key and private key (or path to a private key) in the Connection. Providing a client ID, Salesforce API version, proxy @@ -173,21 +280,7 @@ def test_get_conn_jwt_auth(self, mock_salesforce): conn_type="salesforce", login=None, password=None, - extra=''' - { - "extra__salesforce__client_id": "my_client3", - "extra__salesforce__consumer_key": "consumer_key", - "extra__salesforce__domain": "login", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "private_key", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "34.0" - } - ''', + extra=extras_string, ) TestSalesforceHook._insert_conn_db_entry(jwt_auth_conn.conn_id, jwt_auth_conn) @@ -199,22 +292,60 @@ def test_get_conn_jwt_auth(self, mock_salesforce): username=jwt_auth_conn.login, password=jwt_auth_conn.password, security_token=None, - domain=extras["extra__salesforce__domain"], + domain=extras.get('extra__salesforce__domain') or extras.get('domain'), session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras["extra__salesforce__version"], + version=extras.get('extra__salesforce__version') or extras.get('version'), proxies=None, session=None, - client_id=extras["extra__salesforce__client_id"], - consumer_key=extras["extra__salesforce__consumer_key"], + client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), + consumer_key=extras.get('extra__salesforce__consumer_key') or extras.get('consumer_key'), privatekey_file=None, - privatekey=extras["extra__salesforce__private_key"], + privatekey=extras.get('extra__salesforce__private_key') or extras.get('private_key'), ) + @parameterized.expand( + [ + ( + "all extras", + ''' + { + "extra__salesforce__client_id": "", + "extra__salesforce__consumer_key": "", + "extra__salesforce__domain": "", + "extra__salesforce__instance": "", + "extra__salesforce__instance_url": "", + "extra__salesforce__organization_id": "my_organization", + "extra__salesforce__private_key": "", + "extra__salesforce__private_key_file_path": "", + "extra__salesforce__proxies": "", + "extra__salesforce__security_token": "", + "extra__salesforce__version": "" + } + ''', + ), + ( + "required extras", + ''' + { + "extra__salesforce__organization_id": "my_organization", + } + ''', + ), + ( + "required extras no prefix", + ''' + { + "organization_id": "my_organization", + } + ''', + ), + ] + ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_ip_filtering_auth(self, mock_salesforce): + def test_get_conn_ip_filtering_auth(self, _, extras_string, mock_salesforce): """ Testing mock IP filtering (aka allow-listing) authentication to Salesforce. Users should provide username, password, and organization ID in the Connection. Providing a client ID, Salesforce API @@ -227,21 +358,7 @@ def test_get_conn_ip_filtering_auth(self, mock_salesforce): conn_type="salesforce", login="username", password="password", - extra=''' - { - "extra__salesforce__client_id": "", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "my_organization", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "" - } - ''', + extra=extras_string, ) TestSalesforceHook._insert_conn_db_entry(ip_filtering_auth_conn.conn_id, ip_filtering_auth_conn) @@ -257,7 +374,7 @@ def test_get_conn_ip_filtering_auth(self, mock_salesforce): session_id=None, instance=None, instance_url=None, - organizationId=extras["extra__salesforce__organization_id"], + organizationId=extras.get('extra__salesforce__organization_id') or extras.get('organization_id'), version=api.DEFAULT_API_VERSION, proxies=None, session=None, From 2eec5d80cd1fd4ec41d6ef9e7652fa384864396a Mon Sep 17 00:00:00 2001 From: Alex Begg Date: Wed, 10 Nov 2021 16:00:51 -0800 Subject: [PATCH 2/5] Using cached_property to cache the Salesforce instance --- .../providers/salesforce/hooks/salesforce.py | 82 ++++++++++--------- 1 file changed, 43 insertions(+), 39 deletions(-) diff --git a/airflow/providers/salesforce/hooks/salesforce.py b/airflow/providers/salesforce/hooks/salesforce.py index 275c79a22b651..196539592f63f 100644 --- a/airflow/providers/salesforce/hooks/salesforce.py +++ b/airflow/providers/salesforce/hooks/salesforce.py @@ -27,6 +27,11 @@ import time from typing import Any, Dict, Iterable, List, Optional +try: + from functools import cached_property +except ImportError: + from cached_property import cached_property + import pandas as pd from requests import Session from simple_salesforce import Salesforce, api @@ -75,7 +80,6 @@ def __init__( ) -> None: super().__init__() self.conn_id = salesforce_conn_id - self.conn = None self.session_id = session_id self.session = session @@ -126,45 +130,45 @@ def get_ui_field_behaviour() -> Dict: }, } + @cached_property + def conn(self) -> api.Salesforce: + """Returns a Salesforce instance. (cached)""" + connection = self.get_connection(self.conn_id) + extras = connection.extra_dejson + # all extras below (besides the version one) are explicitly defaulted to None + # because simple-salesforce has a built-in authentication-choosing method that + # relies on which arguments are None and without "or None" setting this connection + # in the UI will result in the blank extras being empty strings instead of None, + # which would break the connection if "get" was used on its own. + conn = Salesforce( + username=connection.login, + password=connection.password, + security_token=extras.get('security_token') + or extras.get('extra__salesforce__security_token') + or None, + domain=extras.get('domain') or extras.get('extra__salesforce__domain') or None, + session_id=self.session_id, + instance=extras.get('instance') or extras.get('extra__salesforce__instance') or None, + instance_url=extras.get('instance_url') or extras.get('extra__salesforce__instance_url') or None, + organizationId=extras.get('organization_id') + or extras.get('extra__salesforce__organization_id') + or None, + version=extras.get('version') + or extras.get('extra__salesforce__version') + or api.DEFAULT_API_VERSION, + proxies=extras.get('proxies') or extras.get('extra__salesforce__proxies') or None, + session=self.session, + client_id=extras.get('client_id') or extras.get('extra__salesforce__client_id') or None, + consumer_key=extras.get('consumer_key') or extras.get('extra__salesforce__consumer_key') or None, + privatekey_file=extras.get('private_key_file_path') + or extras.get('extra__salesforce__private_key_file_path') + or None, + privatekey=extras.get('private_key') or extras.get('extra__salesforce__private_key') or None, + ) + return conn + def get_conn(self) -> api.Salesforce: - """Sign into Salesforce, only if we are not already signed in.""" - if not self.conn: - connection = self.get_connection(self.conn_id) - extras = connection.extra_dejson - # all extras below (besides the version one) are explicitly defaulted to None - # because simple-salesforce has a built-in authentication-choosing method that - # relies on which arguments are None and without "or None" setting this connection - # in the UI will result in the blank extras being empty strings instead of None, - # which would break the connection if "get" was used on its own. - self.conn = Salesforce( - username=connection.login, - password=connection.password, - security_token=extras.get('security_token') - or extras.get('extra__salesforce__security_token') - or None, - domain=extras.get('domain') or extras.get('extra__salesforce__domain') or None, - session_id=self.session_id, - instance=extras.get('instance') or extras.get('extra__salesforce__instance') or None, - instance_url=extras.get('instance_url') - or extras.get('extra__salesforce__instance_url') - or None, - organizationId=extras.get('organization_id') - or extras.get('extra__salesforce__organization_id') - or None, - version=extras.get('version') - or extras.get('extra__salesforce__version') - or api.DEFAULT_API_VERSION, - proxies=extras.get('proxies') or extras.get('extra__salesforce__proxies') or None, - session=self.session, - client_id=extras.get('client_id') or extras.get('extra__salesforce__client_id') or None, - consumer_key=extras.get('consumer_key') - or extras.get('extra__salesforce__consumer_key') - or None, - privatekey_file=extras.get('private_key_file_path') - or extras.get('extra__salesforce__private_key_file_path') - or None, - privatekey=extras.get('private_key') or extras.get('extra__salesforce__private_key') or None, - ) + """Returns a Salesforce instance. (cached)""" return self.conn def make_query( From 5e565b5953a7105abd725d77994a60ffeb3bd731 Mon Sep 17 00:00:00 2001 From: Alex Begg Date: Wed, 10 Nov 2021 21:34:44 -0800 Subject: [PATCH 3/5] Removing ability to allow extras without prefix for SalesforceHook --- .../providers/salesforce/hooks/salesforce.py | 30 +++----- .../salesforce/hooks/test_salesforce.py | 70 ++++--------------- 2 files changed, 25 insertions(+), 75 deletions(-) diff --git a/airflow/providers/salesforce/hooks/salesforce.py b/airflow/providers/salesforce/hooks/salesforce.py index 196539592f63f..b7f7f297c948c 100644 --- a/airflow/providers/salesforce/hooks/salesforce.py +++ b/airflow/providers/salesforce/hooks/salesforce.py @@ -143,27 +143,19 @@ def conn(self) -> api.Salesforce: conn = Salesforce( username=connection.login, password=connection.password, - security_token=extras.get('security_token') - or extras.get('extra__salesforce__security_token') - or None, - domain=extras.get('domain') or extras.get('extra__salesforce__domain') or None, + security_token=extras.get('extra__salesforce__security_token') or None, + domain=extras.get('extra__salesforce__domain') or None, session_id=self.session_id, - instance=extras.get('instance') or extras.get('extra__salesforce__instance') or None, - instance_url=extras.get('instance_url') or extras.get('extra__salesforce__instance_url') or None, - organizationId=extras.get('organization_id') - or extras.get('extra__salesforce__organization_id') - or None, - version=extras.get('version') - or extras.get('extra__salesforce__version') - or api.DEFAULT_API_VERSION, - proxies=extras.get('proxies') or extras.get('extra__salesforce__proxies') or None, + instance=extras.get('extra__salesforce__instance') or None, + instance_url=extras.get('extra__salesforce__instance_url') or None, + organizationId=extras.get('extra__salesforce__organization_id') or None, + version=extras.get('extra__salesforce__version') or api.DEFAULT_API_VERSION, + proxies=extras.get('extra__salesforce__proxies') or None, session=self.session, - client_id=extras.get('client_id') or extras.get('extra__salesforce__client_id') or None, - consumer_key=extras.get('consumer_key') or extras.get('extra__salesforce__consumer_key') or None, - privatekey_file=extras.get('private_key_file_path') - or extras.get('extra__salesforce__private_key_file_path') - or None, - privatekey=extras.get('private_key') or extras.get('extra__salesforce__private_key') or None, + client_id=extras.get('extra__salesforce__client_id') or None, + consumer_key=extras.get('extra__salesforce__consumer_key') or None, + privatekey_file=extras.get('extra__salesforce__private_key_file_path') or None, + privatekey=extras.get('extra__salesforce__private_key') or None, ) return conn diff --git a/tests/providers/salesforce/hooks/test_salesforce.py b/tests/providers/salesforce/hooks/test_salesforce.py index 03c71e76731c7..884aa6de62e73 100644 --- a/tests/providers/salesforce/hooks/test_salesforce.py +++ b/tests/providers/salesforce/hooks/test_salesforce.py @@ -80,17 +80,6 @@ def test_get_conn_exists(self): } ''', ), - ( - "required extras no prefix", - ''' - { - "client_id": "my_client", - "domain": "test", - "security_token": "token", - "version": "42.0" - } - ''', - ), ] ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") @@ -117,16 +106,16 @@ def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): mock_salesforce.assert_called_once_with( username=password_auth_conn.login, password=password_auth_conn.password, - security_token=extras.get('extra__salesforce__security_token') or extras.get('security_token'), - domain=extras.get('extra__salesforce__domain') or extras.get('domain'), + security_token=extras.get('extra__salesforce__security_token'), + domain=extras.get('extra__salesforce__domain'), session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras.get('extra__salesforce__version') or extras.get('version'), + version=extras.get('extra__salesforce__version'), proxies=None, session=None, - client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), + client_id=extras.get('extra__salesforce__client_id'), consumer_key=None, privatekey_file=None, privatekey=None, @@ -163,17 +152,6 @@ def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): } ''', ), - ( - "required extras no prefix", - ''' - { - "client_id": "my_client2", - "domain": "test", - "instance_url": "https://my.salesforce.com", - "version": "29.0" - } - ''', - ), ] ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") @@ -206,15 +184,15 @@ def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce) username=direct_access_conn.login, password=direct_access_conn.password, security_token=None, - domain=extras.get('extra__salesforce__domain') or extras.get('domain'), + domain=extras.get('extra__salesforce__domain'), session_id=self.salesforce_hook.session_id, instance=None, - instance_url=extras.get('extra__salesforce__instance_url') or extras.get('instance_url'), + instance_url=extras.get('extra__salesforce__instance_url'), organizationId=None, - version=extras.get('extra__salesforce__version') or extras.get('version'), + version=extras.get('extra__salesforce__version'), proxies=None, session=self.salesforce_hook.session, - client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), + client_id=extras.get('extra__salesforce__client_id'), consumer_key=None, privatekey_file=None, privatekey=None, @@ -252,18 +230,6 @@ def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce) } ''', ), - ( - "required extras no prefix", - ''' - { - "client_id": "my_client3", - "consumer_key": "consumer_key", - "domain": "login", - "private_key": "private_key", - "version": "34.0" - } - ''', - ), ] ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") @@ -292,18 +258,18 @@ def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): username=jwt_auth_conn.login, password=jwt_auth_conn.password, security_token=None, - domain=extras.get('extra__salesforce__domain') or extras.get('domain'), + domain=extras.get('extra__salesforce__domain'), session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras.get('extra__salesforce__version') or extras.get('version'), + version=extras.get('extra__salesforce__version'), proxies=None, session=None, - client_id=extras.get('extra__salesforce__client_id') or extras.get('client_id'), - consumer_key=extras.get('extra__salesforce__consumer_key') or extras.get('consumer_key'), + client_id=extras.get('extra__salesforce__client_id'), + consumer_key=extras.get('extra__salesforce__consumer_key'), privatekey_file=None, - privatekey=extras.get('extra__salesforce__private_key') or extras.get('private_key'), + privatekey=extras.get('extra__salesforce__private_key'), ) @parameterized.expand( @@ -334,14 +300,6 @@ def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): } ''', ), - ( - "required extras no prefix", - ''' - { - "organization_id": "my_organization", - } - ''', - ), ] ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") @@ -374,7 +332,7 @@ def test_get_conn_ip_filtering_auth(self, _, extras_string, mock_salesforce): session_id=None, instance=None, instance_url=None, - organizationId=extras.get('extra__salesforce__organization_id') or extras.get('organization_id'), + organizationId=extras.get('extra__salesforce__organization_id'), version=api.DEFAULT_API_VERSION, proxies=None, session=None, From bfc723a3dae7e980b3980c53cd4843cec49db750 Mon Sep 17 00:00:00 2001 From: Alex Begg Date: Wed, 10 Nov 2021 23:10:46 -0800 Subject: [PATCH 4/5] Replacing parameterized tests with one test for default to None --- .../providers/salesforce/hooks/salesforce.py | 6 +- .../salesforce/hooks/test_salesforce.py | 264 +++++++----------- 2 files changed, 106 insertions(+), 164 deletions(-) diff --git a/airflow/providers/salesforce/hooks/salesforce.py b/airflow/providers/salesforce/hooks/salesforce.py index b7f7f297c948c..f67b52088231c 100644 --- a/airflow/providers/salesforce/hooks/salesforce.py +++ b/airflow/providers/salesforce/hooks/salesforce.py @@ -27,15 +27,11 @@ import time from typing import Any, Dict, Iterable, List, Optional -try: - from functools import cached_property -except ImportError: - from cached_property import cached_property - import pandas as pd from requests import Session from simple_salesforce import Salesforce, api +from airflow.compat.functools import cached_property from airflow.hooks.base import BaseHook log = logging.getLogger(__name__) diff --git a/tests/providers/salesforce/hooks/test_salesforce.py b/tests/providers/salesforce/hooks/test_salesforce.py index 884aa6de62e73..d1edb0792d64d 100644 --- a/tests/providers/salesforce/hooks/test_salesforce.py +++ b/tests/providers/salesforce/hooks/test_salesforce.py @@ -23,7 +23,6 @@ import pandas as pd import pytest from numpy import nan -from parameterized import parameterized from requests import Session as request_session from simple_salesforce import Salesforce, api @@ -49,45 +48,13 @@ def test_get_conn_exists(self): assert self.salesforce_hook.conn.return_value is not None - @parameterized.expand( - [ - ( - "all extras", - ''' - { - "extra__salesforce__client_id": "my_client", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "test", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "token", - "extra__salesforce__version": "42.0" - } - ''', - ), - ( - "required extras", - ''' - { - "extra__salesforce__client_id": "my_client", - "extra__salesforce__domain": "test", - "extra__salesforce__security_token": "token", - "extra__salesforce__version": "42.0" - } - ''', - ), - ] - ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): + def test_get_conn_password_auth(self, mock_salesforce): """ Testing mock password authentication to Salesforce. Users should provide a username, password, and security token in the Connection. Providing a client ID, Salesforce API version, proxy mapping, and - domain are optional. Connection params set as empty strings should be converted to `None`. + domain are optional. Connection params not provided or set as empty strings should be converted to + `None`. """ password_auth_conn = Connection( @@ -95,7 +62,14 @@ def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): conn_type="salesforce", login=None, password=None, - extra=extras_string, + extra=''' + { + "extra__salesforce__client_id": "my_client", + "extra__salesforce__domain": "test", + "extra__salesforce__security_token": "token", + "extra__salesforce__version": "42.0" + } + ''', ) TestSalesforceHook._insert_conn_db_entry(password_auth_conn.conn_id, password_auth_conn) @@ -106,61 +80,28 @@ def test_get_conn_password_auth(self, _, extras_string, mock_salesforce): mock_salesforce.assert_called_once_with( username=password_auth_conn.login, password=password_auth_conn.password, - security_token=extras.get('extra__salesforce__security_token'), - domain=extras.get('extra__salesforce__domain'), + security_token=extras["extra__salesforce__security_token"], + domain=extras["extra__salesforce__domain"], session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras.get('extra__salesforce__version'), + version=extras["extra__salesforce__version"], proxies=None, session=None, - client_id=extras.get('extra__salesforce__client_id'), + client_id=extras["extra__salesforce__client_id"], consumer_key=None, privatekey_file=None, privatekey=None, ) - @parameterized.expand( - [ - ( - "all extras", - ''' - { - "extra__salesforce__client_id": "my_client2", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "test", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "https://my.salesforce.com", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "29.0" - } - ''', - ), - ( - "required extras", - ''' - { - "extra__salesforce__client_id": "my_client2", - "extra__salesforce__domain": "test", - "extra__salesforce__instance_url": "https://my.salesforce.com", - "extra__salesforce__version": "29.0" - } - ''', - ), - ] - ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce): + def test_get_conn_direct_session_access(self, mock_salesforce): """ Testing mock direct session access to Salesforce. Users should provide an instance (or instance URL) in the Connection and set a `session_id` value when calling `SalesforceHook`. Providing a client ID, Salesforce API version, proxy mapping, and domain are optional. Connection - params set as empty strings should be converted to `None`. + params not provided or set as empty strings should be converted to `None`. """ direct_access_conn = Connection( @@ -168,7 +109,14 @@ def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce) conn_type="salesforce", login=None, password=None, - extra=extras_string, + extra=''' + { + "extra__salesforce__client_id": "my_client2", + "extra__salesforce__domain": "test", + "extra__salesforce__instance_url": "https://my.salesforce.com", + "extra__salesforce__version": "29.0" + } + ''', ) TestSalesforceHook._insert_conn_db_entry(direct_access_conn.conn_id, direct_access_conn) @@ -184,61 +132,27 @@ def test_get_conn_direct_session_access(self, _, extras_string, mock_salesforce) username=direct_access_conn.login, password=direct_access_conn.password, security_token=None, - domain=extras.get('extra__salesforce__domain'), + domain=extras["extra__salesforce__domain"], session_id=self.salesforce_hook.session_id, instance=None, - instance_url=extras.get('extra__salesforce__instance_url'), + instance_url=extras["extra__salesforce__instance_url"], organizationId=None, - version=extras.get('extra__salesforce__version'), + version=extras["extra__salesforce__version"], proxies=None, session=self.salesforce_hook.session, - client_id=extras.get('extra__salesforce__client_id'), + client_id=extras["extra__salesforce__client_id"], consumer_key=None, privatekey_file=None, privatekey=None, ) - @parameterized.expand( - [ - ( - "all extras", - ''' - { - "extra__salesforce__client_id": "my_client3", - "extra__salesforce__consumer_key": "consumer_key", - "extra__salesforce__domain": "login", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "", - "extra__salesforce__private_key": "private_key", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "34.0" - } - ''', - ), - ( - "required extras", - ''' - { - "extra__salesforce__client_id": "my_client3", - "extra__salesforce__consumer_key": "consumer_key", - "extra__salesforce__domain": "login", - "extra__salesforce__private_key": "private_key", - "extra__salesforce__version": "34.0" - } - ''', - ), - ] - ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): + def test_get_conn_jwt_auth(self, mock_salesforce): """ Testing mock JWT bearer authentication to Salesforce. Users should provide consumer key and private key (or path to a private key) in the Connection. Providing a client ID, Salesforce API version, proxy - mapping, and domain are optional. Connection params set as empty strings should be converted to - `None`. + mapping, and domain are optional. Connection params not provided or set as empty strings should be + converted to `None`. """ jwt_auth_conn = Connection( @@ -246,7 +160,15 @@ def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): conn_type="salesforce", login=None, password=None, - extra=extras_string, + extra=''' + { + "extra__salesforce__client_id": "my_client3", + "extra__salesforce__consumer_key": "consumer_key", + "extra__salesforce__domain": "login", + "extra__salesforce__private_key": "private_key", + "extra__salesforce__version": "34.0" + } + ''', ) TestSalesforceHook._insert_conn_db_entry(jwt_auth_conn.conn_id, jwt_auth_conn) @@ -258,57 +180,27 @@ def test_get_conn_jwt_auth(self, _, extras_string, mock_salesforce): username=jwt_auth_conn.login, password=jwt_auth_conn.password, security_token=None, - domain=extras.get('extra__salesforce__domain'), + domain=extras["extra__salesforce__domain"], session_id=None, instance=None, instance_url=None, organizationId=None, - version=extras.get('extra__salesforce__version'), + version=extras["extra__salesforce__version"], proxies=None, session=None, - client_id=extras.get('extra__salesforce__client_id'), - consumer_key=extras.get('extra__salesforce__consumer_key'), + client_id=extras["extra__salesforce__client_id"], + consumer_key=extras["extra__salesforce__consumer_key"], privatekey_file=None, - privatekey=extras.get('extra__salesforce__private_key'), + privatekey=extras["extra__salesforce__private_key"], ) - @parameterized.expand( - [ - ( - "all extras", - ''' - { - "extra__salesforce__client_id": "", - "extra__salesforce__consumer_key": "", - "extra__salesforce__domain": "", - "extra__salesforce__instance": "", - "extra__salesforce__instance_url": "", - "extra__salesforce__organization_id": "my_organization", - "extra__salesforce__private_key": "", - "extra__salesforce__private_key_file_path": "", - "extra__salesforce__proxies": "", - "extra__salesforce__security_token": "", - "extra__salesforce__version": "" - } - ''', - ), - ( - "required extras", - ''' - { - "extra__salesforce__organization_id": "my_organization", - } - ''', - ), - ] - ) @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") - def test_get_conn_ip_filtering_auth(self, _, extras_string, mock_salesforce): + def test_get_conn_ip_filtering_auth(self, mock_salesforce): """ Testing mock IP filtering (aka allow-listing) authentication to Salesforce. Users should provide username, password, and organization ID in the Connection. Providing a client ID, Salesforce API - version, proxy mapping, and domain are optional. Connection params set as empty strings should be - converted to `None`. + version, proxy mapping, and domain are optional. Connection params not provided or set as empty + strings should be converted to `None`. """ ip_filtering_auth_conn = Connection( @@ -316,7 +208,11 @@ def test_get_conn_ip_filtering_auth(self, _, extras_string, mock_salesforce): conn_type="salesforce", login="username", password="password", - extra=extras_string, + extra=''' + { + "extra__salesforce__organization_id": "my_organization" + } + ''', ) TestSalesforceHook._insert_conn_db_entry(ip_filtering_auth_conn.conn_id, ip_filtering_auth_conn) @@ -332,7 +228,57 @@ def test_get_conn_ip_filtering_auth(self, _, extras_string, mock_salesforce): session_id=None, instance=None, instance_url=None, - organizationId=extras.get('extra__salesforce__organization_id'), + organizationId=extras["extra__salesforce__organization_id"], + version=api.DEFAULT_API_VERSION, + proxies=None, + session=None, + client_id=None, + consumer_key=None, + privatekey_file=None, + privatekey=None, + ) + + @patch("airflow.providers.salesforce.hooks.salesforce.Salesforce") + def test_get_conn_default_to_none(self, mock_salesforce): + """ + Testing mock authentication to Salesforce so that every extra connection param set as an empty + string will be converted to `None`. + """ + + default_to_none_conn = Connection( + conn_id="default_to_none_conn", + conn_type="salesforce", + login=None, + password=None, + extra=''' + { + "extra__salesforce__client_id": "", + "extra__salesforce__consumer_key": "", + "extra__salesforce__domain": "", + "extra__salesforce__instance": "", + "extra__salesforce__instance_url": "", + "extra__salesforce__organization_id": "", + "extra__salesforce__private_key": "", + "extra__salesforce__private_key_file_path": "", + "extra__salesforce__proxies": "", + "extra__salesforce__security_token": "" + } + ''', + ) + TestSalesforceHook._insert_conn_db_entry(default_to_none_conn.conn_id, default_to_none_conn) + + self.salesforce_hook = SalesforceHook(salesforce_conn_id="default_to_none_conn") + self.salesforce_hook.get_conn() + + mock_salesforce.assert_called_once_with( + username=default_to_none_conn.login, + password=default_to_none_conn.password, + security_token=None, + domain=None, + session_id=None, + instance=None, + instance_url=None, + organizationId=None, version=api.DEFAULT_API_VERSION, proxies=None, session=None, From 96c2227888739bc124441069d2b9bb44e41aa92e Mon Sep 17 00:00:00 2001 From: Alex Begg Date: Thu, 11 Nov 2021 11:42:52 -0800 Subject: [PATCH 5/5] Not including airflow.compat.* in a provider --- airflow/providers/salesforce/hooks/salesforce.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/airflow/providers/salesforce/hooks/salesforce.py b/airflow/providers/salesforce/hooks/salesforce.py index f67b52088231c..b7f7f297c948c 100644 --- a/airflow/providers/salesforce/hooks/salesforce.py +++ b/airflow/providers/salesforce/hooks/salesforce.py @@ -27,11 +27,15 @@ import time from typing import Any, Dict, Iterable, List, Optional +try: + from functools import cached_property +except ImportError: + from cached_property import cached_property + import pandas as pd from requests import Session from simple_salesforce import Salesforce, api -from airflow.compat.functools import cached_property from airflow.hooks.base import BaseHook log = logging.getLogger(__name__)