如何在Asyncio Queue中捕获异常并提前抛出首个异常?
问题:Asyncio队列任务遇首个异常立即终止的实现问题
问题背景
使用asyncio固定大小队列处理大量可变任务,无异常时运行正常,但需要实现首个异常出现时立即终止程序的逻辑。当前代码会静默跳过异常,尝试修改后异常场景能触发终止,但无异常时程序会永久挂起。
原代码
import asyncio import threading from typing import Awaitable, Callable, List import aiohttp import aiostream def async_wrap_iter(it): """Wrap blocking iterator into an asynchronous one. Source: https://stackoverflow.com/a/62297994/7619676 """ loop = asyncio.get_event_loop() q = asyncio.Queue(1) exception = None _END = object() async def yield_queue_items(): while True: next_item = await q.get() if next_item is _END: break yield next_item if exception is not None: # the iterator has raised, propagate the exception raise exception def iter_to_queue(): nonlocal exception try: for item in it: # This runs outside the event loop thread, so we # must use thread-safe API to talk to the queue. asyncio.run_coroutine_threadsafe(q.put(item), loop).result() except Exception as e: exception = e finally: asyncio.run_coroutine_threadsafe(q.put(_END), loop).result() threading.Thread(target=iter_to_queue).start() return yield_queue_items() async def main( rows, func: Callable[[List], Awaitable[None]], batch_size: int = 20, max_workers: int = 50, ) -> List: """Adapted from https://stackoverflow.com/a/62404509/7619676""" queue = asyncio.Queue(max_workers) results = [] async def worker(func, queue, results): while True: batch = await queue.get() try: results.append(await func(batch)) except Exception as e: raise e finally: queue.task_done() # create `max_workers` workers and feed them tasks. workers = [ asyncio.create_task(worker(func, queue, results)) for _ in range(max_workers) ] # Feed the database rows to the workers. # The fixed-capacity of the queue ensures that we never hold all rows in memory at the same time. # When the queue reaches full capacity, this will block until a worker dequeues an item. rows = async_wrap_iter(rows) async with aiostream.stream.chunks(rows, batch_size).stream() as chunks: async for batch in chunks: await queue.put(batch) # enqueue a batch of `batch_size` rows await queue.join() for worker in workers: worker.cancel() return results async def func_that_errors_on_evens(batch): i = batch[0] print(i) if i % 2 == 0: raise Exception("fake") return i rows = [1, 2, 3, 4] asyncio.run(main(rows=rows, func=func_that_errors_on_evens, batch_size=1, max_workers=2))
尝试的修改及问题
替换await queue.join()为以下代码后,异常场景能触发终止,但无异常时程序永久挂起:
done, _ = await asyncio.wait( [queue.join(), *workers], return_when=asyncio.FIRST_EXCEPTION ) # alternatively, use asyncio.ALL_COMPLETED to raise "late" consumers_raised = set(done) & set(workers) if consumers_raised: await consumers_raised.pop() # propagate the exception
解决方法
核心问题是:无异常时,queue.join()会完成,但worker是无限循环(while True),永远不会结束,导致asyncio.wait一直等待worker完成。需通过添加终止信号+调整等待逻辑解决。
修改后的完整代码
import asyncio import threading from typing import Awaitable, Callable, List import aiostream def async_wrap_iter(it): """Wrap blocking iterator into an asynchronous one.""" loop = asyncio.get_event_loop() q = asyncio.Queue(1) exception = None _END = object() async def yield_queue_items(): while True: next_item = await q.get() if next_item is _END: break yield next_item if exception is not None: raise exception def iter_to_queue(): nonlocal exception try: for item in it: asyncio.run_coroutine_threadsafe(q.put(item), loop).result() except Exception as e: exception = e finally: asyncio.run_coroutine_threadsafe(q.put(_END), loop).result() threading.Thread(target=iter_to_queue).start() return yield_queue_items() async def main( rows, func: Callable[[List], Awaitable[None]], batch_size: int = 20, max_workers: int = 50, ) -> List: queue = asyncio.Queue(max_workers) results = [] # 定义worker终止标记 _TERMINATE = object() async def worker(func, queue, results): while True: batch = await queue.get() try: if batch is _TERMINATE: break # 收到终止信号,退出循环 results.append(await func(batch)) except Exception as e: # 抛出异常前取消所有worker,快速终止程序 for task in workers: task.cancel() raise e finally: queue.task_done() workers = [ asyncio.create_task(worker(func, queue, results)) for _ in range(max_workers) ] # 喂入任务,捕获取消信号提前终止 rows = async_wrap_iter(rows) try: async with aiostream.stream.chunks(rows, batch_size).stream() as chunks: async for batch in chunks: await queue.put(batch) except asyncio.CancelledError: pass # 给每个worker发送终止信号 for _ in range(max_workers): await queue.put(_TERMINATE) # 等待首个异常或所有worker完成 done, pending = await asyncio.wait( workers, return_when=asyncio.FIRST_EXCEPTION ) # 处理异常:抛出首个异常并清理剩余任务 for task in done: if task.exception() is not None: for t in pending: t.cancel() raise task.exception() # 清理已取消的任务 await asyncio.gather(*pending, return_exceptions=True) return results async def func_that_errors_on_evens(batch): i = batch[0] print(i) if i % 2 == 0: raise Exception("fake") return i rows = [1, 2, 3, 4] asyncio.run(main(rows=rows, func=func_that_errors_on_evens, batch_size=1, max_workers=2))
修改说明
- 添加
_TERMINATE终止标记:任务入队完成后给每个worker发送该标记,让worker退出无限循环,解决无异常时的挂起问题 - worker异常处理:捕获到异常时立即取消所有其他worker,确保程序快速终止
- 调整等待逻辑:直接等待worker任务,用
FIRST_EXCEPTION模式,一旦有worker抛出异常就立即处理 - 任务喂入逻辑:捕获
CancelledError提前终止任务入队,避免不必要的资源消耗
内容的提问来源于stack exchange,提问作者ZaxR
相关产品推荐
相关产品推荐

