From b511d48af0066c3bd676df1f8e81bd5252975583 Mon Sep 17 00:00:00 2001 From: Kydoimos97 Date: Mon, 3 Aug 2026 18:58:26 -0600 Subject: [PATCH] feat(rds): transparent stale-connection reconnect in RdsServiceGateway MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit RdsServiceGateway (single-connection mode) held one psycopg2 connection for the life of the process with no liveness check, so a long-running consumer whose connection was dropped server-side would fail its next query with 'connection already closed' — and the except-block rollback would itself throw on the dead socket, masking the original error. Adds reconnect() (rebuilds the single connection or the pool from db_uri) and classifies connection-level errors; get_data and update_database now reconnect and retry once on such an error only. Non-connection errors (IntegrityError, DataError, etc.) keep their exact prior behavior, and rollback calls are guarded so a dead-connection rollback can't mask the cause. update_database commits once at the end of each path, so a whole-operation retry cannot double-write. --- WrenchCL/Connect/RdsServiceGateway.py | 352 ++++++++++++++++---------- tests/test_connect.py | 132 +++++++++- 2 files changed, 354 insertions(+), 130 deletions(-) diff --git a/WrenchCL/Connect/RdsServiceGateway.py b/WrenchCL/Connect/RdsServiceGateway.py index 12b4ed6..88ff496 100644 --- a/WrenchCL/Connect/RdsServiceGateway.py +++ b/WrenchCL/Connect/RdsServiceGateway.py @@ -54,6 +54,8 @@ def __init__( self.client_manager = AwsClientHub() self.config = self.client_manager.config self.db_uri = self.client_manager.db_uri + self._min_pool_size = min_pool_size + self._max_pool_size = max_pool_size if self.multithreaded: # Initialize a threaded connection pool using the URI @@ -64,6 +66,54 @@ def __init__( # Establish a single connection if multithreading is not enabled self.connection: Optional["RDSClient"] = self.client_manager.db + def reconnect(self) -> None: + """Re-establish the configured database connection or connection pool.""" + if self.test_mode: + logger.info("Test mode active; skipping database reconnect.") + return + + logger.info("Reconnecting to the RDS database.") + if self.multithreaded: + try: + self.pool.closeall() + except (psycopg2.OperationalError, psycopg2.InterfaceError): + pass + self.pool = ThreadedConnectionPool( + minconn=self._min_pool_size, + maxconn=self._max_pool_size, + dsn=self.db_uri, + ) + else: + psycopg2.extras.register_uuid() + new_conn = psycopg2.connect(self.db_uri) + old_conn = self.connection + if old_conn is not None and old_conn is not new_conn: + try: + old_conn.close() + except (psycopg2.OperationalError, psycopg2.InterfaceError): + pass + self.connection = new_conn + logger.info("RDS database reconnect complete.") + + @staticmethod + def _is_connection_level_error(exc: BaseException) -> bool: + """Return whether an exception indicates that the database connection is stale.""" + if isinstance(exc, (psycopg2.OperationalError, psycopg2.InterfaceError)): + return True + if isinstance(exc, psycopg2.Error): + return False + message = str(exc).lower() + return any( + phrase in message + for phrase in ( + "connection already closed", + "server closed the connection", + "ssl connection has been closed", + "connection not open", + "terminating connection", + ) + ) + def set_test_mode(self, test_mode: bool = False): logger.warning("Test mode activated, database commits will not be commited.") self.test_mode = test_mode @@ -100,38 +150,62 @@ def get_data( """ Fetch data from the database based on the input query and parameters. """ - conn = self.get_connection() - try: - with conn.cursor(cursor_factory=psycopg2.extras.DictCursor) as cursor: - if show_query: + def _attempt() -> Optional[Any]: + conn = self.get_connection() + try: + with conn.cursor(cursor_factory=psycopg2.extras.DictCursor) as cursor: + if show_query: + logger._internal.log_internal( + "Mogrified Query:\n", cursor.mogrify(query, payload) + ) + else: + logger._internal.log_internal( + "Mogrified Query:\n", cursor.mogrify(query, payload) + ) + cursor.execute(query, payload) + data = cursor.fetchall() if fetchall else cursor.fetchone() logger._internal.log_internal( - "Mogrified Query:\n", cursor.mogrify(query, payload) + "Fetched data\n: %s", str(data)[:100] if fetchall else str(data) ) + if return_dict and data is not None: + return [dict(row) for row in data] if fetchall else dict(data) + elif data is None: + raise ValueError("None returned") else: - logger._internal.log_internal( - "Mogrified Query:\n", cursor.mogrify(query, payload) - ) - cursor.execute(query, payload) - data = cursor.fetchall() if fetchall else cursor.fetchone() - logger._internal.log_internal( - "Fetched data\n: %s", str(data)[:100] if fetchall else str(data) - ) - if return_dict and data is not None: - return [dict(row) for row in data] if fetchall else dict(data) - elif data is None: - raise ValueError("None returned") - else: - return data + return data + except Exception as e: + try: + conn.rollback() + except (psycopg2.OperationalError, psycopg2.InterfaceError): + pass + if self._is_connection_level_error(e): + raise e + if raise_on_error: + logger.warning(f"Error executing query: {e}") + raise e + else: + logger._internal.log_internal(f"Query returned None: {e}") + return None + finally: + self.release_connection(conn) + + try: + return _attempt() except Exception as e: - conn.rollback() - if raise_on_error: - logger.warning(f"Error executing query: {e}") - raise e - else: - logger._internal.log_internal(f"Query returned None: {e}") - return None - finally: - self.release_connection(conn) + if not self._is_connection_level_error(e): + raise + try: + self.reconnect() + return _attempt() + except Exception as retry_error: + if not self._is_connection_level_error(retry_error): + raise + if raise_on_error: + logger.warning(f"Error executing query: {retry_error}") + raise retry_error + else: + logger._internal.log_internal(f"Query returned None: {retry_error}") + return None def update_database( self, @@ -165,108 +239,132 @@ def update_database( ValueError: If column order is missing for DataFrame payloads or if the payload has incompatible data for batch processing. psycopg2.DataError: If no data was committed in batch processing. """ - conn = self.get_connection() - - if self.test_mode: - test_mode = True - - if test_mode: - logger.warning("Running RDSServiceGateway in test mode.") + original_payload = payload + + def _attempt() -> Optional[List[tuple]]: + nonlocal test_mode + conn = self.get_connection() + if self.test_mode: + test_mode = True + + if test_mode: + logger.warning("Running RDSServiceGateway in test mode.") + + try: + # Convert payload into a tuple if it's a single value or list + payload = original_payload + payload = self.convert_payload(payload) + logger._internal.log_internal(f"Converted payload: {payload}") + + if isinstance(payload, tuple): + logger._internal.log_internal("Payload is a single tuple.") + # Execute query for single tuple payload + with conn.cursor() as cursor: + cursor.execute(query, payload) + return_value = cursor.fetchall() if returning else None + if not test_mode: + conn.commit() + logger._internal.log_internal("Transaction committed successfully.") + else: + conn.rollback() + logger._internal.log_internal("Transaction rolled back in test mode.") + return return_value + + elif isinstance(payload, list) and all(isinstance(item, tuple) for item in payload): + logger._internal.log_internal("Payload is a list of tuples.") + # Execute batch query for list of tuples payload + with conn.cursor() as cursor: + psycopg2.extras.execute_values( + cursor, query, payload, page_size=self.config.db_batch_size + ) + return_value = cursor.fetchall() if returning else None + if not test_mode: + conn.commit() + logger._internal.log_internal("Transaction committed successfully.") + else: + conn.rollback() + logger._internal.log_internal("Transaction rolled back in test mode.") + return return_value + + elif isinstance(payload, pd.DataFrame) and column_order: + logger._internal.log_internal("Payload is a DataFrame with specified column order.") + # Batch processing for DataFrame payloads with specified column order + if returning: + raise ValueError( + "Returning values not compatible with batch processing, please use dictionary input" + ) + if not set(column_order).issubset(payload.columns): + missing_columns = set(column_order) - set(payload.columns) + raise ValueError( + f"The following columns are missing from the payload: {missing_columns}" + ) + + with conn.cursor() as cursor: + data_batch = [] + batch_counter = 1 + total_batches = math.ceil(len(payload) / self.config.db_batch_size) + + for i, row in enumerate(payload.itertuples(index=False, name="Row")): + data_batch.append(tuple(getattr(row, col) for col in column_order)) + + if len(data_batch) == self.config.db_batch_size or i == len(payload) - 1: + psycopg2.extras.execute_values( + cursor, query, data_batch, page_size=self.config.db_batch_size + ) + data_batch = [] + logger._internal.log_internal( + f"Processed batch {batch_counter}/{total_batches} successfully" + ) + batch_counter += 1 + + if batch_counter == 1: + raise psycopg2.DataError("Nothing to commit") + + if not test_mode: + conn.commit() + logger._internal.log_internal("Transaction committed successfully.") + else: + conn.rollback() + logger._internal.log_internal("Transaction rolled back in test mode.") + + except Exception as e: + try: + conn.rollback() + except (psycopg2.OperationalError, psycopg2.InterfaceError): + pass + if isinstance(e, IndexError): + try: + logger.warning( + f"Error processing batch: IndexError | Got {query.count('%s')} placeholders and {len(payload)} values. {e}" + ) + except Exception as nested_exception: + logger.warning( + f"Error processing batch: {str(e)}; Nested error: {str(nested_exception)}", + exc_info=True, + ) + else: + logger.warning(f"Error processing batch: {str(e)}", exc_info=True) + if self._is_connection_level_error(e): + raise e + if raise_on_error: + raise e + finally: + self.release_connection(conn) try: - # Convert payload into a tuple if it's a single value or list - payload = self.convert_payload(payload) - logger._internal.log_internal(f"Converted payload: {payload}") - - if isinstance(payload, tuple): - logger._internal.log_internal("Payload is a single tuple.") - # Execute query for single tuple payload - with conn.cursor() as cursor: - cursor.execute(query, payload) - return_value = cursor.fetchall() if returning else None - if not test_mode: - conn.commit() - logger._internal.log_internal("Transaction committed successfully.") - else: - conn.rollback() - logger._internal.log_internal("Transaction rolled back in test mode.") - return return_value - - elif isinstance(payload, list) and all(isinstance(item, tuple) for item in payload): - logger._internal.log_internal("Payload is a list of tuples.") - # Execute batch query for list of tuples payload - with conn.cursor() as cursor: - psycopg2.extras.execute_values( - cursor, query, payload, page_size=self.config.db_batch_size - ) - return_value = cursor.fetchall() if returning else None - if not test_mode: - conn.commit() - logger._internal.log_internal("Transaction committed successfully.") - else: - conn.rollback() - logger._internal.log_internal("Transaction rolled back in test mode.") - return return_value - - elif isinstance(payload, pd.DataFrame) and column_order: - logger._internal.log_internal("Payload is a DataFrame with specified column order.") - # Batch processing for DataFrame payloads with specified column order - if returning: - raise ValueError( - "Returning values not compatible with batch processing, please use dictionary input" - ) - if not set(column_order).issubset(payload.columns): - missing_columns = set(column_order) - set(payload.columns) - raise ValueError( - f"The following columns are missing from the payload: {missing_columns}" - ) - - with conn.cursor() as cursor: - data_batch = [] - batch_counter = 1 - total_batches = math.ceil(len(payload) / self.config.db_batch_size) - - for i, row in enumerate(payload.itertuples(index=False, name="Row")): - data_batch.append(tuple(getattr(row, col) for col in column_order)) - - if len(data_batch) == self.config.db_batch_size or i == len(payload) - 1: - psycopg2.extras.execute_values( - cursor, query, data_batch, page_size=self.config.db_batch_size - ) - data_batch = [] - logger._internal.log_internal( - f"Processed batch {batch_counter}/{total_batches} successfully" - ) - batch_counter += 1 - - if batch_counter == 1: - raise psycopg2.DataError("Nothing to commit") - - if not test_mode: - conn.commit() - logger._internal.log_internal("Transaction committed successfully.") - else: - conn.rollback() - logger._internal.log_internal("Transaction rolled back in test mode.") - + return _attempt() except Exception as e: - conn.rollback() - if isinstance(e, IndexError): - try: - logger.warning( - f"Error processing batch: IndexError | Got {query.count('%s')} placeholders and {len(payload)} values. {e}" - ) - except Exception as nested_exception: - logger.warning( - f"Error processing batch: {str(e)}; Nested error: {str(nested_exception)}", - exc_info=True, - ) - else: - logger.warning(f"Error processing batch: {str(e)}", exc_info=True) - if raise_on_error: - raise e - finally: - self.release_connection(conn) + if not self._is_connection_level_error(e): + raise + try: + self.reconnect() + return _attempt() + except Exception as retry_error: + if not self._is_connection_level_error(retry_error): + raise + if raise_on_error: + raise retry_error + return None def format_sql_query(self, query: str, payload: tuple) -> None: """ diff --git a/tests/test_connect.py b/tests/test_connect.py index 08f6cf9..9e40415 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -1,6 +1,9 @@ # tests/test_connect.py -from unittest.mock import patch, PropertyMock +from unittest.mock import MagicMock, PropertyMock, patch + +import psycopg2 +import pytest try: from WrenchCL.Connect import AwsClientHub, RdsServiceGateway, S3ServiceGateway @@ -60,11 +63,16 @@ def test_get_s3_client(mock_session_prop, mock_config_prop, ): # RdsServiceGateway Tests # ───────────────────────────────────────────────────────────── -from unittest.mock import patch, MagicMock - from WrenchCL.Connect import RdsServiceGateway +@pytest.fixture(autouse=True) +def reset_rds_gateway_singleton(): + RdsServiceGateway._SingletonWrapper__cls_instance = None + yield + RdsServiceGateway._SingletonWrapper__cls_instance = None + + @patch("WrenchCL.Connect.RdsServiceGateway.ThreadedConnectionPool") @patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") def test_rds_multithreaded_connection(mock_hub_cls, mock_pool_cls): @@ -104,3 +112,121 @@ def test_rds_update_tuple_commit(mock_hub_cls): svc = RdsServiceGateway(multithreaded=False) result = svc.update_database("UPDATE table SET x = %s", payload=("val",), returning=True) assert result == [{'id': 1}] + + +@patch("WrenchCL.Connect.RdsServiceGateway.psycopg2.connect") +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_reconnect_swaps_single_connection(mock_hub_cls, mock_connect): + old_conn = MagicMock() + new_conn = MagicMock() + mock_hub = MagicMock() + mock_hub.db = old_conn + mock_hub.db_uri = "postgresql://u:p@h:5432/d" + mock_hub.config.db_batch_size = 100 + mock_hub_cls.return_value = mock_hub + mock_connect.return_value = new_conn + + svc = RdsServiceGateway(multithreaded=False) + svc.reconnect() + + mock_connect.assert_called_once_with(mock_hub.db_uri) + old_conn.close.assert_called_once_with() + assert svc.connection is new_conn + + +@patch("WrenchCL.Connect.RdsServiceGateway.psycopg2.connect") +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_reconnect_is_noop_in_test_mode(mock_hub_cls, mock_connect): + mock_hub = MagicMock() + mock_hub.db = MagicMock() + mock_hub.db_uri = "postgresql://u:p@h:5432/d" + mock_hub.config.db_batch_size = 100 + mock_hub_cls.return_value = mock_hub + + svc = RdsServiceGateway(multithreaded=False) + svc.set_test_mode(True) + svc.reconnect() + + mock_connect.assert_not_called() + + +@pytest.mark.parametrize( + ("exception", "expected"), + [ + (psycopg2.OperationalError("stale"), True), + (psycopg2.InterfaceError("stale"), True), + (psycopg2.IntegrityError("constraint"), False), + (Exception("connection already closed"), True), + (Exception("unrelated failure"), False), + ], +) +def test_is_connection_level_error(exception, expected): + assert RdsServiceGateway._is_connection_level_error(exception) is expected + + +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_get_data_reconnects_and_retries_once(mock_hub_cls): + first_cursor = MagicMock() + first_cursor.__enter__.return_value = first_cursor + first_cursor.execute.side_effect = psycopg2.OperationalError("connection already closed") + second_cursor = MagicMock() + second_cursor.__enter__.return_value = second_cursor + second_cursor.fetchall.return_value = [{"id": 1}] + + conn = MagicMock() + conn.cursor.side_effect = [first_cursor, second_cursor] + mock_hub = MagicMock() + mock_hub.db = conn + mock_hub.db_uri = "postgresql://u:p@h:5432/d" + mock_hub.config.db_batch_size = 100 + mock_hub_cls.return_value = mock_hub + + svc = RdsServiceGateway(multithreaded=False) + with patch.object(svc, "reconnect") as reconnect: + result = svc.get_data("SELECT * FROM foo", payload=None) + + reconnect.assert_called_once_with() + assert result == [{"id": 1}] + assert conn.cursor.call_count == 2 + + +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_get_data_integrity_error_does_not_reconnect(mock_hub_cls): + cursor = MagicMock() + cursor.__enter__.return_value = cursor + cursor.execute.side_effect = psycopg2.IntegrityError("constraint") + conn = MagicMock() + conn.cursor.return_value = cursor + mock_hub = MagicMock() + mock_hub.db = conn + mock_hub.db_uri = "postgresql://u:p@h:5432/d" + mock_hub.config.db_batch_size = 100 + mock_hub_cls.return_value = mock_hub + + svc = RdsServiceGateway(multithreaded=False) + with patch.object(svc, "reconnect") as reconnect: + result = svc.get_data("SELECT * FROM foo", payload=None) + + reconnect.assert_not_called() + assert result is None + + +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_update_integrity_error_does_not_reconnect(mock_hub_cls): + cursor = MagicMock() + cursor.__enter__.return_value = cursor + cursor.execute.side_effect = psycopg2.IntegrityError("constraint") + conn = MagicMock() + conn.cursor.return_value = cursor + mock_hub = MagicMock() + mock_hub.db = conn + mock_hub.db_uri = "postgresql://u:p@h:5432/d" + mock_hub.config.db_batch_size = 100 + mock_hub_cls.return_value = mock_hub + + svc = RdsServiceGateway(multithreaded=False) + with patch.object(svc, "reconnect") as reconnect: + with pytest.raises(psycopg2.IntegrityError): + svc.update_database("UPDATE table SET x = %s", payload=("val",)) + + reconnect.assert_not_called()