如何基于非阻塞FastAPI WebSocket实现多进程任务处理
核心问题定位
你的代码存在4个致命错误,导致无法满足非阻塞运行、可终止进程的要求:
run_in_executor调用语法错误:你直接传入update_cicada(data, s, queue),等价于在主事件循环里同步执行完该函数、再把返回值传给执行器,任务根本没有被投递到进程池,必然阻塞WebSocket。正确用法是将函数对象作为第一个参数,后续依次传入函数入参,且必须保存返回的Future句柄用于后续取消操作。- 跨进程传递对象不兼容:
asyncio.Queue是事件循环线程专属对象,不支持跨进程pickle序列化,子进程根本无法写入该队列,主进程永远拿不到执行结果。同时你在主进程实例化的CICADA对象跨进程传递时会丢失运行状态,无法正常工作。 - 消息监听逻辑串行阻塞:你按顺序写
await websocket.receive_json()再await queue.get(),两个操作是串行等待的,在等待队列返回结果的阶段,WebSocket完全无法接收新消息(包括stop终止指令),根本做不到随时终止运行中的进程。 - 资源泄漏风险:你将
ProcessPoolExecutor定义在WebSocket接口内部,每次客户端连接都会新建一个进程池,连接断开时没有手动shutdown回收资源,会产生大量僵尸进程占用CPU、内存。
正确实现方案
必须遵守的配置规则
- 全局初始化ProcessPoolExecutor,不要在接口函数内创建,服务启动时初始化、服务关闭时统一回收。
- 跨进程通信使用
multiprocessing.Queue,跨进程终止信号使用multiprocessing.Event,禁止跨进程传递asyncio队列、事件循环对象、数据库连接、打开的文件句柄等不可序列化资源。 - CPU密集任务依赖的CICADA实例在子进程内部初始化,不要在主进程实例化后跨进程传递。
- 并行监听WebSocket消息和任务结果队列,不要串行await导致阻塞。
- 读
multiprocessing.Queue的阻塞操作要丢到默认线程池执行,避免卡住主事件循环。
可直接运行的修正代码
import asyncio import time import multiprocessing from concurrent.futures import ProcessPoolExecutor from fastapi import FastAPI, WebSocket app = FastAPI() # 全局初始化进程池,max_workers根据CPU核心数调整 MAX_WORKERS = multiprocessing.cpu_count() # Windows环境将mp_context替换为multiprocessing.get_context("spawn") process_pool = ProcessPoolExecutor( max_workers=MAX_WORKERS, mp_context=multiprocessing.get_context("fork") ) # 注意:CPU密集任务函数必须定义在模块顶层,保证可序列化 def update_cicada(data, cicada_init_args, mp_queue: multiprocessing.Queue, stop_event: multiprocessing.Event): # 子进程内部初始化CICADA实例,避免跨进程传对象丢失状态 clip_class = cicada_init_args[0] s = CICADA(clip_class, mp_queue) while not stop_event.is_set(): # 执行CPU密集计算逻辑 calc_result = s.run_step(data) # 结果写入跨进程队列 mp_queue.put(calc_result) # 按需加短休眠,避免占满CPU核心 time.sleep(0.05) @app.websocket_route("/ws") async def websocket_endpoint(websocket: WebSocket): await websocket.accept() loop = asyncio.get_running_loop() # 跨进程通信组件 mp_queue = multiprocessing.Queue() stop_event = multiprocessing.Event() # 主进程协程间通信用异步队列 async_queue = asyncio.Queue() # 保存运行中进程任务的句柄 running_task = None # 协程:持续消费跨进程队列的数据,转存到异步队列 async def consume_mp_queue(): while not stop_event.is_set(): try: # mp_queue.get是阻塞操作,丢到默认线程池执行避免卡事件循环 result = await loop.run_in_executor(None, mp_queue.get, True, 0.1) await async_queue.put(result) except Exception: continue consumer_task = asyncio.create_task(consume_mp_queue()) try: while True: # 并行等待两个事件:收到WebSocket消息、队列有新结果,不会阻塞 recv_msg_task = asyncio.create_task(websocket.receive_json()) get_result_task = asyncio.create_task(async_queue.get()) done, pending = await asyncio.wait( [recv_msg_task, get_result_task], return_when=asyncio.FIRST_COMPLETED ) for finished_task in done: # 处理收到的WebSocket消息 if finished_task is recv_msg_task: data = finished_task.result() # 收到停止指令 if data.get("action") == "stop": stop_event.set() if running_task and not running_task.done(): running_task.cancel() await running_task await websocket.send_json({"status": "stopped"}) return # 收到启动/更新任务指令 else: # 如果已有运行中的任务,先终止再启动新任务 if running_task and not running_task.done(): stop_event.set() await running_task stop_event.clear() # 提交任务到进程池,注意函数和参数分开传递 running_task = loop.run_in_executor( process_pool, update_cicada, data, (clip_class,), mp_queue, stop_event ) # 处理队列返回的计算结果,推给客户端 else: result = finished_task.result() await websocket.send_json(result) # 清理未完成的pending任务,避免泄漏 for pending_task in pending: pending_task.cancel() finally: # 连接断开时强制清理所有资源 stop_event.set() if running_task and not running_task.done(): running_task.cancel() consumer_task.cancel() mp_queue.close() mp_queue.join_thread() await websocket.close() # 服务关闭时回收进程池资源 @app.on_event("shutdown") async def shutdown_pool(): process_pool.shutdown(wait=True)
注意:Windows环境下spawn模式对序列化要求更严格,所有传给进程池的函数、类必须定义在模块顶层,不能嵌套在接口函数内部,否则会抛出序列化错误。
关键逻辑说明
- 终止进程的核心是
multiprocessing.Event,子进程轮询该事件状态,收到信号后主动退出,比直接强制杀进程更安全,不会留下僵尸进程。 - 用
asyncio.wait同时监听消息和结果两个源,保证WebSocket任何时候都能响应客户端的stop指令,不会被计算任务阻塞。 - 跨进程队列只传可序列化的纯数据(字典、字符串、数字等),不要传任何对象实例,避免序列化失败或状态丢失。
内容的提问来源于stack exchange,提问作者user18697242
相关产品推荐
相关产品推荐

