Skip to content

Commit 6bdf2fb

Browse files
committed
Add paging for ControlConnection topology queries
1 parent bcc2d3d commit 6bdf2fb

3 files changed

Lines changed: 201 additions & 26 deletions

File tree

cassandra/cluster.py

Lines changed: 38 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -3929,23 +3929,32 @@ def _try_connect(self, endpoint):
39293929
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
39303930
sel_local = self._SELECT_LOCAL if self._token_meta_enabled else self._SELECT_LOCAL_NO_TOKENS
39313931
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3932-
consistency_level=ConsistencyLevel.ONE)
3932+
consistency_level=ConsistencyLevel.ONE,
3933+
fetch_size=self._schema_meta_page_size)
39333934
local_query = QueryMessage(query=maybe_add_timeout_to_query(sel_local, self._metadata_request_timeout),
3934-
consistency_level=ConsistencyLevel.ONE)
3935-
(peers_success, peers_result), (local_success, local_result) = connection.wait_for_responses(
3936-
peers_query, local_query, timeout=self._timeout, fail_on_error=False)
3937-
3938-
if not local_success:
3939-
raise local_result
3940-
3941-
if not peers_success:
3942-
# error with the peers v2 query, fallback to peers v1
3943-
self._uses_peers_v2 = False
3944-
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
3945-
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3946-
consistency_level=ConsistencyLevel.ONE)
3947-
peers_result = connection.wait_for_response(
3948-
peers_query, timeout=self._timeout)
3935+
consistency_level=ConsistencyLevel.ONE,
3936+
fetch_size=self._schema_meta_page_size)
3937+
3938+
with ThreadPoolExecutor(max_workers=2) as executor:
3939+
local_future = executor.submit(
3940+
connection.fetch_all_pages, local_query, self._timeout, False)
3941+
peers_future = executor.submit(
3942+
connection.fetch_all_pages, peers_query, self._timeout, False)
3943+
3944+
local_success, local_result = local_future.result()
3945+
3946+
if not local_success:
3947+
raise local_result
3948+
3949+
peers_success, peers_result = peers_future.result()
3950+
3951+
if not peers_success:
3952+
self._uses_peers_v2 = False
3953+
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
3954+
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3955+
consistency_level=ConsistencyLevel.ONE,
3956+
fetch_size=self._schema_meta_page_size)
3957+
peers_result = connection.fetch_all_pages(peers_query, self._timeout)
39493958

39503959
shared_results = (peers_result, local_result)
39513960
self._refresh_node_list_and_token_map(connection, preloaded_results=shared_results)
@@ -4088,11 +4097,20 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None,
40884097
log.debug("[control connection] Refreshing node list and token map")
40894098
sel_local = self._SELECT_LOCAL
40904099
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
4091-
consistency_level=cl)
4100+
consistency_level=cl,
4101+
fetch_size=self._schema_meta_page_size)
40924102
local_query = QueryMessage(query=maybe_add_timeout_to_query(sel_local, self._metadata_request_timeout),
4093-
consistency_level=cl)
4094-
peers_result, local_result = connection.wait_for_responses(
4095-
peers_query, local_query, timeout=self._timeout)
4103+
consistency_level=cl,
4104+
fetch_size=self._schema_meta_page_size)
4105+
4106+
with ThreadPoolExecutor(max_workers=2) as executor:
4107+
peers_future = executor.submit(
4108+
connection.fetch_all_pages, peers_query, self._timeout)
4109+
local_future = executor.submit(
4110+
connection.fetch_all_pages, local_query, self._timeout)
4111+
4112+
peers_result = peers_future.result()
4113+
local_result = local_future.result()
40964114

40974115
peers_result = dict_factory(peers_result.column_names, peers_result.parsed_rows)
40984116

cassandra/connection.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1288,6 +1288,63 @@ def wait_for_responses(self, *msgs, **kwargs):
12881288
self.defunct(exc)
12891289
raise
12901290

1291+
def fetch_all_pages(self, query_msg, timeout, fail_on_error=True):
1292+
"""Fetch all pages for a query, following paging_state until exhausted.
1293+
1294+
Runs the given query and, if the response has a paging_state,
1295+
continues fetching subsequent pages until no paging_state remains.
1296+
Concatenates all parsed_rows into a single result.
1297+
1298+
Args:
1299+
query_msg: QueryMessage to execute (paging_state may be pre-set).
1300+
timeout: Per-request timeout passed to wait_for_response.
1301+
fail_on_error: If True (default), raises on error. If False,
1302+
returns (success, result_or_error).
1303+
1304+
Returns:
1305+
When fail_on_error=True: the fully accumulated ResultMessage.
1306+
When fail_on_error=False: (True, ResultMessage) or
1307+
(False, Exception).
1308+
"""
1309+
response = self.wait_for_response(query_msg, timeout=timeout, fail_on_error=fail_on_error)
1310+
1311+
if not fail_on_error:
1312+
success, result = response
1313+
if not success:
1314+
return response
1315+
else:
1316+
result = response
1317+
1318+
if not result or not result.paging_state:
1319+
return response if not fail_on_error else result
1320+
1321+
all_rows = result.parsed_rows
1322+
if all_rows is None:
1323+
all_rows = []
1324+
original_paging_state = query_msg.paging_state
1325+
1326+
try:
1327+
while result and result.paging_state:
1328+
query_msg.paging_state = result.paging_state
1329+
page_response = self.wait_for_response(query_msg, timeout=timeout, fail_on_error=fail_on_error)
1330+
1331+
if not fail_on_error:
1332+
page_success, page_result = page_response
1333+
if not page_success:
1334+
return page_response
1335+
result = page_result
1336+
else:
1337+
result = page_response
1338+
1339+
if result and result.parsed_rows:
1340+
all_rows.extend(result.parsed_rows)
1341+
finally:
1342+
query_msg.paging_state = original_paging_state
1343+
1344+
result.parsed_rows = all_rows
1345+
1346+
return (True, result) if not fail_on_error else result
1347+
12911348
def register_watcher(self, event_type, callback, register_timeout=None):
12921349
"""
12931350
Register a callback for a given event type.

tests/unit/test_control_connection.py

Lines changed: 106 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -15,10 +15,10 @@
1515
import unittest
1616

1717
from concurrent.futures import ThreadPoolExecutor
18-
from unittest.mock import Mock, ANY, call, patch
18+
from unittest.mock import Mock, ANY, call, patch, MagicMock
1919

20-
from cassandra import OperationTimedOut, SchemaTargetType, SchemaChangeType
21-
from cassandra.protocol import ResultMessage, RESULT_KIND_ROWS
20+
from cassandra import OperationTimedOut, SchemaTargetType, SchemaChangeType, ConsistencyLevel
21+
from cassandra.protocol import ResultMessage, RESULT_KIND_ROWS, QueryMessage
2222
from cassandra.cluster import ControlConnection, _Scheduler, ProfileManager, EXEC_PROFILE_DEFAULT, ExecutionProfile
2323
from cassandra.pool import Host
2424
from cassandra.connection import EndPoint, DefaultEndPoint, DefaultEndPointFactory
@@ -168,6 +168,21 @@ def __init__(self):
168168
]
169169
self.wait_for_responses = Mock(return_value=_node_meta_results(self.local_results, self.peer_results))
170170

171+
def wait_for_response_side_effect(query_msg, timeout=None, fail_on_error=True):
172+
result = ResultMessage(kind=RESULT_KIND_ROWS)
173+
if "peers" in query_msg.query.lower():
174+
result.column_names = self.peer_results[0]
175+
result.parsed_rows = self.peer_results[1]
176+
else:
177+
result.column_names = self.local_results[0]
178+
result.parsed_rows = self.local_results[1]
179+
result.paging_state = None
180+
return result
181+
self.wait_for_response = Mock(side_effect=wait_for_response_side_effect)
182+
183+
def fetch_all_pages(self, query_msg, timeout, fail_on_error=True):
184+
return self.wait_for_response(query_msg, timeout=timeout, fail_on_error=fail_on_error)
185+
171186

172187
class FakeTime(object):
173188

@@ -312,6 +327,91 @@ def test_wait_for_schema_agreement_none_timeout(self):
312327
cc._time = self.time
313328
assert cc._wait_for_schema_agreement()
314329

330+
def test_topology_queries_use_paging(self):
331+
self.control_connection.refresh_node_list_and_token_map()
332+
assert self.connection.wait_for_response.called
333+
calls = self.connection.wait_for_response.call_args_list
334+
for call in calls:
335+
query_msg = call[0][0]
336+
assert isinstance(query_msg, QueryMessage)
337+
assert query_msg.fetch_size == self.control_connection._schema_meta_page_size
338+
339+
def test_topology_queries_fetch_all_pages(self):
340+
from cassandra.connection import Connection as RealConnection
341+
mock_connection = MagicMock()
342+
mock_connection.endpoint = DefaultEndPoint("192.168.1.0")
343+
mock_connection.original_endpoint = mock_connection.endpoint
344+
mock_connection.fetch_all_pages = RealConnection.fetch_all_pages.__get__(mock_connection, RealConnection)
345+
first_page = ResultMessage(kind=RESULT_KIND_ROWS)
346+
first_page.column_names = ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"]
347+
first_page.parsed_rows = [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"]]
348+
first_page.paging_state = b"has_more_pages"
349+
second_page = ResultMessage(kind=RESULT_KIND_ROWS)
350+
second_page.column_names = ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"]
351+
second_page.parsed_rows = [["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], "uuid3"]]
352+
second_page.paging_state = None
353+
mock_connection.wait_for_response.side_effect = [first_page, second_page]
354+
self.control_connection._connection = mock_connection
355+
query_msg = QueryMessage(query="SELECT * FROM system.peers",
356+
consistency_level=ConsistencyLevel.ONE,
357+
fetch_size=self.control_connection._schema_meta_page_size)
358+
result = mock_connection.fetch_all_pages(query_msg, timeout=5)
359+
assert len(result.parsed_rows) == 2
360+
assert result.parsed_rows[0][0] == "192.168.1.1"
361+
assert result.parsed_rows[1][0] == "192.168.1.2"
362+
assert result.paging_state is None
363+
assert mock_connection.wait_for_response.call_count == 2
364+
365+
def test_topology_queries_fetch_all_pages_fail_on_error_false(self):
366+
from cassandra.connection import Connection as RealConnection
367+
mock_connection = MagicMock()
368+
mock_connection.endpoint = DefaultEndPoint("192.168.1.0")
369+
mock_connection.original_endpoint = mock_connection.endpoint
370+
mock_connection.fetch_all_pages = RealConnection.fetch_all_pages.__get__(mock_connection, RealConnection)
371+
first_page = ResultMessage(kind=RESULT_KIND_ROWS)
372+
first_page.column_names = ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"]
373+
first_page.parsed_rows = [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"]]
374+
first_page.paging_state = b"has_more_pages"
375+
second_page = ResultMessage(kind=RESULT_KIND_ROWS)
376+
second_page.column_names = ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"]
377+
second_page.parsed_rows = [["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], "uuid3"]]
378+
second_page.paging_state = None
379+
mock_connection.wait_for_response.side_effect = [
380+
(True, first_page),
381+
(True, second_page),
382+
]
383+
query_msg = QueryMessage(query="SELECT * FROM system.peers",
384+
consistency_level=ConsistencyLevel.ONE,
385+
fetch_size=self.control_connection._schema_meta_page_size)
386+
success, result = mock_connection.fetch_all_pages(query_msg, timeout=5, fail_on_error=False)
387+
assert success
388+
assert len(result.parsed_rows) == 2
389+
assert result.parsed_rows[0][0] == "192.168.1.1"
390+
assert result.parsed_rows[1][0] == "192.168.1.2"
391+
assert mock_connection.wait_for_response.call_count == 2
392+
393+
def test_topology_queries_fetch_all_pages_page_failure(self):
394+
from cassandra.connection import Connection as RealConnection
395+
mock_connection = MagicMock()
396+
mock_connection.endpoint = DefaultEndPoint("192.168.1.0")
397+
mock_connection.original_endpoint = mock_connection.endpoint
398+
mock_connection.fetch_all_pages = RealConnection.fetch_all_pages.__get__(mock_connection, RealConnection)
399+
first_page = ResultMessage(kind=RESULT_KIND_ROWS)
400+
first_page.column_names = ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"]
401+
first_page.parsed_rows = [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"]]
402+
first_page.paging_state = b"has_more_pages"
403+
mock_connection.wait_for_response.side_effect = [
404+
(True, first_page),
405+
(False, OperationTimedOut()),
406+
]
407+
query_msg = QueryMessage(query="SELECT * FROM system.peers",
408+
consistency_level=ConsistencyLevel.ONE,
409+
fetch_size=self.control_connection._schema_meta_page_size)
410+
success, error = mock_connection.fetch_all_pages(query_msg, timeout=5, fail_on_error=False)
411+
assert not success
412+
assert isinstance(error, OperationTimedOut)
413+
assert mock_connection.wait_for_response.call_count == 2
414+
315415
def test_refresh_nodes_and_tokens(self):
316416
self.control_connection.refresh_node_list_and_token_map()
317417
meta = self.cluster.metadata
@@ -328,7 +428,7 @@ def test_refresh_nodes_and_tokens(self):
328428
assert host.datacenter == "dc1"
329429
assert host.rack == "rack1"
330430

331-
assert self.connection.wait_for_responses.call_count == 1
431+
assert self.connection.wait_for_response.call_count == 2
332432

333433
def test_refresh_nodes_and_tokens_with_invalid_peers(self):
334434
def refresh_and_validate_added_hosts():
@@ -444,11 +544,11 @@ def test_refresh_nodes_and_tokens_remove_host(self):
444544

445545
def test_refresh_nodes_and_tokens_timeout(self):
446546

447-
def bad_wait_for_responses(*args, **kwargs):
547+
def bad_wait_for_response(*args, **kwargs):
448548
assert kwargs['timeout'] == self.control_connection._timeout
449549
raise OperationTimedOut()
450550

451-
self.connection.wait_for_responses = bad_wait_for_responses
551+
self.connection.wait_for_response = Mock(side_effect=bad_wait_for_response)
452552
self.control_connection.refresh_node_list_and_token_map()
453553
self.cluster.executor.submit.assert_called_with(self.control_connection._reconnect)
454554

0 commit comments

Comments
 (0)