From ec38134f4b18c9aea369c9187050215a22054124 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Wed, 21 Sep 2022 16:30:53 -0400 Subject: [PATCH 1/8] Add oracledb thick mode support for oracle provider --- airflow/providers/oracle/hooks/oracle.py | 24 +++++++++++++++++++ airflow/utils/db.py | 12 ++++++++++ .../connections/oracle.rst | 9 +++++++ tests/providers/oracle/hooks/test_oracle.py | 16 +++++++++++++ 4 files changed, 61 insertions(+) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index a66d12dc4d33c..0cc2b1e42b40d 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -28,6 +28,7 @@ except ImportError: numpy = None # type: ignore +from airflow.exceptions import AirflowException from airflow.providers.common.sql.hooks.sql import DbApiHook PARAM_TYPES = {bool, float, int, str} @@ -57,6 +58,29 @@ class OracleHook(DbApiHook): supports_autocommit = True + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + if self.oracle_conn_id is not None: + conn = self.get_connection(self.oracle_conn_id) + if conn.extra is not None: + extra_options = conn.extra_dejson + + # Check if thick_mode should be enabled in python-oracledb + thick_mode = extra_options.get('thick_mode', False) + if not isinstance(thick_mode, bool): + raise AirflowException(f'thick_mode should be a boolean but type was {type(thick_mode)}') + if thick_mode: + thick_mode_lib_dir = extra_options.get('thick_mode_lib_dir') + thick_mode_config_dir = extra_options.get('thick_mode_config_dir') + oracledb.init_oracle_client(lib_dir=thick_mode_lib_dir, config_dir=thick_mode_config_dir) + + # Check if python-oracledb Defaults attributes should be set + # Values default to the initial values used by python-oracledb + fetch_decimals = extra_options.get('fetch_decimals', False) + oracledb.defaults.fetch_decimals = fetch_decimals + fetch_lobs = extra_options.get('fetch_lobs', True) + oracledb.defaults.fetch_lobs = fetch_lobs + def get_conn(self) -> oracledb.Connection: """ Returns a oracle connection object diff --git a/airflow/utils/db.py b/airflow/utils/db.py index 6aa284026402f..5c24502f46250 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -421,6 +421,18 @@ def create_default_connections(session: Session = NEW_SESSION): ), session, ) + merge_conn( + Connection( + conn_id="oracle_default", + conn_type="oracle", + host="localhost", + login="root", + password="password", + schema="schema", + port=1521, + ), + session, + ) merge_conn( Connection( conn_id="oss_default", diff --git a/docs/apache-airflow-providers-oracle/connections/oracle.rst b/docs/apache-airflow-providers-oracle/connections/oracle.rst index 3934b6648e3ac..c2a55885b6d9e 100644 --- a/docs/apache-airflow-providers-oracle/connections/oracle.rst +++ b/docs/apache-airflow-providers-oracle/connections/oracle.rst @@ -58,6 +58,15 @@ Extra (optional) configuration parameter. * ``dsn``. Specify a Data Source Name (and ignore Host). * ``sid`` or ``service_name``. Use to form DSN instead of Schema. + * ``thick_mode`` (bool) - Specify whether to use python-oracledb in thick mode. Defaults to False. + If set to True, you must have the Oracle Client libraries installed. + See `oracledb docs` for more info. + * ``thick_mode_lib_dir`` (str) - Path to use to find the Oracle Client libraries when using thick mode. + If not specified, defaults to the standard way of locating the Oracle Client library on the OS. + See `oracledb docs` for more info. + * ``thick_mode_config_dir`` (str) - Path to use to find the Oracle Client library configuration files when using thick mode. + If not specified, defaults to the standard way of locating the Oracle Client library configuration files on the OS. + See `oracledb docs` for more info. Connect using `dsn`, Host and `sid`, Host and `service_name`, or only Host `(OracleHook.getconn Documentation) `_. diff --git a/tests/providers/oracle/hooks/test_oracle.py b/tests/providers/oracle/hooks/test_oracle.py index 124381450907e..76f52bf66a1d6 100644 --- a/tests/providers/oracle/hooks/test_oracle.py +++ b/tests/providers/oracle/hooks/test_oracle.py @@ -144,6 +144,21 @@ def test_set_current_schema(self, mock_connect): self.connection.extra = json.dumps({'service_name': 'service_name'}) assert self.db_hook.get_conn().current_schema == self.connection.schema + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') + def test_set_thick_mode(self, mock_init_client): + thick_mode_test = { + 'thick_mode': True, + 'thick_mode_lib_dir': '/opt/oracle/instantclient', + 'thick_mode_config_dir': '/opt/oracle/config', + } + self.connection.extra = json.dumps(thick_mode_test) + self.db_hook.__init__() + assert mock_init_client.call_count == 1 + args, kwargs = mock_init_client.call_args + assert args == () + assert kwargs['lib_dir'] == thick_mode_test['thick_mode_lib_dir'] + assert kwargs['config_dir'] == thick_mode_test['thick_mode_config_dir'] + @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHook(unittest.TestCase): @@ -157,6 +172,7 @@ def setUp(self): class UnitTestOracleHook(OracleHook): conn_name_attr = 'test_conn_id' + oracle_conn_id = None def get_conn(self): return conn From 0fae53aa261183b1f74618381bd28dff5259c0d7 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Wed, 21 Sep 2022 17:02:45 -0400 Subject: [PATCH 2/8] Update docs for fetch_decimals and fetch_lobs defaults --- docs/apache-airflow-providers-oracle/connections/oracle.rst | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/docs/apache-airflow-providers-oracle/connections/oracle.rst b/docs/apache-airflow-providers-oracle/connections/oracle.rst index c2a55885b6d9e..3c5899e61440e 100644 --- a/docs/apache-airflow-providers-oracle/connections/oracle.rst +++ b/docs/apache-airflow-providers-oracle/connections/oracle.rst @@ -67,6 +67,10 @@ Extra (optional) * ``thick_mode_config_dir`` (str) - Path to use to find the Oracle Client library configuration files when using thick mode. If not specified, defaults to the standard way of locating the Oracle Client library configuration files on the OS. See `oracledb docs` for more info. + * ``fetch_decimals`` (bool) - Specify whether numbers should be fetched as ``decimal.Decimal`` values. Defaults to False. + See `defaults.fetch_decimals` for more info. + * ``fetch_lobs`` (bool) - Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. Defaults to True. + See `defaults.fetch_lobs` for more info. Connect using `dsn`, Host and `sid`, Host and `service_name`, or only Host `(OracleHook.getconn Documentation) `_. From 23e59748d6fe7c720966f9dd9351d73d372fa6ff Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Thu, 22 Sep 2022 01:26:21 -0400 Subject: [PATCH 3/8] Consolidate setting parameters in init --- airflow/providers/oracle/hooks/oracle.py | 252 +++++++++++------- .../connections/oracle.rst | 12 +- tests/providers/oracle/hooks/test_oracle.py | 31 ++- 3 files changed, 192 insertions(+), 103 deletions(-) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index 0cc2b1e42b40d..0f08018332565 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -20,6 +20,7 @@ import math import warnings from datetime import datetime +from typing import Any import oracledb @@ -49,42 +50,8 @@ class OracleHook(DbApiHook): :param oracle_conn_id: The :ref:`Oracle connection id ` used for Oracle credentials. - """ - - conn_name_attr = 'oracle_conn_id' - default_conn_name = 'oracle_default' - conn_type = 'oracle' - hook_name = 'Oracle' - - supports_autocommit = True - - def __init__(self, *args, **kwargs) -> None: - super().__init__(*args, **kwargs) - if self.oracle_conn_id is not None: - conn = self.get_connection(self.oracle_conn_id) - if conn.extra is not None: - extra_options = conn.extra_dejson - # Check if thick_mode should be enabled in python-oracledb - thick_mode = extra_options.get('thick_mode', False) - if not isinstance(thick_mode, bool): - raise AirflowException(f'thick_mode should be a boolean but type was {type(thick_mode)}') - if thick_mode: - thick_mode_lib_dir = extra_options.get('thick_mode_lib_dir') - thick_mode_config_dir = extra_options.get('thick_mode_config_dir') - oracledb.init_oracle_client(lib_dir=thick_mode_lib_dir, config_dir=thick_mode_config_dir) - - # Check if python-oracledb Defaults attributes should be set - # Values default to the initial values used by python-oracledb - fetch_decimals = extra_options.get('fetch_decimals', False) - oracledb.defaults.fetch_decimals = fetch_decimals - fetch_lobs = extra_options.get('fetch_lobs', True) - oracledb.defaults.fetch_lobs = fetch_lobs - - def get_conn(self) -> oracledb.Connection: - """ - Returns a oracle connection object - Optional parameters for using a custom DSN connection + Optional parameters for using a custom DSN connection (instead of using a server alias from tnsnames.ora) The dsn (data source name) is the TNS entry (from the Oracle names server or tnsnames.ora file) @@ -105,76 +72,179 @@ def get_conn(self) -> oracledb.Connection: see more param detail in `oracledb.connect `_ + """ + + conn_name_attr = 'oracle_conn_id' + default_conn_name = 'oracle_default' + conn_type = 'oracle' + hook_name = 'Oracle' + supports_autocommit = True - """ - conn = self.get_connection(self.oracle_conn_id) # type: ignore[attr-defined] - conn_config = {'user': conn.login, 'password': conn.password} - sid = conn.extra_dejson.get('sid') - mod = conn.extra_dejson.get('module') - schema = conn.schema - - service_name = conn.extra_dejson.get('service_name') - port = conn.port if conn.port else 1521 - if conn.host and sid and not service_name: - conn_config['dsn'] = oracledb.makedsn(conn.host, port, sid) - elif conn.host and service_name and not sid: - conn_config['dsn'] = oracledb.makedsn(conn.host, port, service_name=service_name) + def __init__( + self, + oracle_conn_id: str | None = 'oracle_default', + host: str | None = None, + schema: str | None = None, + login: str | None = None, + password: str | None = None, + sid: str | None = None, + mod: str | None = None, + service_name: str | None = None, + port: int = 1521, + dsn: str | None = None, + events: bool = False, + mode: int = oracledb.AUTH_MODE_DEFAULT, + purity: int = oracledb.PURITY_DEFAULT, + threaded: bool = True, + thick_mode: bool = False, + thick_mode_lib_dir: str | None = None, + thick_mode_config_dir: str | None = None, + fetch_decimals: bool = False, + fetch_lobs: bool = True, + *args, + **kwargs, + ) -> None: + super().__init__(*args, **kwargs) + + self.oracle_conn_id = oracle_conn_id + self.host = host + self.schema = schema + self.login = login + self.password = password + self.sid = sid + self.mod = mod + self.service_name = service_name + self.port = port + self.dsn = dsn + self.events = events + self.mode = mode + self.purity = purity + self.threaded = threaded + self.thick_mode = thick_mode + self.thick_mode_lib_dir = thick_mode_lib_dir + self.thick_mode_config_dir = thick_mode_config_dir + self.fetch_decimals = fetch_decimals + self.fetch_lobs = fetch_lobs + + self.conn_config: dict[str, Any] = {} + + # Use connection to override defaults + if self.oracle_conn_id is not None: + conn = self.get_connection(self.oracle_conn_id) + self.host = conn.host if conn.host else self.host + self.login = conn.login if conn.login else self.login + self.password = conn.password if conn.password else self.password + self.port = conn.port if conn.port else self.port + self.schema = conn.schema if conn.schema else self.schema + + if conn.extra is not None: + extra_options = conn.extra_dejson + + sid = extra_options.get('sid') + self.sid = sid if sid else self.sid + + mod = extra_options.get('mod') + self.mod = mod if mod else self.mod + + service_name = extra_options.get('service_name') + self.service_name = service_name if service_name else self.service_name + + dsn = extra_options.get('dsn') + self.dsn = dsn if dsn else self.dsn + + events = extra_options.get('events') + self.events = events if events is not None else self.events + + threaded = extra_options.get('threaded') + self.threaded = threaded if threaded is not None else self.threaded + + mode = extra_options.get('mode', '').lower() + if mode == 'sysdba': + self.mode = oracledb.AUTH_MODE_SYSDBA + elif mode == 'sysasm': + self.mode = oracledb.AUTH_MODE_SYSASM + elif mode == 'sysoper': + self.mode = oracledb.AUTH_MODE_SYSOPER + elif mode == 'sysbkp': + self.mode = oracledb.AUTH_MODE_SYSBKP + elif mode == 'sysdgd': + self.mode = oracledb.AUTH_MODE_SYSDGD + elif mode == 'syskmt': + self.mode = oracledb.AUTH_MODE_SYSKMT + elif mode == 'sysrac': + self.mode = oracledb.AUTH_MODE_SYSRAC + + purity = extra_options.get('purity', '').lower() + if purity == 'new': + self.purity = oracledb.PURITY_NEW + elif purity == 'self': + self.purity = oracledb.PURITY_SELF + elif purity == 'default': + self.purity = oracledb.PURITY_DEFAULT + + # Check if thick_mode should be enabled in python-oracledb + thick_mode = extra_options.get('thick_mode') + self.thick_mode = thick_mode if thick_mode is not None else self.thick_mode + if not isinstance(self.thick_mode, bool): + raise AirflowException( + f'thick_mode should be a boolean but type was {type(self.thick_mode)}' + ) + if self.thick_mode: + thick_mode_lib_dir = extra_options.get('thick_mode_lib_dir') + thick_mode_config_dir = extra_options.get('thick_mode_config_dir') + oracledb.init_oracle_client(lib_dir=thick_mode_lib_dir, config_dir=thick_mode_config_dir) + + # Check if python-oracledb Defaults attributes should be set + # Values default to the initial values used by python-oracledb + fetch_decimals = extra_options.get('fetch_decimals') + self.fetch_decimals = fetch_decimals if fetch_decimals is not None else self.fetch_decimals + oracledb.defaults.fetch_decimals = self.fetch_decimals + fetch_lobs = extra_options.get('fetch_lobs') + self.fetch_lobs = fetch_lobs if fetch_lobs is not None else self.fetch_lobs + oracledb.defaults.fetch_lobs = self.fetch_lobs + + self.conn_config['user'] = self.login + self.conn_config['password'] = self.password + self.conn_config['events'] = self.events + self.conn_config['mode'] = self.mode + self.conn_config['purity'] = self.purity + self.conn_config['threaded'] = self.threaded + + if self.host and self.sid and not self.service_name: + self.conn_config['dsn'] = oracledb.makedsn(self.host, self.port, self.sid) + elif self.host and self.service_name and not self.sid: + self.conn_config['dsn'] = oracledb.makedsn(self.host, self.port, service_name=self.service_name) else: - dsn = conn.extra_dejson.get('dsn') - if dsn is None: - dsn = conn.host - if conn.port is not None: - dsn += ":" + str(conn.port) - if service_name: - dsn += "/" + service_name - elif conn.schema: + if self.dsn is None: + dsn = str(self.host) if self.host is not None else '' + if self.port is not None: + dsn += ":" + str(self.port) + if self.service_name: + dsn += "/" + str(service_name) + elif self.schema: warnings.warn( """Using conn.schema to pass the Oracle Service Name is deprecated. Please use conn.extra.service_name instead.""", DeprecationWarning, stacklevel=2, ) - dsn += "/" + conn.schema - conn_config['dsn'] = dsn - - if 'events' in conn.extra_dejson: - conn_config['events'] = conn.extra_dejson.get('events') - - mode = conn.extra_dejson.get('mode', '').lower() - if mode == 'sysdba': - conn_config['mode'] = oracledb.AUTH_MODE_SYSDBA - elif mode == 'sysasm': - conn_config['mode'] = oracledb.AUTH_MODE_SYSASM - elif mode == 'sysoper': - conn_config['mode'] = oracledb.AUTH_MODE_SYSOPER - elif mode == 'sysbkp': - conn_config['mode'] = oracledb.AUTH_MODE_SYSBKP - elif mode == 'sysdgd': - conn_config['mode'] = oracledb.AUTH_MODE_SYSDGD - elif mode == 'syskmt': - conn_config['mode'] = oracledb.AUTH_MODE_SYSKMT - elif mode == 'sysrac': - conn_config['mode'] = oracledb.AUTH_MODE_SYSRAC - - purity = conn.extra_dejson.get('purity', '').lower() - if purity == 'new': - conn_config['purity'] = oracledb.PURITY_NEW - elif purity == 'self': - conn_config['purity'] = oracledb.PURITY_SELF - elif purity == 'default': - conn_config['purity'] = oracledb.PURITY_DEFAULT - - conn = oracledb.connect(**conn_config) - if mod is not None: - conn.module = mod + dsn += "/" + self.schema + self.dsn = dsn if dsn else self.dsn + self.conn_config['dsn'] = self.dsn + + def get_conn(self) -> oracledb.Connection: + """Returns a oracle connection object""" + conn = oracledb.connect(**self.conn_config) + if self.mod is not None: + conn.module = self.mod # if Connection.schema is defined, set schema after connecting successfully # cannot be part of conn_config # https://python-oracledb.readthedocs.io/en/latest/api_manual/connection.html?highlight=schema#Connection.current_schema # Only set schema when not using conn.schema as Service Name - if schema and service_name: - conn.current_schema = schema + if self.schema and self.service_name: + conn.current_schema = self.schema return conn diff --git a/docs/apache-airflow-providers-oracle/connections/oracle.rst b/docs/apache-airflow-providers-oracle/connections/oracle.rst index 3c5899e61440e..3a7f1f37f0fd1 100644 --- a/docs/apache-airflow-providers-oracle/connections/oracle.rst +++ b/docs/apache-airflow-providers-oracle/connections/oracle.rst @@ -42,13 +42,6 @@ Extra (optional) Specify the extra parameters (as json dictionary) that can be used in Oracle connection. The following parameters are supported: - * ``encoding`` - The encoding to use for regular database strings. If not specified, - the environment variable ``NLS_LANG`` is used. If the environment variable ``NLS_LANG`` - is not set, ``ASCII`` is used. - * ``nencoding`` - The encoding to use for national character set database strings. - If not specified, the environment variable ``NLS_NCHAR`` is used. If the environment - variable ``NLS_NCHAR`` is not used, the environment variable ``NLS_LANG`` is used instead, - and if the environment variable ``NLS_LANG`` is not set, ``ASCII`` is used. * ``threaded`` - Whether or not Oracle should wrap accesses to connections with a mutex. Default value is False. * ``events`` - Whether or not to initialize Oracle in events mode. @@ -58,6 +51,8 @@ Extra (optional) configuration parameter. * ``dsn``. Specify a Data Source Name (and ignore Host). * ``sid`` or ``service_name``. Use to form DSN instead of Schema. + * ``module`` (str) - This write-only attribute sets the module column in the v$session table. + The maximum length for this string is 48 and if you exceed this length you will get ORA-24960. * ``thick_mode`` (bool) - Specify whether to use python-oracledb in thick mode. Defaults to False. If set to True, you must have the Oracle Client libraries installed. See `oracledb docs` for more info. @@ -72,6 +67,7 @@ Extra (optional) * ``fetch_lobs`` (bool) - Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. Defaults to True. See `defaults.fetch_lobs` for more info. + Connect using `dsn`, Host and `sid`, Host and `service_name`, or only Host `(OracleHook.getconn Documentation) `_. For example: @@ -106,8 +102,6 @@ Extra (optional) .. code-block:: json { - "encoding": "UTF-8", - "nencoding": "UTF-8", "threaded": false, "events": false, "mode": "sysdba", diff --git a/tests/providers/oracle/hooks/test_oracle.py b/tests/providers/oracle/hooks/test_oracle.py index 76f52bf66a1d6..c94bb6920df2f 100644 --- a/tests/providers/oracle/hooks/test_oracle.py +++ b/tests/providers/oracle/hooks/test_oracle.py @@ -27,6 +27,8 @@ from airflow.models import Connection from airflow.providers.oracle.hooks.oracle import OracleHook +from airflow.utils import db +from airflow.utils.session import create_session try: import oracledb @@ -36,14 +38,30 @@ @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHookConn(unittest.TestCase): + CONN_ORACLE_WITH_NO_EXTRA = 'oracle_with_no_extra' + + @classmethod + def tearDownClass(self) -> None: + with create_session() as session: + conns_to_reset = [self.CONN_ORACLE_WITH_NO_EXTRA] + connections = session.query(Connection).filter(Connection.conn_id.in_(conns_to_reset)) + connections.delete(synchronize_session=False) + session.commit() + def setUp(self): super().setUp() - self.connection = Connection( - login='login', password='password', host='host', schema='schema', port=1521 + conn_id=self.CONN_ORACLE_WITH_NO_EXTRA, + conn_type='oracle', + login='login', + password='password', + host='host', + schema='schema', + port=1521, ) + db.merge_conn(self.connection) - self.db_hook = OracleHook() + self.db_hook = OracleHook(oracle_conn_id=self.CONN_ORACLE_WITH_NO_EXTRA) self.db_hook.get_connection = mock.Mock() self.db_hook.get_connection.return_value = self.connection @@ -60,6 +78,7 @@ def test_get_conn_host(self, mock_connect): @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_get_conn_host_alternative_port(self, mock_connect): self.connection.port = 1522 + self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -72,6 +91,7 @@ def test_get_conn_host_alternative_port(self, mock_connect): def test_get_conn_sid(self, mock_connect): dsn_sid = {'dsn': 'ignored', 'sid': 'sid'} self.connection.extra = json.dumps(dsn_sid) + self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -82,6 +102,7 @@ def test_get_conn_sid(self, mock_connect): def test_get_conn_service_name(self, mock_connect): dsn_service_name = {'dsn': 'ignored', 'service_name': 'service_name'} self.connection.extra = json.dumps(dsn_service_name) + self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -103,6 +124,7 @@ def test_get_conn_mode(self, mock_connect): first = True for mod in mode: self.connection.extra = json.dumps({'mode': mod}) + self.db_hook.__init__() self.db_hook.get_conn() if first: assert mock_connect.call_count == 1 @@ -114,6 +136,7 @@ def test_get_conn_mode(self, mock_connect): @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_get_conn_events(self, mock_connect): self.connection.extra = json.dumps({'events': True}) + self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -130,6 +153,7 @@ def test_get_conn_purity(self, mock_connect): first = True for pur in purity: self.connection.extra = json.dumps({'purity': pur}) + self.db_hook.__init__() self.db_hook.get_conn() if first: assert mock_connect.call_count == 1 @@ -142,6 +166,7 @@ def test_get_conn_purity(self, mock_connect): def test_set_current_schema(self, mock_connect): self.connection.schema = "schema_name" self.connection.extra = json.dumps({'service_name': 'service_name'}) + self.db_hook.__init__() assert self.db_hook.get_conn().current_schema == self.connection.schema @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') From cc0e018ad201f1109a84d2d9750bda45573e49a6 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Thu, 22 Sep 2022 03:08:07 -0400 Subject: [PATCH 4/8] Remove support for threaded since it is ignored by python-oracledb --- airflow/providers/oracle/hooks/oracle.py | 6 ------ docs/apache-airflow-providers-oracle/connections/oracle.rst | 3 --- 2 files changed, 9 deletions(-) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index 0f08018332565..dc409fda6f7f6 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -96,7 +96,6 @@ def __init__( events: bool = False, mode: int = oracledb.AUTH_MODE_DEFAULT, purity: int = oracledb.PURITY_DEFAULT, - threaded: bool = True, thick_mode: bool = False, thick_mode_lib_dir: str | None = None, thick_mode_config_dir: str | None = None, @@ -120,7 +119,6 @@ def __init__( self.events = events self.mode = mode self.purity = purity - self.threaded = threaded self.thick_mode = thick_mode self.thick_mode_lib_dir = thick_mode_lib_dir self.thick_mode_config_dir = thick_mode_config_dir @@ -156,9 +154,6 @@ def __init__( events = extra_options.get('events') self.events = events if events is not None else self.events - threaded = extra_options.get('threaded') - self.threaded = threaded if threaded is not None else self.threaded - mode = extra_options.get('mode', '').lower() if mode == 'sysdba': self.mode = oracledb.AUTH_MODE_SYSDBA @@ -209,7 +204,6 @@ def __init__( self.conn_config['events'] = self.events self.conn_config['mode'] = self.mode self.conn_config['purity'] = self.purity - self.conn_config['threaded'] = self.threaded if self.host and self.sid and not self.service_name: self.conn_config['dsn'] = oracledb.makedsn(self.host, self.port, self.sid) diff --git a/docs/apache-airflow-providers-oracle/connections/oracle.rst b/docs/apache-airflow-providers-oracle/connections/oracle.rst index 3a7f1f37f0fd1..f0d750e672364 100644 --- a/docs/apache-airflow-providers-oracle/connections/oracle.rst +++ b/docs/apache-airflow-providers-oracle/connections/oracle.rst @@ -42,8 +42,6 @@ Extra (optional) Specify the extra parameters (as json dictionary) that can be used in Oracle connection. The following parameters are supported: - * ``threaded`` - Whether or not Oracle should wrap accesses to connections with a mutex. - Default value is False. * ``events`` - Whether or not to initialize Oracle in events mode. * ``mode`` - one of ``sysdba``, ``sysasm``, ``sysoper``, ``sysbkp``, ``sysdgd``, ``syskmt`` or ``sysrac`` which are defined at the module level, Default mode is connecting. @@ -102,7 +100,6 @@ Extra (optional) .. code-block:: json { - "threaded": false, "events": false, "mode": "sysdba", "purity": "new" From 16f19c6fbb4b545a735ef5d4b219289df83b2aca Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Thu, 22 Sep 2022 23:14:01 -0400 Subject: [PATCH 5/8] Add hook params but process connection config extra in get_conn --- airflow/providers/oracle/hooks/oracle.py | 295 +++++++++----------- tests/providers/oracle/hooks/test_oracle.py | 119 ++++++-- 2 files changed, 227 insertions(+), 187 deletions(-) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index dc409fda6f7f6..09a29f7120772 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -20,7 +20,6 @@ import math import warnings from datetime import datetime -from typing import Any import oracledb @@ -29,7 +28,6 @@ except ImportError: numpy = None # type: ignore -from airflow.exceptions import AirflowException from airflow.providers.common.sql.hooks.sql import DbApiHook PARAM_TYPES = {bool, float, int, str} @@ -50,28 +48,32 @@ class OracleHook(DbApiHook): :param oracle_conn_id: The :ref:`Oracle connection id ` used for Oracle credentials. - - Optional parameters for using a custom DSN connection - (instead of using a server alias from tnsnames.ora) - The dsn (data source name) is the TNS entry - (from the Oracle names server or tnsnames.ora file) - or is a string like the one returned from makedsn(). - - :param dsn: the data source name for the Oracle server - :param service_name: the db_unique_name of the database - that you are connecting to (CONNECT_DATA part of TNS) - :param sid: Oracle System ID that identifies a particular - database on a system - - You can set these parameters in the extra fields of your connection - as in - - .. code-block:: python - - {"dsn": ("(DESCRIPTION=(ADDRESS=(PROTOCOL=TCP)(HOST=host)(PORT=1521))(CONNECT_DATA=(SID=sid)))")} - - see more param detail in `oracledb.connect - `_ + :param thick_mode: Specify whether to use python-oracledb in thick mode. Defaults to False. + If set to True, you must have the Oracle Client libraries installed. + See `oracledb docs` + for more info. + :param thick_mode_lib_dir: Path to use to find the Oracle Client libraries when using thick mode. + If not specified, defaults to the standard way of locating the Oracle Client library on the OS. + See `oracledb docs + ` + for more info. + :param thick_mode_config_dir: Path to use to find the Oracle Client library + configuration files when using thick mode. + If not specified, defaults to the standard way of locating the Oracle Client + library configuration files on the OS. + See `oracledb docs + ` + for more info. + :param fetch_decimals: Specify whether numbers should be fetched as ``decimal.Decimal`` values. + Defaults to False. + See `defaults.fetch_decimals + ` + for more info. + :param fetch_lobs: Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. + Defaults to True. + See `defaults.fetch_lobs + ` + for more info. """ conn_name_attr = 'oracle_conn_id' @@ -83,162 +85,143 @@ class OracleHook(DbApiHook): def __init__( self, - oracle_conn_id: str | None = 'oracle_default', - host: str | None = None, - schema: str | None = None, - login: str | None = None, - password: str | None = None, - sid: str | None = None, - mod: str | None = None, - service_name: str | None = None, - port: int = 1521, - dsn: str | None = None, - events: bool = False, - mode: int = oracledb.AUTH_MODE_DEFAULT, - purity: int = oracledb.PURITY_DEFAULT, - thick_mode: bool = False, + *args, + thick_mode: bool | None = None, thick_mode_lib_dir: str | None = None, thick_mode_config_dir: str | None = None, - fetch_decimals: bool = False, - fetch_lobs: bool = True, - *args, + fetch_decimals: bool | None = None, + fetch_lobs: bool | None = None, **kwargs, ) -> None: super().__init__(*args, **kwargs) - self.oracle_conn_id = oracle_conn_id - self.host = host - self.schema = schema - self.login = login - self.password = password - self.sid = sid - self.mod = mod - self.service_name = service_name - self.port = port - self.dsn = dsn - self.events = events - self.mode = mode - self.purity = purity self.thick_mode = thick_mode self.thick_mode_lib_dir = thick_mode_lib_dir self.thick_mode_config_dir = thick_mode_config_dir self.fetch_decimals = fetch_decimals self.fetch_lobs = fetch_lobs - self.conn_config: dict[str, Any] = {} - - # Use connection to override defaults - if self.oracle_conn_id is not None: - conn = self.get_connection(self.oracle_conn_id) - self.host = conn.host if conn.host else self.host - self.login = conn.login if conn.login else self.login - self.password = conn.password if conn.password else self.password - self.port = conn.port if conn.port else self.port - self.schema = conn.schema if conn.schema else self.schema - - if conn.extra is not None: - extra_options = conn.extra_dejson - - sid = extra_options.get('sid') - self.sid = sid if sid else self.sid - - mod = extra_options.get('mod') - self.mod = mod if mod else self.mod - - service_name = extra_options.get('service_name') - self.service_name = service_name if service_name else self.service_name - - dsn = extra_options.get('dsn') - self.dsn = dsn if dsn else self.dsn - - events = extra_options.get('events') - self.events = events if events is not None else self.events - - mode = extra_options.get('mode', '').lower() - if mode == 'sysdba': - self.mode = oracledb.AUTH_MODE_SYSDBA - elif mode == 'sysasm': - self.mode = oracledb.AUTH_MODE_SYSASM - elif mode == 'sysoper': - self.mode = oracledb.AUTH_MODE_SYSOPER - elif mode == 'sysbkp': - self.mode = oracledb.AUTH_MODE_SYSBKP - elif mode == 'sysdgd': - self.mode = oracledb.AUTH_MODE_SYSDGD - elif mode == 'syskmt': - self.mode = oracledb.AUTH_MODE_SYSKMT - elif mode == 'sysrac': - self.mode = oracledb.AUTH_MODE_SYSRAC - - purity = extra_options.get('purity', '').lower() - if purity == 'new': - self.purity = oracledb.PURITY_NEW - elif purity == 'self': - self.purity = oracledb.PURITY_SELF - elif purity == 'default': - self.purity = oracledb.PURITY_DEFAULT - - # Check if thick_mode should be enabled in python-oracledb - thick_mode = extra_options.get('thick_mode') - self.thick_mode = thick_mode if thick_mode is not None else self.thick_mode - if not isinstance(self.thick_mode, bool): - raise AirflowException( - f'thick_mode should be a boolean but type was {type(self.thick_mode)}' - ) - if self.thick_mode: - thick_mode_lib_dir = extra_options.get('thick_mode_lib_dir') - thick_mode_config_dir = extra_options.get('thick_mode_config_dir') - oracledb.init_oracle_client(lib_dir=thick_mode_lib_dir, config_dir=thick_mode_config_dir) - - # Check if python-oracledb Defaults attributes should be set - # Values default to the initial values used by python-oracledb - fetch_decimals = extra_options.get('fetch_decimals') - self.fetch_decimals = fetch_decimals if fetch_decimals is not None else self.fetch_decimals - oracledb.defaults.fetch_decimals = self.fetch_decimals - fetch_lobs = extra_options.get('fetch_lobs') - self.fetch_lobs = fetch_lobs if fetch_lobs is not None else self.fetch_lobs - oracledb.defaults.fetch_lobs = self.fetch_lobs - - self.conn_config['user'] = self.login - self.conn_config['password'] = self.password - self.conn_config['events'] = self.events - self.conn_config['mode'] = self.mode - self.conn_config['purity'] = self.purity - - if self.host and self.sid and not self.service_name: - self.conn_config['dsn'] = oracledb.makedsn(self.host, self.port, self.sid) - elif self.host and self.service_name and not self.sid: - self.conn_config['dsn'] = oracledb.makedsn(self.host, self.port, service_name=self.service_name) + def get_conn(self) -> oracledb.Connection: + """ + Returns a oracle connection object + Optional parameters for using a custom DSN connection + (instead of using a server alias from tnsnames.ora) + The dsn (data source name) is the TNS entry + (from the Oracle names server or tnsnames.ora file) + or is a string like the one returned from makedsn(). + + :param dsn: the data source name for the Oracle server + :param service_name: the db_unique_name of the database + that you are connecting to (CONNECT_DATA part of TNS) + :param sid: Oracle System ID that identifies a particular + database on a system + + You can set these parameters in the extra fields of your connection + as in + + .. code-block:: python + + {"dsn": ("(DESCRIPTION=(ADDRESS=(PROTOCOL=TCP)(HOST=host)(PORT=1521))(CONNECT_DATA=(SID=sid)))")} + + see more param detail in `oracledb.connect + `_ + + + """ + conn = self.get_connection(self.oracle_conn_id) # type: ignore[attr-defined] + conn_config = {'user': conn.login, 'password': conn.password} + sid = conn.extra_dejson.get('sid') + mod = conn.extra_dejson.get('module') + schema = conn.schema + + # Enable oracledb thick mode if thick_mode is set to True, defaults to False + # Parameters take precedence over connection config extra + # Defaults to False (use thin mode) if not provided in params or connection config extra + if self.thick_mode is None: + self.thick_mode = conn.extra_dejson.get('thick_mode', False) + if self.thick_mode: + if self.thick_mode_lib_dir is None: + self.thick_mode_lib_dir = conn.extra_dejson.get('thick_mode_lib_dir') + if self.thick_mode_config_dir is None: + self.thick_mode_config_dir = conn.extra_dejson.get('thick_mode_config_dir') + oracledb.init_oracle_client( + lib_dir=self.thick_mode_lib_dir, config_dir=self.thick_mode_config_dir + ) + + # Set oracledb Defaults Attributes + # Default to the initial values + # if not provided in params or connection config extra + # (https://python-oracledb.readthedocs.io/en/latest/api_manual/defaults.html) + if self.fetch_decimals is None: + self.fetch_decimals = conn.extra_dejson.get('fetch_decimals', False) + oracledb.defaults.fetch_decimals = self.fetch_decimals + + if self.fetch_lobs is None: + self.fetch_lobs = conn.extra_dejson.get('fetch_lobs', True) + oracledb.defaults.fetch_lobs = self.fetch_lobs + + # Set up DSN + service_name = conn.extra_dejson.get('service_name') + port = conn.port if conn.port else 1521 + if conn.host and sid and not service_name: + conn_config['dsn'] = oracledb.makedsn(conn.host, port, sid) + elif conn.host and service_name and not sid: + conn_config['dsn'] = oracledb.makedsn(conn.host, port, service_name=service_name) else: - if self.dsn is None: - dsn = str(self.host) if self.host is not None else '' - if self.port is not None: - dsn += ":" + str(self.port) - if self.service_name: - dsn += "/" + str(service_name) - elif self.schema: + dsn = conn.extra_dejson.get('dsn') + if dsn is None: + dsn = conn.host + if conn.port is not None: + dsn += ":" + str(conn.port) + if service_name: + dsn += "/" + service_name + elif conn.schema: warnings.warn( """Using conn.schema to pass the Oracle Service Name is deprecated. Please use conn.extra.service_name instead.""", DeprecationWarning, stacklevel=2, ) - dsn += "/" + self.schema - self.dsn = dsn if dsn else self.dsn - self.conn_config['dsn'] = self.dsn - - def get_conn(self) -> oracledb.Connection: - """Returns a oracle connection object""" - conn = oracledb.connect(**self.conn_config) - if self.mod is not None: - conn.module = self.mod + dsn += "/" + conn.schema + conn_config['dsn'] = dsn + + if 'events' in conn.extra_dejson: + conn_config['events'] = conn.extra_dejson.get('events') + + mode = conn.extra_dejson.get('mode', '').lower() + if mode == 'sysdba': + conn_config['mode'] = oracledb.AUTH_MODE_SYSDBA + elif mode == 'sysasm': + conn_config['mode'] = oracledb.AUTH_MODE_SYSASM + elif mode == 'sysoper': + conn_config['mode'] = oracledb.AUTH_MODE_SYSOPER + elif mode == 'sysbkp': + conn_config['mode'] = oracledb.AUTH_MODE_SYSBKP + elif mode == 'sysdgd': + conn_config['mode'] = oracledb.AUTH_MODE_SYSDGD + elif mode == 'syskmt': + conn_config['mode'] = oracledb.AUTH_MODE_SYSKMT + elif mode == 'sysrac': + conn_config['mode'] = oracledb.AUTH_MODE_SYSRAC + + purity = conn.extra_dejson.get('purity', '').lower() + if purity == 'new': + conn_config['purity'] = oracledb.PURITY_NEW + elif purity == 'self': + conn_config['purity'] = oracledb.PURITY_SELF + elif purity == 'default': + conn_config['purity'] = oracledb.PURITY_DEFAULT + + conn = oracledb.connect(**conn_config) + if mod is not None: + conn.module = mod # if Connection.schema is defined, set schema after connecting successfully # cannot be part of conn_config # https://python-oracledb.readthedocs.io/en/latest/api_manual/connection.html?highlight=schema#Connection.current_schema # Only set schema when not using conn.schema as Service Name - if self.schema and self.service_name: - conn.current_schema = self.schema + if schema and service_name: + conn.current_schema = schema return conn diff --git a/tests/providers/oracle/hooks/test_oracle.py b/tests/providers/oracle/hooks/test_oracle.py index c94bb6920df2f..14a98b23c3484 100644 --- a/tests/providers/oracle/hooks/test_oracle.py +++ b/tests/providers/oracle/hooks/test_oracle.py @@ -27,8 +27,6 @@ from airflow.models import Connection from airflow.providers.oracle.hooks.oracle import OracleHook -from airflow.utils import db -from airflow.utils.session import create_session try: import oracledb @@ -38,30 +36,14 @@ @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHookConn(unittest.TestCase): - CONN_ORACLE_WITH_NO_EXTRA = 'oracle_with_no_extra' - - @classmethod - def tearDownClass(self) -> None: - with create_session() as session: - conns_to_reset = [self.CONN_ORACLE_WITH_NO_EXTRA] - connections = session.query(Connection).filter(Connection.conn_id.in_(conns_to_reset)) - connections.delete(synchronize_session=False) - session.commit() - def setUp(self): super().setUp() + self.connection = Connection( - conn_id=self.CONN_ORACLE_WITH_NO_EXTRA, - conn_type='oracle', - login='login', - password='password', - host='host', - schema='schema', - port=1521, + login='login', password='password', host='host', schema='schema', port=1521 ) - db.merge_conn(self.connection) - self.db_hook = OracleHook(oracle_conn_id=self.CONN_ORACLE_WITH_NO_EXTRA) + self.db_hook = OracleHook() self.db_hook.get_connection = mock.Mock() self.db_hook.get_connection.return_value = self.connection @@ -78,7 +60,6 @@ def test_get_conn_host(self, mock_connect): @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_get_conn_host_alternative_port(self, mock_connect): self.connection.port = 1522 - self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -91,7 +72,6 @@ def test_get_conn_host_alternative_port(self, mock_connect): def test_get_conn_sid(self, mock_connect): dsn_sid = {'dsn': 'ignored', 'sid': 'sid'} self.connection.extra = json.dumps(dsn_sid) - self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -102,7 +82,6 @@ def test_get_conn_sid(self, mock_connect): def test_get_conn_service_name(self, mock_connect): dsn_service_name = {'dsn': 'ignored', 'service_name': 'service_name'} self.connection.extra = json.dumps(dsn_service_name) - self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -124,7 +103,6 @@ def test_get_conn_mode(self, mock_connect): first = True for mod in mode: self.connection.extra = json.dumps({'mode': mod}) - self.db_hook.__init__() self.db_hook.get_conn() if first: assert mock_connect.call_count == 1 @@ -136,7 +114,6 @@ def test_get_conn_mode(self, mock_connect): @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_get_conn_events(self, mock_connect): self.connection.extra = json.dumps({'events': True}) - self.db_hook.__init__() self.db_hook.get_conn() assert mock_connect.call_count == 1 args, kwargs = mock_connect.call_args @@ -153,7 +130,6 @@ def test_get_conn_purity(self, mock_connect): first = True for pur in purity: self.connection.extra = json.dumps({'purity': pur}) - self.db_hook.__init__() self.db_hook.get_conn() if first: assert mock_connect.call_count == 1 @@ -166,24 +142,106 @@ def test_get_conn_purity(self, mock_connect): def test_set_current_schema(self, mock_connect): self.connection.schema = "schema_name" self.connection.extra = json.dumps({'service_name': 'service_name'}) - self.db_hook.__init__() assert self.db_hook.get_conn().current_schema == self.connection.schema @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') - def test_set_thick_mode(self, mock_init_client): + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_thick_mode_extra(self, mock_connect, mock_init_client): thick_mode_test = { 'thick_mode': True, 'thick_mode_lib_dir': '/opt/oracle/instantclient', 'thick_mode_config_dir': '/opt/oracle/config', } self.connection.extra = json.dumps(thick_mode_test) - self.db_hook.__init__() + self.db_hook.get_conn() + assert mock_connect.call_count == 1 assert mock_init_client.call_count == 1 args, kwargs = mock_init_client.call_args assert args == () assert kwargs['lib_dir'] == thick_mode_test['thick_mode_lib_dir'] assert kwargs['config_dir'] == thick_mode_test['thick_mode_config_dir'] + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_thick_mode_params(self, mock_connect, mock_init_client): + # Verify params overrides connection config extra + thick_mode_test = { + 'thick_mode': False, + 'thick_mode_lib_dir': '/opt/oracle/instantclient', + 'thick_mode_config_dir': '/opt/oracle/config', + } + self.connection.extra = json.dumps(thick_mode_test) + db_hook = OracleHook(thick_mode=True, thick_mode_lib_dir='/test', thick_mode_config_dir='/test_conf') + db_hook.get_connection = mock.Mock() + db_hook.get_connection.return_value = self.connection + db_hook.get_conn() + assert mock_connect.call_count == 1 + assert mock_init_client.call_count == 1 + args, kwargs = mock_init_client.call_args + assert args == () + assert kwargs['lib_dir'] == '/test' + assert kwargs['config_dir'] == '/test_conf' + + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_thick_mode_defaults_to_false(self, mock_connect, mock_init_client): + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + assert mock_init_client.call_count == 0 + + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_thick_mode_dirs_defaults(self, mock_connect, mock_init_client): + thick_mode_test = {'thick_mode': True} + self.connection.extra = json.dumps(thick_mode_test) + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + assert mock_init_client.call_count == 1 + args, kwargs = mock_init_client.call_args + assert args == () + assert kwargs['lib_dir'] is None + assert kwargs['config_dir'] is None + + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_oracledb_defaults_attributes_default_values(self, mock_connect): + # Check that oracledb defaults are what we expect + assert oracledb.defaults.fetch_decimals is False + assert oracledb.defaults.fetch_lobs is True + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + # Check that OracleHook.get_conn() properly defaults values + assert self.db_hook.fetch_decimals is False + assert self.db_hook.fetch_lobs is True + # Check that oracledb defaults are still correct + assert oracledb.defaults.fetch_decimals is False + assert oracledb.defaults.fetch_lobs is True + + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_oracledb_defaults_attributes_extra(self, mock_connect): + defaults_test = {'fetch_decimals': True, 'fetch_lobs': False} + self.connection.extra = json.dumps(defaults_test) + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + assert self.db_hook.fetch_decimals == defaults_test['fetch_decimals'] + assert self.db_hook.fetch_lobs == defaults_test['fetch_lobs'] + assert oracledb.defaults.fetch_decimals == defaults_test['fetch_decimals'] + assert oracledb.defaults.fetch_lobs == defaults_test['fetch_lobs'] + + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_oracledb_defaults_attributes_params(self, mock_connect): + # Verify params overrides connection config extra + defaults_test = {'fetch_decimals': False, 'fetch_lobs': True} + self.connection.extra = json.dumps(defaults_test) + db_hook = OracleHook(fetch_decimals=True, fetch_lobs=False) + db_hook.get_connection = mock.Mock() + db_hook.get_connection.return_value = self.connection + db_hook.get_conn() + assert mock_connect.call_count == 1 + assert db_hook.fetch_decimals is True + assert db_hook.fetch_lobs is False + assert oracledb.defaults.fetch_decimals is True + assert oracledb.defaults.fetch_lobs is False + @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHook(unittest.TestCase): @@ -197,7 +255,6 @@ def setUp(self): class UnitTestOracleHook(OracleHook): conn_name_attr = 'test_conn_id' - oracle_conn_id = None def get_conn(self): return conn From f53792fe3c3a805d4ebc7d70e446daee8600a400 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Fri, 23 Sep 2022 00:23:11 -0400 Subject: [PATCH 6/8] Add type checking and tests --- airflow/providers/oracle/hooks/oracle.py | 16 +++++++++++ tests/providers/oracle/hooks/test_oracle.py | 30 +++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index 09a29f7120772..20b52ab9f0715 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -139,11 +139,23 @@ def get_conn(self) -> oracledb.Connection: # Defaults to False (use thin mode) if not provided in params or connection config extra if self.thick_mode is None: self.thick_mode = conn.extra_dejson.get('thick_mode', False) + if not isinstance(self.thick_mode, bool): + raise TypeError(f'thick_mode expected bool, got {type(self.thick_mode).__name__}') if self.thick_mode: if self.thick_mode_lib_dir is None: self.thick_mode_lib_dir = conn.extra_dejson.get('thick_mode_lib_dir') + if not isinstance(self.thick_mode_lib_dir, (str, type(None))): + raise TypeError( + f'thick_mode_lib_dir expected str or None, ' + f'got {type(self.thick_mode_lib_dir).__name__}' + ) if self.thick_mode_config_dir is None: self.thick_mode_config_dir = conn.extra_dejson.get('thick_mode_config_dir') + if not isinstance(self.thick_mode_config_dir, (str, type(None))): + raise TypeError( + f'thick_mode_config_dir expected str or None, ' + f'got {type(self.thick_mode_config_dir).__name__}' + ) oracledb.init_oracle_client( lib_dir=self.thick_mode_lib_dir, config_dir=self.thick_mode_config_dir ) @@ -154,10 +166,14 @@ def get_conn(self) -> oracledb.Connection: # (https://python-oracledb.readthedocs.io/en/latest/api_manual/defaults.html) if self.fetch_decimals is None: self.fetch_decimals = conn.extra_dejson.get('fetch_decimals', False) + if not isinstance(self.fetch_decimals, bool): + raise TypeError(f'fetch_decimals expected bool, got {type(self.fetch_decimals).__name__}') oracledb.defaults.fetch_decimals = self.fetch_decimals if self.fetch_lobs is None: self.fetch_lobs = conn.extra_dejson.get('fetch_lobs', True) + if not isinstance(self.fetch_lobs, bool): + raise TypeError(f'fetch_lobs expected bool, got {type(self.fetch_lobs).__name__}') oracledb.defaults.fetch_lobs = self.fetch_lobs # Set up DSN diff --git a/tests/providers/oracle/hooks/test_oracle.py b/tests/providers/oracle/hooks/test_oracle.py index 14a98b23c3484..1889fd3a798b5 100644 --- a/tests/providers/oracle/hooks/test_oracle.py +++ b/tests/providers/oracle/hooks/test_oracle.py @@ -242,6 +242,36 @@ def test_set_oracledb_defaults_attributes_params(self, mock_connect): assert oracledb.defaults.fetch_decimals is True assert oracledb.defaults.fetch_lobs is False + def test_type_checking_thick_mode(self): + with pytest.raises(TypeError, match=r"thick_mode expected bool, got.*"): + thick_mode_test = {'thick_mode': 'bad'} + self.connection.extra = json.dumps(thick_mode_test) + self.db_hook.get_conn() + + def test_type_checking_thick_mode_lib_dir(self): + with pytest.raises(TypeError, match=r"thick_mode_lib_dir expected str or None, got.*"): + thick_mode_lib_dir_test = {'thick_mode': True, 'thick_mode_lib_dir': 1} + self.connection.extra = json.dumps(thick_mode_lib_dir_test) + self.db_hook.get_conn() + + def test_type_checking_thick_mode_config_dir(self): + with pytest.raises(TypeError, match=r"thick_mode_config_dir expected str or None, got.*"): + thick_mode_config_dir_test = {'thick_mode': True, 'thick_mode_config_dir': 1} + self.connection.extra = json.dumps(thick_mode_config_dir_test) + self.db_hook.get_conn() + + def test_type_checking_fetch_decimals(self): + with pytest.raises(TypeError, match=r"fetch_decimals expected bool, got.*"): + fetch_decimals_test = {'fetch_decimals': 'bad'} + self.connection.extra = json.dumps(fetch_decimals_test) + self.db_hook.get_conn() + + def test_type_checking_fetch_lobs(self): + with pytest.raises(TypeError, match=r"fetch_lobs expected bool, got.*"): + fetch_lobs_test = {'fetch_lobs': 'bad'} + self.connection.extra = json.dumps(fetch_lobs_test) + self.db_hook.get_conn() + @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHook(unittest.TestCase): From 28eefd5aaee1e568c38f575f7f075147b9018f12 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Fri, 23 Sep 2022 03:08:30 -0400 Subject: [PATCH 7/8] Revise parsing of conn config extra booleans and only set oracledb defaults if provided --- airflow/providers/oracle/hooks/oracle.py | 55 ++++++++++++--------- tests/providers/oracle/hooks/test_oracle.py | 54 +++++++++----------- 2 files changed, 55 insertions(+), 54 deletions(-) diff --git a/airflow/providers/oracle/hooks/oracle.py b/airflow/providers/oracle/hooks/oracle.py index 20b52ab9f0715..a3850f9dd181f 100644 --- a/airflow/providers/oracle/hooks/oracle.py +++ b/airflow/providers/oracle/hooks/oracle.py @@ -42,6 +42,26 @@ def _map_param(value): return value +def _get_bool(val): + if isinstance(val, bool): + return val + if isinstance(val, str): + val = val.lower().strip() + if val == 'true': + return True + if val == 'false': + return False + return None + + +def _get_first_bool(*vals): + for val in vals: + converted = _get_bool(val) + if isinstance(converted, bool): + return converted + return None + + class OracleHook(DbApiHook): """ Interact with Oracle SQL. @@ -65,12 +85,10 @@ class OracleHook(DbApiHook): ` for more info. :param fetch_decimals: Specify whether numbers should be fetched as ``decimal.Decimal`` values. - Defaults to False. See `defaults.fetch_decimals ` for more info. :param fetch_lobs: Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. - Defaults to True. See `defaults.fetch_lobs ` for more info. @@ -134,14 +152,11 @@ def get_conn(self) -> oracledb.Connection: mod = conn.extra_dejson.get('module') schema = conn.schema - # Enable oracledb thick mode if thick_mode is set to True, defaults to False + # Enable oracledb thick mode if thick_mode is set to True # Parameters take precedence over connection config extra - # Defaults to False (use thin mode) if not provided in params or connection config extra - if self.thick_mode is None: - self.thick_mode = conn.extra_dejson.get('thick_mode', False) - if not isinstance(self.thick_mode, bool): - raise TypeError(f'thick_mode expected bool, got {type(self.thick_mode).__name__}') - if self.thick_mode: + # Defaults to use thin mode if not provided in params or connection config extra + thick_mode = _get_first_bool(self.thick_mode, conn.extra_dejson.get('thick_mode')) + if thick_mode is True: if self.thick_mode_lib_dir is None: self.thick_mode_lib_dir = conn.extra_dejson.get('thick_mode_lib_dir') if not isinstance(self.thick_mode_lib_dir, (str, type(None))): @@ -160,21 +175,15 @@ def get_conn(self) -> oracledb.Connection: lib_dir=self.thick_mode_lib_dir, config_dir=self.thick_mode_config_dir ) - # Set oracledb Defaults Attributes - # Default to the initial values - # if not provided in params or connection config extra + # Set oracledb Defaults Attributes if provided # (https://python-oracledb.readthedocs.io/en/latest/api_manual/defaults.html) - if self.fetch_decimals is None: - self.fetch_decimals = conn.extra_dejson.get('fetch_decimals', False) - if not isinstance(self.fetch_decimals, bool): - raise TypeError(f'fetch_decimals expected bool, got {type(self.fetch_decimals).__name__}') - oracledb.defaults.fetch_decimals = self.fetch_decimals - - if self.fetch_lobs is None: - self.fetch_lobs = conn.extra_dejson.get('fetch_lobs', True) - if not isinstance(self.fetch_lobs, bool): - raise TypeError(f'fetch_lobs expected bool, got {type(self.fetch_lobs).__name__}') - oracledb.defaults.fetch_lobs = self.fetch_lobs + fetch_decimals = _get_first_bool(self.fetch_decimals, conn.extra_dejson.get('fetch_decimals')) + if isinstance(fetch_decimals, bool): + oracledb.defaults.fetch_decimals = fetch_decimals + + fetch_lobs = _get_first_bool(self.fetch_lobs, conn.extra_dejson.get('fetch_lobs')) + if isinstance(fetch_lobs, bool): + oracledb.defaults.fetch_lobs = fetch_lobs # Set up DSN service_name = conn.extra_dejson.get('service_name') diff --git a/tests/providers/oracle/hooks/test_oracle.py b/tests/providers/oracle/hooks/test_oracle.py index 1889fd3a798b5..0e6b6ee96f61b 100644 --- a/tests/providers/oracle/hooks/test_oracle.py +++ b/tests/providers/oracle/hooks/test_oracle.py @@ -161,6 +161,15 @@ def test_set_thick_mode_extra(self, mock_connect, mock_init_client): assert kwargs['lib_dir'] == thick_mode_test['thick_mode_lib_dir'] assert kwargs['config_dir'] == thick_mode_test['thick_mode_config_dir'] + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_thick_mode_extra_str(self, mock_connect, mock_init_client): + thick_mode_test = {'thick_mode': 'True'} + self.connection.extra = json.dumps(thick_mode_test) + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + assert mock_init_client.call_count == 1 + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.init_oracle_client') @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_set_thick_mode_params(self, mock_connect, mock_init_client): @@ -204,17 +213,13 @@ def test_thick_mode_dirs_defaults(self, mock_connect, mock_init_client): @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_oracledb_defaults_attributes_default_values(self, mock_connect): - # Check that oracledb defaults are what we expect - assert oracledb.defaults.fetch_decimals is False - assert oracledb.defaults.fetch_lobs is True + default_fetch_decimals = oracledb.defaults.fetch_decimals + default_fetch_lobs = oracledb.defaults.fetch_lobs self.db_hook.get_conn() assert mock_connect.call_count == 1 - # Check that OracleHook.get_conn() properly defaults values - assert self.db_hook.fetch_decimals is False - assert self.db_hook.fetch_lobs is True - # Check that oracledb defaults are still correct - assert oracledb.defaults.fetch_decimals is False - assert oracledb.defaults.fetch_lobs is True + # Check that OracleHook.get_conn() doesn't try to set defaults if not provided + assert oracledb.defaults.fetch_decimals == default_fetch_decimals + assert oracledb.defaults.fetch_lobs == default_fetch_lobs @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_set_oracledb_defaults_attributes_extra(self, mock_connect): @@ -222,11 +227,18 @@ def test_set_oracledb_defaults_attributes_extra(self, mock_connect): self.connection.extra = json.dumps(defaults_test) self.db_hook.get_conn() assert mock_connect.call_count == 1 - assert self.db_hook.fetch_decimals == defaults_test['fetch_decimals'] - assert self.db_hook.fetch_lobs == defaults_test['fetch_lobs'] assert oracledb.defaults.fetch_decimals == defaults_test['fetch_decimals'] assert oracledb.defaults.fetch_lobs == defaults_test['fetch_lobs'] + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') + def test_set_oracledb_defaults_attributes_extra_str(self, mock_connect): + defaults_test = {'fetch_decimals': 'True', 'fetch_lobs': 'False'} + self.connection.extra = json.dumps(defaults_test) + self.db_hook.get_conn() + assert mock_connect.call_count == 1 + assert oracledb.defaults.fetch_decimals is True + assert oracledb.defaults.fetch_lobs is False + @mock.patch('airflow.providers.oracle.hooks.oracle.oracledb.connect') def test_set_oracledb_defaults_attributes_params(self, mock_connect): # Verify params overrides connection config extra @@ -237,17 +249,9 @@ def test_set_oracledb_defaults_attributes_params(self, mock_connect): db_hook.get_connection.return_value = self.connection db_hook.get_conn() assert mock_connect.call_count == 1 - assert db_hook.fetch_decimals is True - assert db_hook.fetch_lobs is False assert oracledb.defaults.fetch_decimals is True assert oracledb.defaults.fetch_lobs is False - def test_type_checking_thick_mode(self): - with pytest.raises(TypeError, match=r"thick_mode expected bool, got.*"): - thick_mode_test = {'thick_mode': 'bad'} - self.connection.extra = json.dumps(thick_mode_test) - self.db_hook.get_conn() - def test_type_checking_thick_mode_lib_dir(self): with pytest.raises(TypeError, match=r"thick_mode_lib_dir expected str or None, got.*"): thick_mode_lib_dir_test = {'thick_mode': True, 'thick_mode_lib_dir': 1} @@ -260,18 +264,6 @@ def test_type_checking_thick_mode_config_dir(self): self.connection.extra = json.dumps(thick_mode_config_dir_test) self.db_hook.get_conn() - def test_type_checking_fetch_decimals(self): - with pytest.raises(TypeError, match=r"fetch_decimals expected bool, got.*"): - fetch_decimals_test = {'fetch_decimals': 'bad'} - self.connection.extra = json.dumps(fetch_decimals_test) - self.db_hook.get_conn() - - def test_type_checking_fetch_lobs(self): - with pytest.raises(TypeError, match=r"fetch_lobs expected bool, got.*"): - fetch_lobs_test = {'fetch_lobs': 'bad'} - self.connection.extra = json.dumps(fetch_lobs_test) - self.db_hook.get_conn() - @unittest.skipIf(oracledb is None, 'oracledb package not present') class TestOracleHook(unittest.TestCase): From df99481ba531620443fb5011afd2bf40cd98c0b6 Mon Sep 17 00:00:00 2001 From: Paul Williams Date: Tue, 27 Sep 2022 01:57:08 +0000 Subject: [PATCH 8/8] Remove references to defaults for fetch_decimals and fetch_lobs --- docs/apache-airflow-providers-oracle/connections/oracle.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/apache-airflow-providers-oracle/connections/oracle.rst b/docs/apache-airflow-providers-oracle/connections/oracle.rst index f0d750e672364..6575dce464e63 100644 --- a/docs/apache-airflow-providers-oracle/connections/oracle.rst +++ b/docs/apache-airflow-providers-oracle/connections/oracle.rst @@ -60,9 +60,9 @@ Extra (optional) * ``thick_mode_config_dir`` (str) - Path to use to find the Oracle Client library configuration files when using thick mode. If not specified, defaults to the standard way of locating the Oracle Client library configuration files on the OS. See `oracledb docs` for more info. - * ``fetch_decimals`` (bool) - Specify whether numbers should be fetched as ``decimal.Decimal`` values. Defaults to False. + * ``fetch_decimals`` (bool) - Specify whether numbers should be fetched as ``decimal.Decimal`` values. See `defaults.fetch_decimals` for more info. - * ``fetch_lobs`` (bool) - Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. Defaults to True. + * ``fetch_lobs`` (bool) - Specify whether to fetch strings/bytes for CLOBs or BLOBs instead of locators. See `defaults.fetch_lobs` for more info.