diff options
| author | Bar Shaul <88437685+barshaul@users.noreply.github.com> | 2022-01-25 09:01:48 +0200 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2022-01-25 09:01:48 +0200 |
| commit | 039257613f523e6a30736f218ca2d37d2f12320f (patch) | |
| tree | 0b8f7bade605d5978d36925ce69db6d6ae7cb4ea | |
| parent | fa5841e50c242a7584eedda59ea918570adbf729 (diff) | |
| download | redis-py-039257613f523e6a30736f218ca2d37d2f12320f.tar.gz | |
Added retry mechanism on socket timeouts when connecting to the server (#1895)
| -rwxr-xr-x | redis/connection.py | 6 | ||||
| -rw-r--r-- | redis/retry.py | 6 | ||||
| -rw-r--r-- | tests/test_connection.py | 50 |
3 files changed, 58 insertions, 4 deletions
diff --git a/redis/connection.py b/redis/connection.py index 5fdac54..508c196 100755 --- a/redis/connection.py +++ b/redis/connection.py @@ -604,7 +604,9 @@ class Connection: if self._sock: return try: - sock = self._connect() + sock = self.retry.call_with_retry( + lambda: self._connect(), lambda error: self.disconnect(error) + ) except socket.timeout: raise TimeoutError("Timeout connecting to server") except OSError as e: @@ -721,7 +723,7 @@ class Connection: if str_if_bytes(self.read_response()) != "OK": raise ConnectionError("Invalid Database") - def disconnect(self): + def disconnect(self, *args): "Disconnects from the Redis server" self._parser.on_disconnect() if self._sock is None: diff --git a/redis/retry.py b/redis/retry.py index 6147fbd..3dced35 100644 --- a/redis/retry.py +++ b/redis/retry.py @@ -1,3 +1,4 @@ +import socket from time import sleep from redis.exceptions import ConnectionError, TimeoutError @@ -7,7 +8,10 @@ class Retry: """Retry a specific number of times after a failure""" def __init__( - self, backoff, retries, supported_errors=(ConnectionError, TimeoutError) + self, + backoff, + retries, + supported_errors=(ConnectionError, TimeoutError, socket.timeout), ): """ Initialize a `Retry` object with a `Backoff` object diff --git a/tests/test_connection.py b/tests/test_connection.py index d94a815..d9251c3 100644 --- a/tests/test_connection.py +++ b/tests/test_connection.py @@ -1,10 +1,14 @@ +import socket import types from unittest import mock +from unittest.mock import patch import pytest +from redis.backoff import NoBackoff from redis.connection import Connection -from redis.exceptions import InvalidResponse +from redis.exceptions import ConnectionError, InvalidResponse, TimeoutError +from redis.retry import Retry from redis.utils import HIREDIS_AVAILABLE from .conftest import skip_if_server_version_lt @@ -74,3 +78,47 @@ class TestConnection: mock_sock.shutdown.assert_called_once() mock_sock.close.assert_called_once() assert conn._sock is None + + def clear(self, conn): + conn.retry_on_error.clear() + + def test_retry_connect_on_timeout_error(self): + """Test that the _connect function is retried in case of a timeout""" + conn = Connection(retry_on_timeout=True, retry=Retry(NoBackoff(), 3)) + origin_connect = conn._connect + conn._connect = mock.Mock() + + def mock_connect(): + # connect only on the last retry + if conn._connect.call_count <= 2: + raise socket.timeout + else: + return origin_connect() + + conn._connect.side_effect = mock_connect + conn.connect() + assert conn._connect.call_count == 3 + self.clear(conn) + + def test_connect_without_retry_on_os_error(self): + """Test that the _connect function is not being retried in case of a OSError""" + with patch.object(Connection, "_connect") as _connect: + _connect.side_effect = OSError("") + conn = Connection(retry_on_timeout=True, retry=Retry(NoBackoff(), 2)) + with pytest.raises(ConnectionError): + conn.connect() + assert _connect.call_count == 1 + self.clear(conn) + + def test_connect_timeout_error_without_retry(self): + """Test that the _connect function is not being retried if retry_on_timeout is + set to False""" + conn = Connection(retry_on_timeout=False) + conn._connect = mock.Mock() + conn._connect.side_effect = socket.timeout + + with pytest.raises(TimeoutError) as e: + conn.connect() + assert conn._connect.call_count == 1 + assert str(e.value) == "Timeout connecting to server" + self.clear(conn) |
