diff --git a/cassandra/concurrent.py b/cassandra/concurrent.py index 0e7bf794e0..8c834e9101 100644 --- a/cassandra/concurrent.py +++ b/cassandra/concurrent.py @@ -14,9 +14,8 @@ from collections import namedtuple -from heapq import heappush, heappop from itertools import cycle -from threading import Condition +from queue import Empty, SimpleQueue from cassandra.cluster import ResultSet, EXEC_PROFILE_DEFAULT @@ -42,10 +41,9 @@ def execute_concurrent(session, statements_and_parameters, concurrency=100, rais * If :const:`False`, the results are returned only after all requests have completed. * If :const:`True`, a generator expression is returned. Using a generator results in a constrained - memory footprint when the results set will be large -- results are yielded - as they return instead of materializing the entire list at once. The trade for lower memory - footprint is marginal CPU overhead (more thread coordination and sorting out-of-order results - on-the-fly). + memory footprint when the results set will be large -- results are yielded as they return + instead of materializing the entire list at once. Results are still returned in the order the + statements were passed in, so out-of-order completions are held until their turn. `execution_profile` argument is the execution profile to use for this request, it is passed directly to :meth:`Session.execute_async`. @@ -74,10 +72,12 @@ def execute_concurrent(session, statements_and_parameters, concurrency=100, rais else: process_user(result[0]) # result will be a list of rows - Note: in the case that `generators` are used, it is important to ensure the consumers do not - block or attempt further synchronous requests, because no further IO will be processed until - the consumer returns. This may also produce a deadlock in the IO event thread. + Requests are submitted from the calling thread, never from the IO event thread. With + `results_generator`, new requests are only submitted while the consumer is asking for the + next result, so a slow consumer lowers the effective concurrency. """ + if not isinstance(concurrency, int): + raise TypeError("concurrency must be an integer") if concurrency <= 0: raise ValueError("concurrency must be greater than 0") @@ -90,131 +90,119 @@ def execute_concurrent(session, statements_and_parameters, concurrency=100, rais class _ConcurrentExecutor(object): + # All submission happens on the calling thread. IO-thread callbacks only + # enqueue the completed result, so they never block on the caller. def __init__(self, session, statements_and_params, execution_profile): self.session = session - self._enum_statements = enumerate(iter(statements_and_params)) + self._statements = iter(statements_and_params) self._execution_profile = execution_profile - self._condition = Condition() + self._done = SimpleQueue() + self._results = {} + self._submitted = 0 + self._in_flight = 0 + self._exhausted = False self._fail_fast = False - self._results_queue = [] - self._current = 0 - self._exec_count = 0 - self._executing = False - - def execute(self, concurrency, fail_fast): - self._fail_fast = fail_fast - self._results_queue = [] - self._current = 0 - self._exec_count = 0 - with self._condition: - for n in range(concurrency): - if not self._execute_next(): - break - return self._results() - - def _execute_next(self): - # lock must be held - try: - (idx, (statement, params)) = next(self._enum_statements) - self._exec_count += 1 - self._execute(idx, statement, params) - return True - except StopIteration: - pass - - def _execute(self, idx, statement, params): - # When execute_async completes synchronously (e.g. immediate timeout), - # the errback fires inline: _on_error -> _put_result -> _execute_next - # -> _execute. Without protection this recurses once per remaining - # statement and blows the stack. - # - # ``_executing`` marks that we are already inside this method higher up - # the call stack. When a synchronous callback re-enters, we just stash - # the pending work in ``_pending_executions`` and let the outermost - # invocation drain it in a loop -- no recursion. - if self._executing: - self._pending_executions.append((idx, statement, params)) - return - - self._executing = True - self._pending_executions = [(idx, statement, params)] - try: - while self._pending_executions: - p_idx, p_statement, p_params = self._pending_executions.pop(0) - try: - future = self.session.execute_async(p_statement, p_params, timeout=None, execution_profile=self._execution_profile) - args = (future, p_idx) - future.add_callbacks( - callback=self._on_success, callback_args=args, - errback=self._on_error, errback_args=args) - except Exception as exc: - self._put_result(exc, p_idx, False) - finally: - self._executing = False + self._first_error = None + + def _window(self): + # Outstanding work that counts against `concurrency`. List mode + # consumes every completion, so only in-flight requests count. + return self._in_flight + + def _submit(self, concurrency): + while self._window() < concurrency and not self._exhausted \ + and not (self._fail_fast and self._first_error is not None): + try: + statement, params = next(self._statements) + except StopIteration: + self._exhausted = True + return + idx = self._submitted + self._submitted += 1 + self._in_flight += 1 + try: + future = self.session.execute_async(statement, params, timeout=None, execution_profile=self._execution_profile) + future.add_callbacks( + callback=self._on_success, callback_args=(future, idx), + errback=self._on_error, errback_args=(idx,)) + except Exception as exc: + self._on_error(exc, idx) + + def _reap(self, concurrency, block): + # Runs on the calling thread. Each statement is recorded and counted + # exactly once, so a future that reports twice (e.g. a late speculative + # response) can neither advance the window nor fail fast on a result + # that was already superseded. + while True: + try: + idx, result = self._done.get(block) + except Empty: + break + block = False + if idx in self._results: + continue + self._results[idx] = result + self._in_flight -= 1 + if self._fail_fast and not result.success and self._first_error is None: + self._first_error = result.result_or_exc + self._submit(concurrency) def _on_success(self, result, future, idx): future.clear_callbacks() - self._put_result(ResultSet(future, result), idx, True) + self._complete(idx, ExecutionResult(True, ResultSet(future, result))) - def _on_error(self, result, future, idx): - self._put_result(result, idx, False) + def _on_error(self, exc, idx): + self._complete(idx, ExecutionResult(False, exc)) + def _complete(self, idx, result): + # Runs on an IO thread. Only enqueue the completion; the calling thread + # does the bookkeeping (dedup, fail-fast) when it reaps. + self._done.put((idx, result)) -class ConcurrentExecutorGenResults(_ConcurrentExecutor): - - def _put_result(self, result, idx, success): - with self._condition: - heappush(self._results_queue, (idx, ExecutionResult(success, result))) - self._execute_next() - self._condition.notify() - - def _results(self): - with self._condition: - while self._current < self._exec_count: - while not self._results_queue or self._results_queue[0][0] != self._current: - self._condition.wait() - while self._results_queue and self._results_queue[0][0] == self._current: - _, res = heappop(self._results_queue) - try: - self._condition.release() - if self._fail_fast and not res[0]: - raise res[1] - yield res - finally: - self._condition.acquire() - self._current += 1 +class ConcurrentExecutorGenResults(_ConcurrentExecutor): -class ConcurrentExecutorListResults(_ConcurrentExecutor): + def __init__(self, session, statements_and_params, execution_profile): + super(ConcurrentExecutorGenResults, self).__init__(session, statements_and_params, execution_profile) + self._current = 0 - _exception = None + def _window(self): + # Completed-but-unyielded results also occupy the window, so a stalled + # head can't let the whole iterable be pulled into memory. + return self._submitted - self._current def execute(self, concurrency, fail_fast): - self._exception = None - return super(ConcurrentExecutorListResults, self).execute(concurrency, fail_fast) - - def _put_result(self, result, idx, success): - self._results_queue.append((idx, ExecutionResult(success, result))) - with self._condition: + self._fail_fast = fail_fast + self._submit(concurrency) + return self._results_gen(concurrency) + + def _results_gen(self, concurrency): + results = self._results + while self._current < self._submitted: + while self._current not in results: + self._reap(concurrency, block=True) + res = results.pop(self._current) self._current += 1 - if not success and self._fail_fast: - if not self._exception: - self._exception = result - self._condition.notify() - elif not self._execute_next() and self._current == self._exec_count: - self._condition.notify() - - def _results(self): - with self._condition: - while self._current < self._exec_count: - self._condition.wait() - if self._exception and self._fail_fast: - raise self._exception - if self._exception and self._fail_fast: # raise the exception even if there was no wait - raise self._exception - return [r[1] for r in sorted(self._results_queue)] + if self._fail_fast and not res.success: + raise res.result_or_exc + # Keep the window full while the consumer works on this result. + self._reap(concurrency, block=False) + yield res + +class ConcurrentExecutorListResults(_ConcurrentExecutor): + + def execute(self, concurrency, fail_fast): + self._fail_fast = fail_fast + self._submit(concurrency) + while True: + if fail_fast and self._first_error is not None: + raise self._first_error + if not self._in_flight: + break + self._reap(concurrency, block=True) + return [self._results[i] for i in range(self._submitted)] def execute_concurrent_with_args(session, statement, parameters, *args, **kwargs): diff --git a/tests/unit/test_concurrent.py b/tests/unit/test_concurrent.py index d3888aa9de..002a4c4db4 100644 --- a/tests/unit/test_concurrent.py +++ b/tests/unit/test_concurrent.py @@ -24,6 +24,7 @@ import platform import uuid +from cassandra import OperationTimedOut from cassandra.cluster import Cluster, Session from cassandra.concurrent import execute_concurrent, execute_concurrent_with_args from cassandra.pool import Host @@ -307,3 +308,276 @@ def clear_callbacks(self): for success, result in results: assert not success assert result is error + + +class _ManualFuture(object): + """Future completed explicitly by the test, from any thread.""" + _query_trace = None + _col_names = None + _col_types = None + has_more_pages = False + + def __init__(self, params, on_registered): + self.params = params + self._on_registered = on_registered + + def add_callbacks(self, callback, errback, callback_args=(), callback_kwargs=None, + errback_args=(), errback_kwargs=None): + self.callback = lambda rows: callback(rows, *callback_args) + self.errback = lambda exc: errback(exc, *errback_args) + self._on_registered(self) + + def clear_callbacks(self): + pass + + +def _session_with(on_registered): + # on_registered(future) runs once execute_concurrent has attached its callbacks + session = Mock() + session.execute_async.side_effect = lambda stmt, params, **kw: _ManualFuture(params, on_registered) + return session + + +class ConcurrentExecutorTest(unittest.TestCase): + + def _run(self, fn, *args, **kwargs): + # Fail instead of hanging the suite on a deadlock regression. + out = {} + + def target(): + try: + out['result'] = fn(*args, **kwargs) + except BaseException as exc: + out['exc'] = exc + t = threading.Thread(target=target, daemon=True) + t.start() + t.join(10) + assert not t.is_alive(), "execute_concurrent hung" + if 'exc' in out: + raise out['exc'] + return out['result'] + + def test_callbacks_do_not_block_on_submitting_thread(self): + # Completions arriving while the caller is still submitting must not wait for it. + for results_generator in (False, True): + futures = [] + + def on_execute(future): + futures.append(future) + if len(futures) == 2: + t = threading.Thread(target=futures[0].callback, args=(['r'],)) + t.start() + t.join(2) + assert not t.is_alive(), "IO-thread callback blocked on the submitter" + future.callback(['r']) + elif len(futures) > 2: + future.callback(['r']) + return future + + results = self._run(lambda: list(execute_concurrent( + _session_with(on_execute), [("q", (i,)) for i in range(10)], + concurrency=10, results_generator=results_generator))) + assert [r.success for r in results] == [True] * 10 + + def test_no_helper_threads(self): + before = threading.active_count() + + def on_execute(future): + assert threading.active_count() == before + future.callback(['r']) + return future + + results = execute_concurrent(_session_with(on_execute), [("q", (i,)) for i in range(50)], concurrency=5) + assert len(results) == 50 + + def test_submission_happens_on_calling_thread(self): + # Submission must run on the thread that called execute_concurrent, + # never on the IO/callback thread that completes a future. + caller = threading.current_thread() + submitter_threads = [] + + def on_registered(future): + # complete from a foreign thread, like a reactor would + threading.Thread(target=future.callback, args=(['r'],), daemon=True).start() + + def execute_async(stmt, params, **kw): + submitter_threads.append(threading.current_thread()) + return _ManualFuture(params, on_registered) + + session = Mock() + session.execute_async.side_effect = execute_async + + results = list(execute_concurrent(session, [("q", (i,)) for i in range(20)], + concurrency=4, results_generator=True)) + assert len(results) == 20 + assert submitter_threads + assert all(t is caller for t in submitter_threads) + + def test_results_in_order_with_out_of_order_completion(self): + for results_generator in (False, True): + pending = [] + + def on_execute(future): + pending.append(future) + if len(pending) == 4: # complete the window in reverse + while pending: + f = pending.pop() + f.callback([f.params[0]]) + return future + + results = self._run(lambda: list(execute_concurrent( + _session_with(on_execute), [("q", (i,)) for i in range(40)], + concurrency=4, results_generator=results_generator))) + assert [r.result_or_exc.current_rows for r in results] == [[i] for i in range(40)] + + def test_generator_backpressure_waits_for_consumer(self): + # In generator mode new requests are submitted only when the consumer + # asks for the next result, so a slow consumer lowers concurrency. + submitted = [] + pending = [] + + def on_execute(future): + submitted.append(future.params[0]) + pending.append(future) + return future + + gen = execute_concurrent(_session_with(on_execute), + [("q", (i,)) for i in range(1000)], + concurrency=3, results_generator=True) + assert len(submitted) == 3 + + # Complete the initial window while the consumer is idle: nothing more + # may be submitted without demand. + for f in list(pending): + f.callback(['r']) + time.sleep(0.05) + assert len(submitted) == 3 + + # Consuming a result tops the window back up (one new request). + assert next(gen).success + assert len(submitted) == 4 + + def test_generator_window_bounded_when_head_stalls(self): + # If the head request is slow while later ones complete, the completed + # results already occupy the window and must not let the whole iterable + # be submitted and buffered. + submitted = [] + pending = [] + + def on_execute(future): + submitted.append(future.params[0]) + pending.append(future) + return future + + gen = execute_concurrent(_session_with(on_execute), + [("q", (i,)) for i in range(1000)], + concurrency=3, results_generator=True) + assert len(submitted) == 3 + pending[1].callback(['r']) + pending[2].callback(['r']) + + out = [] + t = threading.Thread(target=lambda: out.append(next(gen)), daemon=True) + t.start() + t.join(0.2) + assert not out + assert len(submitted) == 3, "submitted past the window while the head stalled" + + pending[0].callback(['r']) + t.join(5) + assert not t.is_alive() + assert out and out[0].success + + def test_duplicate_completion_counted_once(self): + # e.g. a speculative response arriving after a client timeout + futures = [] + + def on_execute(future): + futures.append(future) + if future.params[0] == 0: + future.errback(OperationTimedOut()) + future.callback(['late']) + return future + + out = [] + t = threading.Thread(target=lambda: out.append(execute_concurrent( + _session_with(on_execute), [("q", (i,)) for i in range(2)], raise_on_first_error=False)), + daemon=True) + t.start() + t.join(0.5) + assert not out, "returned before the second request completed" + futures[1].callback(['r']) + t.join(5) + assert not t.is_alive(), "execute_concurrent hung" + assert [r.success for r in out[0]] == [False, True] + + def test_late_error_after_success_is_ignored(self): + # A duplicate (late) error for a statement whose retained result was a + # success must not be recorded and must not fail-fast, in either mode. + def on_execute(future): + future.callback(['ok']) + if future.params[0] == 0: + future.errback(OperationTimedOut()) + return future + + for results_generator in (False, True): + results = self._run(lambda: list(execute_concurrent( + _session_with(on_execute), [("q", (i,)) for i in range(3)], + concurrency=1, raise_on_first_error=True, + results_generator=results_generator))) + assert [r.success for r in results] == [True, True, True] + + def test_iterable_errors_propagate(self): + def broken(exc): + for i in range(5): + yield ("q", (i,)) + raise exc + + def on_execute(future): + future.callback(['r']) + return future + + for results_generator in (False, True): + for exc in (ValueError("boom"), GeneratorExit()): + with pytest.raises(type(exc)): + self._run(lambda: list(execute_concurrent( + _session_with(on_execute), broken(exc), concurrency=2, + raise_on_first_error=False, results_generator=results_generator))) + + def test_fail_fast_stops_consuming_input(self): + consumed = [] + + def statements(): + for i in range(20000): + consumed.append(i) + yield ("q", (i,)) + + def on_execute(future): + if future.params[0] == 0: + future.errback(ValueError("first")) + else: + future.callback(['r']) + return future + + for results_generator in (False, True): + del consumed[:] + with pytest.raises(ValueError, match="first"): + self._run(lambda: list(execute_concurrent( + _session_with(on_execute), statements(), concurrency=5, + raise_on_first_error=True, results_generator=results_generator))) + assert len(consumed) <= 5 + + def test_execute_async_raising_is_recorded(self): + session = Mock() + session.execute_async.side_effect = RuntimeError("no hosts") + results = execute_concurrent(session, [("q", ())] * 3, raise_on_first_error=False) + assert [(r.success, type(r.result_or_exc)) for r in results] == [(False, RuntimeError)] * 3 + + def test_invalid_concurrency_rejected(self): + session = _session_with(lambda future: future.callback(['r'])) + for bad in (1.5, float('inf'), float('nan'), "10"): + with pytest.raises(TypeError): + execute_concurrent(session, [("q", ())], concurrency=bad) + for bad in (0, -1): + with pytest.raises(ValueError): + execute_concurrent(session, [("q", ())], concurrency=bad)