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

如何在FastAPI WebSocket中并发运行两个阻塞函数?

解决方案:FastAPI WebSocket中并发运行阻塞函数并实现可控停止

核心思路是用线程包装阻塞函数,避免阻塞asyncio事件循环,同时用线程安全的信号和数据结构实现任务停止与数据共享。

步骤1:改造阻塞函数,支持停止信号与数据传递

示例1:用线程安全队列替代磁盘(推荐,避免磁盘IO同步问题)

# foo所在模块
import time
def foo(url, stop_event, data_queue):
    while not stop_event.is_set():
        # 原有业务逻辑:处理url生成数据
        raw_data = f"[{time.strftime('%H:%M:%S')}] Processed from {url}"
        # 将数据放入线程安全队列
        data_queue.put(raw_data)
        # 模拟原有逻辑的延迟(比如IO操作耗时)
        time.sleep(1)
    print("foo task stopped")

# bar所在模块
import queue
import time
def bar(stop_event, data_queue):
    while not stop_event.is_set():
        try:
            # 从队列取数据,超时等待避免永久阻塞
            raw_data = data_queue.get(timeout=0.5)
            # 原有业务逻辑:处理数据
            processed_data = f"Bar processed: {raw_data}"
            yield processed_data
        except queue.Empty:
            continue
    print("bar task stopped")

示例2:保留磁盘读写逻辑

如果必须使用磁盘存储,修改函数如下:

# foo所在模块
import time
import os
def foo(url, stop_event, client_id):
    file_path = f"{client_id}_temp_data.txt"
    while not stop_event.is_set():
        raw_data = f"[{time.strftime('%H:%M:%S')}] Processed from {url}"
        # 写入磁盘
        with open(file_path, "w", encoding="utf-8") as f:
            f.write(raw_data)
        time.sleep(1)
    # 清理临时文件
    if os.path.exists(file_path):
        os.remove(file_path)
    print("foo task stopped")

# bar所在模块
import time
import os
def bar(stop_event, client_id):
    file_path = f"{client_id}_temp_data.txt"
    while not stop_event.is_set():
        if os.path.exists(file_path):
            with open(file_path, "r", encoding="utf-8") as f:
                raw_data = f.read()
            processed_data = f"Bar processed: {raw_data}"
            yield processed_data
        # 定期扫描磁盘
        time.sleep(0.5)
    print("bar task stopped")

步骤2:实现FastAPI WebSocket端点

队列版完整代码

import asyncio
import queue
import threading
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
import json

app = FastAPI()

# 存储每个客户端的任务资源:client_id -> (stop_event, data_queue, foo_thread, bar_thread)
CLIENT_RESOURCES = {}

async def run_bar(websocket: WebSocket, stop_event: threading.Event, data_queue: queue.Queue):
    """包装bar函数,在后台线程运行并处理WebSocket消息发送"""
    def bar_thread_func():
        for processed_data in bar(stop_event, data_queue):
            if stop_event.is_set():
                break
            # 将发送任务提交到asyncio事件循环(线程中不能直接await)
            asyncio.run_coroutine_threadsafe(
                websocket.send_text(processed_data),
                asyncio.get_event_loop()
            )
    
    bar_thread = threading.Thread(target=bar_thread_func, daemon=True)
    bar_thread.start()
    return bar_thread

@app.websocket("/live-transcription")
async def websocket_endpoint(websocket: WebSocket):
    await websocket.accept()
    client_id = id(websocket)

    try:
        while True:
            try:
                # 超时等待客户端消息,避免永久阻塞
                msg_raw = await asyncio.wait_for(websocket.receive_text(), timeout=1)
                message = json.loads(msg_raw)
                command = message.get("command")
                url = message.get("url", "")

                if command == "STOP":
                    # 停止并清理当前客户端的任务
                    if client_id in CLIENT_RESOURCES:
                        stop_event, data_queue, foo_thread, bar_thread = CLIENT_RESOURCES.pop(client_id)
                        stop_event.set()
                        # 等待线程结束(超时2秒避免僵死)
                        if foo_thread.is_alive():
                            foo_thread.join(timeout=2)
                        if bar_thread.is_alive():
                            bar_thread.join(timeout=2)
                    await websocket.send_text("All tasks stopped")

                elif command == "START":
                    # 先清理之前的任务(如果存在)
                    if client_id in CLIENT_RESOURCES:
                        stop_event, data_queue, foo_thread, bar_thread = CLIENT_RESOURCES.pop(client_id)
                        stop_event.set()
                        if foo_thread.is_alive():
                            foo_thread.join(timeout=2)
                        if bar_thread.is_alive():
                            bar_thread.join(timeout=2)
                    
                    # 创建新的任务资源
                    stop_event = threading.Event()
                    data_queue = queue.Queue(maxsize=10)  # 限制队列大小,防止内存溢出
                    # 启动foo线程
                    foo_thread = threading.Thread(
                        target=foo,
                        args=(url, stop_event, data_queue),
                        daemon=True
                    )
                    foo_thread.start()
                    # 启动bar线程
                    bar_thread = await run_bar(websocket, stop_event, data_queue)
                    # 保存资源
                    CLIENT_RESOURCES[client_id] = (stop_event, data_queue, foo_thread, bar_thread)
                    await websocket.send_text("Tasks started successfully")

            except asyncio.TimeoutError:
                # 超时后检查客户端是否已断开
                if websocket.client_state.value == 2:
                    break
                continue

    except WebSocketDisconnect:
        print(f"Client {client_id} disconnected")
    finally:
        # 最终清理所有资源
        if client_id in CLIENT_RESOURCES:
            stop_event, data_queue, foo_thread, bar_thread = CLIENT_RESOURCES.pop(client_id)
            stop_event.set()
            if foo_thread.is_alive():
                foo_thread.join(timeout=2)
            if bar_thread.is_alive():
                bar_thread.join(timeout=2)
        await websocket.close()

磁盘版代码调整

只需将run_bar和foo_thread的参数改为client_id即可:

# 调整run_bar函数
async def run_bar(websocket: WebSocket, stop_event: threading.Event, client_id: int):
    def bar_thread_func():
        for processed_data in bar(stop_event, client_id):
            if stop_event.is_set():
                break
            asyncio.run_coroutine_threadsafe(
                websocket.send_text(processed_data),
                asyncio.get_event_loop()
            )
    
    bar_thread = threading.Thread(target=bar_thread_func, daemon=True)
    bar_thread.start()
    return bar_thread

# 调整START分支的代码
foo_thread = threading.Thread(
    target=foo,
    args=(url, stop_event, client_id),
    daemon=True
)
foo_thread.start()
bar_thread = await run_bar(websocket, stop_event, client_id)
CLIENT_RESOURCES[client_id] = (stop_event, foo_thread, bar_thread)

问题解决说明

  1. 共享数据为空问题:

    • 队列版使用queue.Queue(线程安全容器),foo生产数据放入队列,bar从队列消费,确保数据能正确传递。
    • 磁盘版通过客户端唯一ID生成专属临时文件,避免多客户端数据冲突,bar定期扫描该文件获取数据。
  2. 任务无法停止问题:

    • 使用threading.Event作为停止信号,foo和bar的循环中每次都会检查stop_event.is_set(),一旦触发就退出循环。
    • 收到STOP命令或客户端断开时,立即设置停止事件并等待线程结束,确保任务能及时终止。
  3. 按需停止任务问题:

    • 用CLIENT_RESOURCES字典存储每个客户端的任务资源,实现任务的隔离管理。
    • 收到START命令时先清理该客户端之前的任务,再创建新任务;收到STOP或断开时直接取出对应资源停止任务,完全按需控制。

内容的提问来源于stack exchange,提问作者Lynob

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 21:55:57