From edd8ef57011879c2303ce30cf53a48dbffc821ee Mon Sep 17 00:00:00 2001 From: Kydoimos97 Date: Tue, 4 Aug 2026 01:38:32 -0600 Subject: [PATCH] fix(rds): add connect_timeout so pool/connection construction fails fast MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit RdsServiceGateway built its ThreadedConnectionPool (and the reconnect() pool rebuild / single-connection reconnect) with a bare DSN and no connect_timeout, so an unreachable or saturated database caused connection construction to block indefinitely — a consumer's startup could hang forever with no recovery path. This turned a transient prod DB CPU saturation into a multi-minute downstream outage (AiAxis 502 boot loop; Datadog BUG-96). Add a connect_timeout parameter (default 10s, floored to the libpq minimum of 2s) threaded through every connection site: the pool constructor, the reconnect() pool rebuild, and the single-connection reconnect. Connections now fail fast instead of hanging, letting consumers surface the error and recover. --- WrenchCL/Connect/RdsServiceGateway.py | 17 +++++++++-- tests/test_connect.py | 43 +++++++++++++++++++++++++-- 2 files changed, 55 insertions(+), 5 deletions(-) diff --git a/WrenchCL/Connect/RdsServiceGateway.py b/WrenchCL/Connect/RdsServiceGateway.py index 0bf6004..8268475 100644 --- a/WrenchCL/Connect/RdsServiceGateway.py +++ b/WrenchCL/Connect/RdsServiceGateway.py @@ -36,7 +36,11 @@ class RdsServiceGateway: """ def __init__( - self, multithreaded: bool = False, min_pool_size: int = 1, max_pool_size: int = 10 + self, + multithreaded: bool = False, + min_pool_size: int = 1, + max_pool_size: int = 10, + connect_timeout: int = 10, ): """ Initializes the RdsServiceGateway by establishing a connection or connection pool @@ -53,6 +57,9 @@ def __init__( self.test_mode = False self._min_pool_size = min_pool_size self._max_pool_size = max_pool_size + # libpq connect_timeout (seconds, min 2) so pool/connection construction fails + # fast instead of blocking forever when the DB is unreachable or saturated. + self._connect_timeout = max(2, int(connect_timeout)) self.client_manager = AwsClientHub() self.config = self.client_manager.config self.db_uri = self.client_manager.db_uri @@ -60,7 +67,10 @@ def __init__( if self.multithreaded: # Initialize a threaded connection pool using the URI self.pool: Optional[ThreadedConnectionPool] = ThreadedConnectionPool( - minconn=min_pool_size, maxconn=max_pool_size, dsn=self.db_uri + minconn=min_pool_size, + maxconn=max_pool_size, + dsn=self.db_uri, + connect_timeout=self._connect_timeout, ) else: # Establish a single connection if multithreading is not enabled @@ -86,11 +96,12 @@ def reconnect(self) -> None: minconn=self._min_pool_size, maxconn=self._max_pool_size, dsn=self.db_uri, + connect_timeout=self._connect_timeout, ) return psycopg2.extras.register_uuid() - new = psycopg2.connect(self.db_uri) + new = psycopg2.connect(self.db_uri, connect_timeout=self._connect_timeout) old = self.connection if old is not None and old is not new: try: diff --git a/tests/test_connect.py b/tests/test_connect.py index 7e37c47..d638b41 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -130,7 +130,7 @@ def test_rds_single_reconnect_swaps_connection_and_skips_in_test_mode(mock_hub_c assert svc.connection is new_conn old_conn.close.assert_called_once_with() - mock_connect.assert_called_once_with("postgresql://u:p@h:5432/d") + mock_connect.assert_called_once_with("postgresql://u:p@h:5432/d", connect_timeout=10) svc.set_test_mode(True) svc.reconnect() @@ -169,7 +169,7 @@ def test_rds_single_get_data_reconnects_once(mock_hub_cls, mock_connect): svc = RdsServiceGateway(multithreaded=False) assert svc.get_data("SELECT * FROM foo", payload=None) == [{"id": 1}] - mock_connect.assert_called_once_with("postgresql://u:p@h:5432/d") + mock_connect.assert_called_once_with("postgresql://u:p@h:5432/d", connect_timeout=10) old_conn.rollback.assert_called_once_with() @@ -238,3 +238,42 @@ def test_rds_pool_get_data_returns_healthy_connection(mock_hub_cls, mock_pool_cl mock_pool.putconn.assert_called_once_with(healthy_conn, close=False) mock_pool.closeall.assert_not_called() + + +# ───────────────────────────────────────────────────────────── +# connect_timeout — pool/connection construction must fail fast +# ───────────────────────────────────────────────────────────── + + +@patch("WrenchCL.Connect.RdsServiceGateway.ThreadedConnectionPool") +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_pool_uses_default_connect_timeout(mock_hub_cls, mock_pool_cls): + mock_hub_cls.return_value = _mock_rds_hub() + mock_pool_cls.return_value = MagicMock() + + RdsServiceGateway(multithreaded=True) + + assert mock_pool_cls.call_args.kwargs["connect_timeout"] == 10 + + +@patch("WrenchCL.Connect.RdsServiceGateway.ThreadedConnectionPool") +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_pool_honours_custom_connect_timeout(mock_hub_cls, mock_pool_cls): + mock_hub_cls.return_value = _mock_rds_hub() + mock_pool_cls.return_value = MagicMock() + + RdsServiceGateway(multithreaded=True, min_pool_size=2, max_pool_size=100, connect_timeout=7) + + assert mock_pool_cls.call_args.kwargs["connect_timeout"] == 7 + + +@patch("WrenchCL.Connect.RdsServiceGateway.ThreadedConnectionPool") +@patch("WrenchCL.Connect.RdsServiceGateway.AwsClientHub") +def test_rds_pool_connect_timeout_floored_at_libpq_minimum(mock_hub_cls, mock_pool_cls): + mock_hub_cls.return_value = _mock_rds_hub() + mock_pool_cls.return_value = MagicMock() + + # libpq rejects connect_timeout < 2; values below are coerced up to the floor. + RdsServiceGateway(multithreaded=True, connect_timeout=0) + + assert mock_pool_cls.call_args.kwargs["connect_timeout"] == 2