summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorBar Shaul <88437685+barshaul@users.noreply.github.com>2022-01-25 09:01:48 +0200
committerGitHub <noreply@github.com>2022-01-25 09:01:48 +0200
commit039257613f523e6a30736f218ca2d37d2f12320f (patch)
tree0b8f7bade605d5978d36925ce69db6d6ae7cb4ea
parentfa5841e50c242a7584eedda59ea918570adbf729 (diff)
downloadredis-py-039257613f523e6a30736f218ca2d37d2f12320f.tar.gz
Added retry mechanism on socket timeouts when connecting to the server (#1895)
-rwxr-xr-xredis/connection.py6
-rw-r--r--redis/retry.py6
-rw-r--r--tests/test_connection.py50
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)