summaryrefslogtreecommitdiff
path: root/tests/test_asyncio/test_cluster.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_asyncio/test_cluster.py')
-rw-r--r--tests/test_asyncio/test_cluster.py40
1 files changed, 40 insertions, 0 deletions
diff --git a/tests/test_asyncio/test_cluster.py b/tests/test_asyncio/test_cluster.py
index 38bcaf6..1997c95 100644
--- a/tests/test_asyncio/test_cluster.py
+++ b/tests/test_asyncio/test_cluster.py
@@ -788,6 +788,27 @@ class TestRedisClusterObj:
)
await rc.close()
+ def test_replace_cluster_node(self, r: RedisCluster) -> None:
+ prev_default_node = r.get_default_node()
+ r.replace_default_node()
+ assert r.get_default_node() != prev_default_node
+ r.replace_default_node(prev_default_node)
+ assert r.get_default_node() == prev_default_node
+
+ async def test_default_node_is_replaced_after_exception(self, r):
+ curr_default_node = r.get_default_node()
+ # CLUSTER NODES command is being executed on the default node
+ nodes = await r.cluster_nodes()
+ assert "myself" in nodes.get(curr_default_node.name).get("flags")
+ # Mock connection error for the default node
+ mock_node_resp_exc(curr_default_node, ConnectionError("error"))
+ # Test that the command succeed from a different node
+ nodes = await r.cluster_nodes()
+ assert "myself" not in nodes.get(curr_default_node.name).get("flags")
+ assert r.get_default_node() != curr_default_node
+ # Rollback to the old default node
+ r.replace_default_node(curr_default_node)
+
class TestClusterRedisCommands:
"""
@@ -2591,6 +2612,25 @@ class TestClusterPipeline:
*(self.test_multi_key_operation_with_multi_slots(r) for i in range(100)),
)
+ @pytest.mark.onlycluster
+ async def test_pipeline_with_default_node_error_command(self, create_redis):
+ """
+ Test that the default node is being replaced when it raises a relevant exception
+ """
+ r = await create_redis(cls=RedisCluster, flushdb=False)
+ curr_default_node = r.get_default_node()
+ err = ConnectionError("error")
+ cmd_count = await r.command_count()
+ mock_node_resp_exc(curr_default_node, err)
+ async with r.pipeline(transaction=False) as pipe:
+ pipe.command_count()
+ result = await pipe.execute(raise_on_error=False)
+ assert result[0] == err
+ assert r.get_default_node() != curr_default_node
+ pipe.command_count()
+ result = await pipe.execute(raise_on_error=False)
+ assert result[0] == cmd_count
+
@pytest.mark.ssl
class TestSSL: