You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何基于非阻塞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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 16:21:20