diff --git a/src/executorlib/standalone/interactive/communication.py b/src/executorlib/standalone/interactive/communication.py index a9f4ea59..5dd640be 100644 --- a/src/executorlib/standalone/interactive/communication.py +++ b/src/executorlib/standalone/interactive/communication.py @@ -5,6 +5,7 @@ import cloudpickle import zmq +import zmq.asyncio class ExecutorlibSocketError(RuntimeError): @@ -32,7 +33,7 @@ def __init__( log_obj_size (boolean): Enable debug mode which reports the size of the communicated objects. time_out_ms (int): Time out for waiting for a message on socket in milliseconds. """ - self._context = zmq.Context() + self._context = self._new_context() self._socket = self._context.socket(zmq.PAIR) self._poller = zmq.Poller() self._poller.register(self._socket, zmq.POLLIN) @@ -46,6 +47,9 @@ def __init__( self._booted_sucessfully: bool = False self._stop_function: Optional[Callable] = None + def _new_context(self) -> zmq.Context: + return zmq.Context() + @property def status(self) -> bool: return self._booted_sucessfully @@ -179,6 +183,87 @@ def __del__(self): self.shutdown(wait=True) +class AsyncSocketInterface(SocketInterface): + """ + Variant of the SocketInterface based on zmq.asyncio. The communication methods are coroutines which have to be + awaited on the event loop the interface is used from. All interfaces of one executor share a single ZMQ context, + which is owned and terminated by the caller. The communication protocol is identical to the SocketInterface. + + Args: + context (zmq.asyncio.Context): shared asyncio ZMQ context + spawner (executorlib.shared.spawner.BaseSpawner): Interface for starting the parallel process + log_obj_size (boolean): Enable debug mode which reports the size of the communicated objects. + time_out_ms (int): Time out for waiting for a message on socket in milliseconds. + """ + + def __init__( + self, + context: zmq.asyncio.Context, + spawner=None, + log_obj_size: bool = False, + time_out_ms: int = 1000, + ): + self._shared_context = context + super().__init__( + spawner=spawner, log_obj_size=log_obj_size, time_out_ms=time_out_ms + ) + + def _new_context(self) -> zmq.Context: + return self._shared_context + + async def send_dict_async(self, input_dict: dict): + data = cloudpickle.dumps(input_dict) + if self._logger is not None: + self._logger.warning("Send dictionary of size: " + str(sys.getsizeof(data))) + await self._socket.send(data) + + async def receive_dict_async(self) -> dict: + response = 0 + while response == 0: + response = await self._socket.poll(self._time_out_ms) + if not self._spawner.poll(): + return { + "error": ExecutorlibSocketError( + "SocketInterface crashed during execution." + ), + } + data = await self._socket.recv() + if self._logger is not None: + self._logger.warning( + "Received dictionary of size: " + str(sys.getsizeof(data)) + ) + return cloudpickle.loads(data) + + async def send_and_receive_dict_async(self, input_dict: dict) -> dict: + await self.send_dict_async(input_dict=input_dict) + return await self.receive_dict_async() + + async def shutdown_async(self, wait: bool = True): + result = None + if self._spawner.poll(): + output = await self.send_and_receive_dict_async( + input_dict={"shutdown": True, "wait": wait} + ) + if "result" in output: + result = output["result"] + self._spawner.shutdown(wait=wait) + self._reset_socket() + return result + + def _reset_socket(self): + # The shared context is terminated by its owner, only the socket is closed here. + if self._socket is not None: + self._socket.close(linger=0) + self._process = None + self._socket = None + self._context = None + + def __del__(self): + # The socket belongs to the event loop thread, so only the worker process is stopped here. + if self._spawner is not None and self._spawner.poll(): + self._spawner.shutdown(wait=False) + + def interface_bootup( command_lst: list[str], connections, @@ -186,6 +271,7 @@ def interface_bootup( log_obj_size: bool = False, worker_id: Optional[int] = None, stop_function: Optional[Callable] = None, + context: Optional[zmq.asyncio.Context] = None, ) -> SocketInterface: """ Start interface for ZMQ communication @@ -205,6 +291,7 @@ def interface_bootup( worker_id (int): Communicate the worker which ID was assigned to it for future reference and resource distribution. stop_function (Callable): Function to stop the interface. + context (zmq.asyncio.Context): if provided an AsyncSocketInterface using this shared context is returned. Returns: executorlib.shared.communication.SocketInterface: socket interface for zmq communication @@ -218,10 +305,17 @@ def interface_bootup( ] if worker_id is not None: command_lst += ["--worker-id", str(worker_id)] - interface = SocketInterface( - spawner=connections, - log_obj_size=log_obj_size, - ) + if context is not None: + interface: SocketInterface = AsyncSocketInterface( + context=context, + spawner=connections, + log_obj_size=log_obj_size, + ) + else: + interface = SocketInterface( + spawner=connections, + log_obj_size=log_obj_size, + ) command_lst += [ "--zmqport", str(interface.bind_to_random_port()), diff --git a/src/executorlib/task_scheduler/interactive/blockallocation.py b/src/executorlib/task_scheduler/interactive/blockallocation.py index e9ee1fd2..3371b9c9 100644 --- a/src/executorlib/task_scheduler/interactive/blockallocation.py +++ b/src/executorlib/task_scheduler/interactive/blockallocation.py @@ -1,3 +1,4 @@ +import os import queue import random from concurrent.futures import Future @@ -17,6 +18,10 @@ from executorlib.standalone.interactive.spawner import BaseSpawner, MpiExecSpawner from executorlib.standalone.queue import cancel_items_in_queue, put_front from executorlib.task_scheduler.base import TaskSchedulerBase, validate_resource_dict +from executorlib.task_scheduler.interactive.blockallocation_async import ( + AsyncWorker, + AsyncWorkerPool, +) from executorlib.task_scheduler.interactive.shared import ( execute_task_dict, reset_task_dict, @@ -61,6 +66,8 @@ class BlockAllocationTaskScheduler(TaskSchedulerBase): """ + _async_pool: Optional[AsyncWorkerPool] = None + def __init__( self, max_workers: int = 1, @@ -86,14 +93,22 @@ def __init__( self._alive_workers_lock = Lock() self._bootup_events = [Event() for _ in range(self._max_workers)] self._bootup_events[0].set() + # Experimental: drive all workers from one private asyncio event loop instead of one thread per worker. + if os.environ.get("EXECUTORLIB_ASYNCIO", "0").lower() in ("1", "true"): + assert self._future_queue is not None + self._async_pool = AsyncWorkerPool(future_queue=self._future_queue) self._set_process( - process=[ - Thread( - target=_execute_multiple_tasks, - kwargs=self._worker_kwargs(worker_id), - ) - for worker_id in range(self._max_workers) - ], + process=[self._new_worker(worker_id) for worker_id in range(max_workers)], + ) + + def _new_worker(self, worker_id: int): + if self._async_pool is not None: + return AsyncWorker( + pool=self._async_pool, kwargs=self._worker_kwargs(worker_id) + ) + return Thread( + target=_execute_multiple_tasks, + kwargs=self._worker_kwargs(worker_id), ) def _worker_kwargs(self, worker_id: int) -> dict: @@ -137,10 +152,7 @@ def max_workers(self, max_workers: int): with self._alive_workers_lock: self._alive_workers[0] += max_workers - old_max_workers new_process_lst = [ - Thread( - target=_execute_multiple_tasks, - kwargs=self._worker_kwargs(worker_id), - ) + self._new_worker(worker_id) for worker_id in range(old_max_workers, max_workers) ] for process_instance in new_process_lst: @@ -210,10 +222,12 @@ def shutdown(self, wait: bool = True, *, cancel_futures: bool = False): for process in self._process: process.join() self._future_queue.join() + if self._async_pool is not None: + self._async_pool.close(wait=wait) self._process = None self._future_queue = None - def _set_process(self, process: list[Thread]): # type: ignore + def _set_process(self, process: list): # type: ignore """ Set the process for the executor. diff --git a/src/executorlib/task_scheduler/interactive/blockallocation_async.py b/src/executorlib/task_scheduler/interactive/blockallocation_async.py new file mode 100644 index 00000000..c7df46b1 --- /dev/null +++ b/src/executorlib/task_scheduler/interactive/blockallocation_async.py @@ -0,0 +1,283 @@ +""" +Experimental asyncio control plane for the BlockAllocationTaskScheduler. + +Instead of one Python thread per worker, all workers of an executor are driven by coroutines on a single private +asyncio event loop, which runs in one dedicated background thread owned by the executor. The ZMQ communication uses +zmq.asyncio. The task queue stays a queue.Queue; a single dispatcher thread hands its items to idle worker coroutines, +so tasks remain in the queue (and can be cancelled) until a worker is ready, exactly like in the threaded version. + +Resulting administrative threads per executor: 2 (event loop + dispatcher), independent of the number of workers. + +The event loop of the calling thread (e.g. the one of a Jupyter kernel) is never used or modified. +""" + +import asyncio +import queue +import traceback +from threading import Event, Lock, Thread +from typing import Callable, Optional + +import zmq.asyncio + +from executorlib.standalone.command import get_interactive_execute_command +from executorlib.standalone.interactive.communication import ( + AsyncSocketInterface, + ExecutorlibSocketError, + interface_bootup, +) +from executorlib.standalone.interactive.spawner import BaseSpawner, MpiExecSpawner +from executorlib.task_scheduler.interactive.shared import ( + execute_task_dict_async, + reset_task_dict, + task_done, +) + + +class AsyncWorkerPool: + """ + Private asyncio event loop in one background thread, plus one dispatcher thread which blocks on the shared + queue.Queue on behalf of the idle worker coroutines. + + Args: + future_queue (queue.Queue): task queue shared with the executor + """ + + def __init__(self, future_queue: queue.Queue): + self._future_queue = future_queue + self._loop = asyncio.new_event_loop() + self.context = zmq.asyncio.Context() + self._idle_workers: queue.Queue = queue.Queue() + self._lock = Lock() + self._active_workers = 0 + self._closing = False + self._close_requested = False + self._loop_thread = Thread(target=self._run_loop) + self._dispatch_thread = Thread(target=self._dispatch) + self._loop_thread.start() + self._dispatch_thread.start() + + def start_worker(self, worker: "AsyncWorker"): + with self._lock: + self._active_workers += 1 + asyncio.run_coroutine_threadsafe(self._run_worker(worker), self._loop) + + def close(self, wait: bool = True): + """ + Stop the event loop and the dispatcher thread as soon as all worker coroutines have finished. + + Args: + wait (bool): block until both threads are stopped + """ + if not self._close_requested: + self._close_requested = True + self._loop.call_soon_threadsafe(self._request_close) + if wait: + self._loop_thread.join() + self._dispatch_thread.join() + + async def get_task(self) -> dict: + """ + Wait for the next item of the shared queue.Queue without blocking the event loop. + """ + waiter = self._loop.create_future() + self._idle_workers.put(waiter) + return await waiter + + def _dispatch(self): + while True: + waiter = self._idle_workers.get() + if waiter is None: + break + task_dict = self._future_queue.get() + self._loop.call_soon_threadsafe(waiter.set_result, task_dict) + + def _run_loop(self): + asyncio.set_event_loop(self._loop) + self._loop.run_forever() + self.context.destroy(linger=0) + self._loop.close() + + async def _run_worker(self, worker: "AsyncWorker"): + try: + await _execute_multiple_tasks_async(pool=self, **worker.kwargs) + except Exception: + traceback.print_exc() + finally: + worker.done.set() + with self._lock: + self._active_workers -= 1 + self._stop_if_finished() + + def _request_close(self): + self._closing = True + self._stop_if_finished() + + def _stop_if_finished(self): + with self._lock: + finished = self._closing and self._active_workers == 0 + if finished: + self._idle_workers.put(None) + self._loop.stop() + + +class AsyncWorker: + """ + Thread-like handle (start(), join(), is_alive()) for a worker coroutine, so the BlockAllocationTaskScheduler can + manage it like the Thread objects it uses otherwise. + """ + + def __init__(self, pool: AsyncWorkerPool, kwargs: dict): + self._pool = pool + self.kwargs = kwargs + self.done = Event() + + def start(self): + self._pool.start_worker(self) + + def join(self, timeout: Optional[float] = None): + self.done.wait(timeout) + + def is_alive(self) -> bool: + return not self.done.is_set() + + +async def _execute_multiple_tasks_async( + pool: AsyncWorkerPool, + future_queue: queue.Queue, + cores: int = 1, + spawner: type[BaseSpawner] = MpiExecSpawner, + hostname_localhost: Optional[bool] = None, + init_function: Optional[Callable] = None, + cache_directory: Optional[str] = None, + cache_key: Optional[str] = None, + log_obj_size: bool = False, + error_log_file: Optional[str] = None, + worker_id: int = 0, + stop_function: Optional[Callable] = None, + restart_limit: int = 0, + next_bootup_event: Optional[Event] = None, + alive_workers: Optional[list] = None, + alive_workers_lock: Optional[Lock] = None, + bootup_event: Optional[Event] = None, + queue_join_on_shutdown: bool = False, + **kwargs, +) -> None: + """ + Coroutine version of blockallocation._execute_multiple_tasks(), see there for the arguments. + + The worker coroutines are started in worker_id order and boot up their process synchronously before their first + await, so the boot order is preserved without waiting on bootup_event (which would block the event loop). For the + same reason queue_join_on_shutdown is not supported; the BlockAllocationTaskScheduler always sets it to False. + """ + interface = interface_bootup( + command_lst=get_interactive_execute_command( + cores=cores, + ), + connections=spawner(cores=cores, worker_id=worker_id, **kwargs), + hostname_localhost=hostname_localhost, + log_obj_size=log_obj_size, + worker_id=worker_id, + stop_function=stop_function, + context=pool.context, + ) + assert isinstance(interface, AsyncSocketInterface) + if next_bootup_event is not None: + next_bootup_event.set() + interface_initialization_exception = await _set_init_function_async( + interface=interface, + init_function=init_function, + ) + restart_counter = 0 + while True: + if not interface.status and restart_counter >= restart_limit: + await _drain_dead_worker_async( + pool=pool, + future_queue=future_queue, + alive_workers=alive_workers, + alive_workers_lock=alive_workers_lock, + ) + break + elif not interface.status: + interface.bootup() + interface_initialization_exception = await _set_init_function_async( + interface=interface, + init_function=init_function, + ) + restart_counter += 1 + else: # interface.status == True + task_dict = await pool.get_task() + if "shutdown" in task_dict and task_dict["shutdown"]: + if interface.status: + await interface.shutdown_async(wait=task_dict["wait"]) + task_done(future_queue=future_queue) + break + elif "fn" in task_dict and "future" in task_dict: + f = task_dict.pop("future") + if interface_initialization_exception is not None: + f.set_exception(exception=interface_initialization_exception) + else: + # The interface failed during the execution + interface.status = await execute_task_dict_async( + task_dict=task_dict, + future_obj=f, + interface=interface, + cache_directory=cache_directory, + cache_key=cache_key, + error_log_file=error_log_file, + ) + if not interface.status: + reset_task_dict( + future_obj=f, future_queue=future_queue, task_dict=task_dict + ) + task_done(future_queue=future_queue) + + +async def _drain_dead_worker_async( + pool: AsyncWorkerPool, + future_queue: queue.Queue, + alive_workers: Optional[list] = None, + alive_workers_lock: Optional[Lock] = None, +) -> None: + """ + Coroutine version of blockallocation._drain_dead_worker(). + """ + if alive_workers is not None and alive_workers_lock is not None: + with alive_workers_lock: + if alive_workers[0] > 0: + alive_workers[0] -= 1 + while True: + task_dict = await pool.get_task() + if "shutdown" in task_dict and task_dict["shutdown"]: + task_done(future_queue=future_queue) + break + elif "fn" in task_dict and "future" in task_dict: + if alive_workers is not None and alive_workers_lock is not None: + with alive_workers_lock: + has_healthy_workers = alive_workers[0] > 0 + else: + has_healthy_workers = False + if has_healthy_workers: + future_queue.put(task_dict) + task_done(future_queue=future_queue) + # give the healthy workers a chance to pick up the recycled task + await asyncio.sleep(0.01) + else: + f = task_dict.pop("future") + f.set_exception( + ExecutorlibSocketError("SocketInterface crashed during execution.") + ) + task_done(future_queue=future_queue) + + +async def _set_init_function_async( + interface: AsyncSocketInterface, + init_function: Optional[Callable] = None, +) -> Optional[Exception]: + interface_initialization_exception = None + if init_function is not None and interface.status: + output = await interface.send_and_receive_dict_async( + input_dict={"init": True, "fn": init_function, "args": (), "kwargs": {}} + ) + if "error" in output: + interface_initialization_exception = output["error"] + return interface_initialization_exception diff --git a/src/executorlib/task_scheduler/interactive/shared.py b/src/executorlib/task_scheduler/interactive/shared.py index abb633ae..370b2ffe 100644 --- a/src/executorlib/task_scheduler/interactive/shared.py +++ b/src/executorlib/task_scheduler/interactive/shared.py @@ -4,9 +4,10 @@ import time from concurrent.futures import Future from concurrent.futures._base import PENDING -from typing import Optional +from typing import Any, Optional from executorlib.standalone.interactive.communication import ( + AsyncSocketInterface, ExecutorlibSocketError, SocketInterface, ) @@ -57,6 +58,63 @@ def execute_task_dict( return True +async def execute_task_dict_async( + task_dict: dict, + future_obj: Future, + interface: AsyncSocketInterface, + cache_directory: Optional[str] = None, + cache_key: Optional[str] = None, + error_log_file: Optional[str] = None, +) -> bool: + """ + Coroutine version of execute_task_dict() for the AsyncSocketInterface, with identical caching and error semantics. + + Returns: + bool: True if the task was submitted successfully, False otherwise. + """ + if future_obj.done() or not future_obj.set_running_or_notify_cancel(): + return True + if error_log_file is not None: + task_dict["error_log_file"] = error_log_file + file_name: Optional[str] = None + data_dict: dict[str, Any] = {} + if cache_directory is not None: + from executorlib.standalone.hdf import dump, get_cache_files, get_output + + task_key, data_dict, serialize_exception = serialize_funct( + fn=task_dict["fn"], + fn_args=task_dict["args"], + fn_kwargs=task_dict["kwargs"], + resource_dict=task_dict.get("resource_dict", {}), + cache_key=cache_key, + ) + if serialize_exception is not None: + future_obj.set_exception(exception=serialize_exception) + return True + file_name = os.path.abspath(os.path.join(cache_directory, task_key + "_o.h5")) + if file_name in get_cache_files(cache_directory=cache_directory): + _, _, result = get_output(file_name=file_name) + future_obj.set_result(result) + return True + time_start = time.time() + try: + output = await interface.send_and_receive_dict_async(input_dict=task_dict) + except Exception as e: + future_obj.set_exception(exception=e) + return True + if "result" in output: + if file_name is not None: + data_dict["output"] = output["result"] + data_dict["runtime"] = time.time() - time_start + dump(file_name=file_name, data_dict=data_dict) + future_obj.set_result(output["result"]) + elif isinstance(output["error"], ExecutorlibSocketError): + return False + else: + future_obj.set_exception(exception=output["error"]) + return True + + def task_done(future_queue: queue.Queue): """ Mark the current task as done in the current queue. diff --git a/tests/unit/task_scheduler/interactive/test_blockallocation_async.py b/tests/unit/task_scheduler/interactive/test_blockallocation_async.py new file mode 100644 index 00000000..71f7ab45 --- /dev/null +++ b/tests/unit/task_scheduler/interactive/test_blockallocation_async.py @@ -0,0 +1,212 @@ +import asyncio +import os +import threading +import time +import unittest +from concurrent.futures import CancelledError, Future +from unittest.mock import patch + +from executorlib import SingleNodeExecutor +from executorlib.standalone.interactive.spawner import MpiExecSpawner +from executorlib.task_scheduler.interactive.blockallocation import ( + BlockAllocationTaskScheduler, +) +from executorlib.task_scheduler.interactive.blockallocation_async import ( + AsyncWorker, +) + + +def add(a, b): + return a + b + + +def sleep_and_return(x, seconds=0.2): + import time + + time.sleep(seconds) + return x + + +def raise_error(msg): + raise ValueError(msg) + + +def get_k(k): + return k + + +def init_k(): + return {"k": 42} + + +def _scheduler(max_workers=1, **executor_kwargs): + return BlockAllocationTaskScheduler( + max_workers=max_workers, + executor_kwargs=executor_kwargs | {"hostname_localhost": True}, + spawner=MpiExecSpawner, + ) + + +@patch.dict(os.environ, {"EXECUTORLIB_ASYNCIO": "1"}) +class TestAsyncBlockAllocation(unittest.TestCase): + def test_uses_async_workers(self): + with _scheduler(max_workers=2) as exe: + self.assertTrue(all(isinstance(p, AsyncWorker) for p in exe._process)) + self.assertEqual(exe.submit(add, 1, 2).result(), 3) + + def test_returns_concurrent_future(self): + with _scheduler() as exe: + f = exe.submit(add, 1, b=2) + self.assertIsInstance(f, Future) + self.assertEqual(f.result(), 3) + self.assertTrue(f.done()) + + def test_many_simultaneous_tasks(self): + with _scheduler(max_workers=4) as exe: + fs = [exe.submit(sleep_and_return, i, seconds=0.1) for i in range(16)] + self.assertEqual([f.result() for f in fs], list(range(16))) + + def test_parallel_execution(self): + with _scheduler(max_workers=4) as exe: + exe.submit(add, 0, 0).result() # wait for all workers to boot + start = time.time() + fs = [exe.submit(sleep_and_return, i, seconds=1.0) for i in range(4)] + [f.result() for f in fs] + self.assertLess(time.time() - start, 3.0) + + def test_exception(self): + with _scheduler(max_workers=2) as exe: + f = exe.submit(raise_error, "boom") + with self.assertRaises(ValueError): + f.result() + # the worker stays usable after an exception + self.assertEqual(exe.submit(add, 2, 2).result(), 4) + + def test_init_function(self): + with _scheduler(max_workers=2, init_function=init_k) as exe: + self.assertEqual(exe.submit(get_k).result(), 42) + + def test_shutdown_wait(self): + exe = _scheduler(max_workers=2) + fs = [exe.submit(sleep_and_return, i) for i in range(4)] + pool = exe._async_pool + exe.shutdown(wait=True) + self.assertTrue(all(f.done() for f in fs)) + self.assertEqual([f.result() for f in fs], list(range(4))) + self.assertFalse(pool._loop_thread.is_alive()) + self.assertFalse(pool._dispatch_thread.is_alive()) + self.assertTrue(pool._loop.is_closed()) + + def test_shutdown_no_wait(self): + exe = _scheduler(max_workers=1) + f = exe.submit(sleep_and_return, 1) + pool = exe._async_pool + exe.shutdown(wait=False) + self.assertEqual(f.result(), 1) + pool._loop_thread.join(timeout=30) + pool._dispatch_thread.join(timeout=30) + self.assertFalse(pool._loop_thread.is_alive()) + self.assertFalse(pool._dispatch_thread.is_alive()) + + def test_shutdown_cancel_futures(self): + exe = _scheduler(max_workers=1) + fs = [exe.submit(sleep_and_return, i, seconds=0.5) for i in range(6)] + fs[0].result() + exe.shutdown(wait=True, cancel_futures=True) + self.assertTrue(any(f.cancelled() for f in fs)) + for f in fs: + if f.cancelled(): + with self.assertRaises(CancelledError): + f.result() + + def test_repeated_creation(self): + base = threading.active_count() + for i in range(5): + with _scheduler(max_workers=2) as exe: + self.assertEqual(exe.submit(add, i, 1).result(), i + 1) + self.assertEqual(threading.active_count(), base) + + def test_resize(self): + with _scheduler(max_workers=1) as exe: + exe.max_workers = 3 + self.assertEqual(len(exe._process), 3) + fs = [exe.submit(add, i, 1) for i in range(6)] + self.assertEqual([f.result() for f in fs], [i + 1 for i in range(6)]) + exe.max_workers = 1 + self.assertEqual(len(exe._process), 1) + self.assertEqual(exe.submit(add, 1, 1).result(), 2) + + def test_single_node_executor(self): + with SingleNodeExecutor(max_workers=2, block_allocation=True) as exe: + fs = [exe.submit(add, i, i) for i in range(8)] + self.assertEqual([f.result() for f in fs], [2 * i for i in range(8)]) + + +@patch.dict(os.environ, {"EXECUTORLIB_ASYNCIO": "1"}) +class TestAsyncRunningEventLoop(unittest.TestCase): + """ + Simulate Jupyter/IPython, where the calling thread already runs an asyncio event loop. + """ + + def test_sync_api_inside_running_loop(self): + async def existing_application(): + running_loop = asyncio.get_running_loop() + with _scheduler(max_workers=2) as exe: + self.assertIsNot(exe._async_pool._loop, running_loop) + fs = [exe.submit(add, i, 1) for i in range(4)] + result = [f.result() for f in fs] + self.assertIs(asyncio.get_running_loop(), running_loop) + return result + + self.assertEqual(asyncio.run(existing_application()), [1, 2, 3, 4]) + + def test_await_wrapped_future_inside_running_loop(self): + async def existing_application(): + with SingleNodeExecutor(max_workers=2, block_allocation=True) as exe: + fs = [exe.submit(sleep_and_return, i, seconds=0.1) for i in range(4)] + # the running loop stays responsive while executorlib works + ticks = 0 + while not all(f.done() for f in fs): + await asyncio.sleep(0.01) + ticks += 1 + result = await asyncio.gather(*[asyncio.wrap_future(f) for f in fs]) + return result, ticks + + result, ticks = asyncio.run(existing_application()) + self.assertEqual(result, [0, 1, 2, 3]) + self.assertGreater(ticks, 0) + + def test_exception_inside_running_loop(self): + async def existing_application(): + with _scheduler() as exe: + with self.assertRaises(ValueError): + await asyncio.wrap_future(exe.submit(raise_error, "boom")) + + asyncio.run(existing_application()) + + +class TestAsyncThreadScaling(unittest.TestCase): + @staticmethod + def _admin_threads(max_workers): + base = threading.active_count() + with _scheduler(max_workers=max_workers) as exe: + fs = [exe.submit(add, i, 1) for i in range(max_workers)] + [f.result() for f in fs] + during = threading.active_count() - base + return during + + def test_thread_count_constant(self): + with patch.dict(os.environ, {"EXECUTORLIB_ASYNCIO": "1"}): + counts = {n: self._admin_threads(n) for n in (1, 4, 16)} + # event loop thread + dispatcher thread, independent of max_workers + self.assertEqual(counts[1], counts[16], counts) + self.assertLessEqual(counts[16], 3, counts) + + def test_thread_count_threaded_reference(self): + with patch.dict(os.environ, {"EXECUTORLIB_ASYNCIO": "0"}): + counts = {n: self._admin_threads(n) for n in (1, 4, 16)} + self.assertEqual(counts[16] - counts[1], 15, counts) + + +if __name__ == "__main__": + unittest.main()