Skip to content

Commit 5aa498e

Browse files
committed
Add paging for ControlConnection topology queries
1 parent c1bfd54 commit 5aa498e

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
@@ -3925,23 +3925,32 @@ def _try_connect(self, endpoint):
39253925
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
39263926
sel_local = self._SELECT_LOCAL if self._token_meta_enabled else self._SELECT_LOCAL_NO_TOKENS
39273927
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3928-
consistency_level=ConsistencyLevel.ONE)
3928+
consistency_level=ConsistencyLevel.ONE,
3929+
fetch_size=self._schema_meta_page_size)
39293930
local_query = QueryMessage(query=maybe_add_timeout_to_query(sel_local, self._metadata_request_timeout),
3930-
consistency_level=ConsistencyLevel.ONE)
3931-
(peers_success, peers_result), (local_success, local_result) = connection.wait_for_responses(
3932-
peers_query, local_query, timeout=self._timeout, fail_on_error=False)
3933-
3934-
if not local_success:
3935-
raise local_result
3936-
3937-
if not peers_success:
3938-
# error with the peers v2 query, fallback to peers v1
3939-
self._uses_peers_v2 = False
3940-
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
3941-
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3942-
consistency_level=ConsistencyLevel.ONE)
3943-
peers_result = connection.wait_for_response(
3944-
peers_query, timeout=self._timeout)
3931+
consistency_level=ConsistencyLevel.ONE,
3932+
fetch_size=self._schema_meta_page_size)
3933+
3934+
with ThreadPoolExecutor(max_workers=2) as executor:
3935+
local_future = executor.submit(
3936+
connection.fetch_all_pages, local_query, self._timeout, False)
3937+
peers_future = executor.submit(
3938+
connection.fetch_all_pages, peers_query, self._timeout, False)
3939+
3940+
local_success, local_result = local_future.result()
3941+
3942+
if not local_success:
3943+
raise local_result
3944+
3945+
peers_success, peers_result = peers_future.result()
3946+
3947+
if not peers_success:
3948+
self._uses_peers_v2 = False
3949+
sel_peers = self._get_peers_query(self.PeersQueryType.PEERS, connection)
3950+
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
3951+
consistency_level=ConsistencyLevel.ONE,
3952+
fetch_size=self._schema_meta_page_size)
3953+
peers_result = connection.fetch_all_pages(peers_query, self._timeout)
39453954

39463955
shared_results = (peers_result, local_result)
39473956
self._refresh_node_list_and_token_map(connection, preloaded_results=shared_results)
@@ -4084,11 +4093,20 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None,
40844093
log.debug("[control connection] Refreshing node list and token map")
40854094
sel_local = self._SELECT_LOCAL
40864095
peers_query = QueryMessage(query=maybe_add_timeout_to_query(sel_peers, self._metadata_request_timeout),
4087-
consistency_level=cl)
4096+
consistency_level=cl,
4097+
fetch_size=self._schema_meta_page_size)
40884098
local_query = QueryMessage(query=maybe_add_timeout_to_query(sel_local, self._metadata_request_timeout),
4089-
consistency_level=cl)
4090-
peers_result, local_result = connection.wait_for_responses(
4091-
peers_query, local_query, timeout=self._timeout)
4099+
consistency_level=cl,
4100+
fetch_size=self._schema_meta_page_size)
4101+
4102+
with ThreadPoolExecutor(max_workers=2) as executor:
4103+
peers_future = executor.submit(
4104+
connection.fetch_all_pages, peers_query, self._timeout)
4105+
local_future = executor.submit(
4106+
connection.fetch_all_pages, local_query, self._timeout)
4107+
4108+
peers_result = peers_future.result()
4109+
local_result = local_future.result()
40924110

40934111
peers_result = dict_factory(peers_result.column_names, peers_result.parsed_rows)
40944112

cassandra/connection.py

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

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