From 6751daa7a548cf4796c5c286ee095a9e5edea970 Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 6 Dec 2024 08:55:09 +0100 Subject: [PATCH 1/7] refactor: Made get_conn in JdbcHook threadsafe to avoid OSError: JVM is already started when used in multithreaded environment --- .../src/airflow/providers/jdbc/hooks/jdbc.py | 2 + providers/tests/jdbc/hooks/test_jdbc.py | 49 ++++++++++++++++++- 2 files changed, 49 insertions(+), 2 deletions(-) diff --git a/providers/src/airflow/providers/jdbc/hooks/jdbc.py b/providers/src/airflow/providers/jdbc/hooks/jdbc.py index 47fcbe8e039c9..51c14b701216b 100644 --- a/providers/src/airflow/providers/jdbc/hooks/jdbc.py +++ b/providers/src/airflow/providers/jdbc/hooks/jdbc.py @@ -28,6 +28,7 @@ from airflow.exceptions import AirflowException from airflow.providers.common.sql.hooks.sql import DbApiHook +from wrapt import synchronized if TYPE_CHECKING: from airflow.models.connection import Connection @@ -177,6 +178,7 @@ def get_sqlalchemy_engine(self, engine_kwargs=None): return super().get_sqlalchemy_engine(engine_kwargs) + @synchronized def get_conn(self) -> jaydebeapi.Connection: conn: Connection = self.connection host: str = conn.host diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index cfb27934d86da..78a368f6f0317 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -20,12 +20,14 @@ import json import logging import sqlite3 +from concurrent.futures import ThreadPoolExecutor, as_completed +from threading import current_thread +from time import sleep from unittest import mock -from unittest.mock import Mock, patch +from unittest.mock import Mock, patch, MagicMock import jaydebeapi import pytest - from airflow.exceptions import AirflowException from airflow.models import Connection from airflow.providers.jdbc.hooks.jdbc import JdbcHook, suppress_and_warn @@ -54,10 +56,17 @@ def get_hook( **conn_params, } ) + jvm_started = False class MockedJdbcHook(JdbcHook): @classmethod def get_connection(cls, conn_id: str) -> Connection: + # nonlocal jvm_started + # + # if jvm_started: + # raise OSError("JVM already started") + # + # jvm_started = True return connection hook = MockedJdbcHook(**hook_params) @@ -229,3 +238,39 @@ def test_get_sqlalchemy_engine_verify_creator_is_being_used(self): jdbc_hook.get_conn = lambda: connection engine = jdbc_hook.get_sqlalchemy_engine() assert engine.connect().connection.connection == connection + + def test_get_conn_thread_safety(self): + mock_conn = MagicMock() + open_connections = 0 + + def connect_side_effect(*args, **kwargs): + nonlocal open_connections + open_connections += 1 + logging.debug("Thread %s has %s open connections", current_thread().name, open_connections) + + try: + if open_connections > 1: + raise OSError("JVM is already started") + finally: + sleep(0.1) # wait a bit before releasing the connection again + open_connections -= 1 + + return mock_conn + + with patch.object(jaydebeapi, "connect", side_effect=connect_side_effect) as mock_connect: + jdbc_hook = get_hook() + + def call_get_conn(): + conn = jdbc_hook.get_conn() + assert conn is mock_conn + + with ThreadPoolExecutor(max_workers=10) as executor: + futures = [] + + for _ in range(0, 10): + futures.append(executor.submit(call_get_conn)) + + for future in as_completed(futures): + future.result() # This will raise OSError if get_conn isn't threadsafe + + assert mock_connect.call_count == 10 From d18abf9754cfb0d645f841aa0b95586d7727b872 Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 6 Dec 2024 22:21:25 +0100 Subject: [PATCH 2/7] Refactor: removed commented code --- providers/tests/jdbc/hooks/test_jdbc.py | 6 ------ 1 file changed, 6 deletions(-) diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index 78a368f6f0317..1e5b16812b3af 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -61,12 +61,6 @@ def get_hook( class MockedJdbcHook(JdbcHook): @classmethod def get_connection(cls, conn_id: str) -> Connection: - # nonlocal jvm_started - # - # if jvm_started: - # raise OSError("JVM already started") - # - # jvm_started = True return connection hook = MockedJdbcHook(**hook_params) From 9f99dd0106cdef172326c82a134ad79d4a8d9870 Mon Sep 17 00:00:00 2001 From: David Blain Date: Mon, 9 Dec 2024 09:43:55 +0100 Subject: [PATCH 3/7] refactor: Reorganized imports --- providers/src/airflow/providers/jdbc/hooks/jdbc.py | 2 +- providers/tests/jdbc/hooks/test_jdbc.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/providers/src/airflow/providers/jdbc/hooks/jdbc.py b/providers/src/airflow/providers/jdbc/hooks/jdbc.py index 51c14b701216b..283cb46e87980 100644 --- a/providers/src/airflow/providers/jdbc/hooks/jdbc.py +++ b/providers/src/airflow/providers/jdbc/hooks/jdbc.py @@ -25,10 +25,10 @@ import jaydebeapi import jpype from sqlalchemy.engine import URL +from wrapt import synchronized from airflow.exceptions import AirflowException from airflow.providers.common.sql.hooks.sql import DbApiHook -from wrapt import synchronized if TYPE_CHECKING: from airflow.models.connection import Connection diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index 1e5b16812b3af..3190a192162aa 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -24,7 +24,7 @@ from threading import current_thread from time import sleep from unittest import mock -from unittest.mock import Mock, patch, MagicMock +from unittest.mock import MagicMock, Mock, patch import jaydebeapi import pytest From e81d8417c576a50ae51fe5b9815b3cce726912a6 Mon Sep 17 00:00:00 2001 From: David Blain Date: Mon, 9 Dec 2024 14:13:03 +0100 Subject: [PATCH 4/7] refactor: Added white line --- providers/tests/jdbc/hooks/test_jdbc.py | 1 + 1 file changed, 1 insertion(+) diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index 3190a192162aa..451df5d7a104d 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -28,6 +28,7 @@ import jaydebeapi import pytest + from airflow.exceptions import AirflowException from airflow.models import Connection from airflow.providers.jdbc.hooks.jdbc import JdbcHook, suppress_and_warn From ca2acfd7099b3120a4e0f5b768b9877ae61ce03d Mon Sep 17 00:00:00 2001 From: David Blain Date: Mon, 9 Dec 2024 15:01:34 +0100 Subject: [PATCH 5/7] refactor: Fixed static checks test JdbcHook --- providers/tests/jdbc/hooks/test_jdbc.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index 451df5d7a104d..06e13b6b932cc 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -28,7 +28,6 @@ import jaydebeapi import pytest - from airflow.exceptions import AirflowException from airflow.models import Connection from airflow.providers.jdbc.hooks.jdbc import JdbcHook, suppress_and_warn @@ -38,6 +37,7 @@ jdbc_conn_mock = Mock(name="jdbc_conn") +logger = logging.getLogger(__name__) def get_hook( @@ -57,7 +57,6 @@ def get_hook( **conn_params, } ) - jvm_started = False class MockedJdbcHook(JdbcHook): @classmethod @@ -241,7 +240,7 @@ def test_get_conn_thread_safety(self): def connect_side_effect(*args, **kwargs): nonlocal open_connections open_connections += 1 - logging.debug("Thread %s has %s open connections", current_thread().name, open_connections) + logger.debug("Thread %s has %s open connections", current_thread().name, open_connections) try: if open_connections > 1: From 02d8cf6f16d9b2cb57b6da064ae3297d84c7a47d Mon Sep 17 00:00:00 2001 From: David Blain Date: Mon, 9 Dec 2024 16:23:50 +0100 Subject: [PATCH 6/7] refactor: Added white line --- providers/tests/jdbc/hooks/test_jdbc.py | 1 + 1 file changed, 1 insertion(+) diff --git a/providers/tests/jdbc/hooks/test_jdbc.py b/providers/tests/jdbc/hooks/test_jdbc.py index 06e13b6b932cc..73015b5b522ab 100644 --- a/providers/tests/jdbc/hooks/test_jdbc.py +++ b/providers/tests/jdbc/hooks/test_jdbc.py @@ -28,6 +28,7 @@ import jaydebeapi import pytest + from airflow.exceptions import AirflowException from airflow.models import Connection from airflow.providers.jdbc.hooks.jdbc import JdbcHook, suppress_and_warn From 553155f73d8a0468bc44894c191246df85fe7b7a Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 13 Dec 2024 12:06:46 +0100 Subject: [PATCH 7/7] refactor: Refactored JdbcHook get_conn method using RLock as suggested by Jarek instead of wrapt synchronized decorator --- .../src/airflow/providers/jdbc/hooks/jdbc.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/providers/src/airflow/providers/jdbc/hooks/jdbc.py b/providers/src/airflow/providers/jdbc/hooks/jdbc.py index 283cb46e87980..808b946bd9762 100644 --- a/providers/src/airflow/providers/jdbc/hooks/jdbc.py +++ b/providers/src/airflow/providers/jdbc/hooks/jdbc.py @@ -20,12 +20,12 @@ import traceback import warnings from contextlib import contextmanager +from threading import RLock from typing import TYPE_CHECKING, Any import jaydebeapi import jpype from sqlalchemy.engine import URL -from wrapt import synchronized from airflow.exceptions import AirflowException from airflow.providers.common.sql.hooks.sql import DbApiHook @@ -99,6 +99,7 @@ def __init__( super().__init__(*args, **kwargs) self._driver_path = driver_path self._driver_class = driver_class + self.lock = RLock() @classmethod def get_ui_field_behaviour(cls) -> dict[str, Any]: @@ -178,20 +179,20 @@ def get_sqlalchemy_engine(self, engine_kwargs=None): return super().get_sqlalchemy_engine(engine_kwargs) - @synchronized def get_conn(self) -> jaydebeapi.Connection: conn: Connection = self.connection host: str = conn.host login: str = conn.login psw: str = conn.password - conn = jaydebeapi.connect( - jclassname=self.driver_class, - url=str(host), - driver_args=[str(login), str(psw)], - jars=self.driver_path.split(",") if self.driver_path else None, - ) - return conn + with self.lock: + conn = jaydebeapi.connect( + jclassname=self.driver_class, + url=str(host), + driver_args=[str(login), str(psw)], + jars=self.driver_path.split(",") if self.driver_path else None, + ) + return conn def set_autocommit(self, conn: jaydebeapi.Connection, autocommit: bool) -> None: """