Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 14 additions & 3 deletions WrenchCL/Connect/RdsServiceGateway.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -53,14 +57,20 @@ 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

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
Expand All @@ -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:
Expand Down
43 changes: 41 additions & 2 deletions tests/test_connect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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()


Expand Down Expand Up @@ -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
Loading