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

FastAPI中如何将WebSocket与带参端点结合实现上传进度通知

FastAPI文件上传结合WebSocket实时推送进度的实现问题

问题背景

需要在FastAPI的/file_upload端点中结合WebSocket实现文件上传进度的实时推送,但直接将WebSocket作为POST端点的参数会导致连接失败——前端无法通过WebSocket协议连接到HTTP POST端点,且FastAPI的POST端点无法处理WebSocket握手请求。

现有错误代码

main.py

from fastapi import FastAPI, UploadFile, Depends, HTTPException, status, WebSocket
from typing import Annotated, Tuple
import logging
import utils

app = FastAPI()

@app.post("/file_upload")
async def file_upload(
      data: UploadFile, name: str, user: Annotated[User, Depends(get_authenticated_user)], ws: WebSocket
):
    """ docs for swagger ... """
    try:
        await ws.accept()  # 错误:POST端点无法处理WebSocket握手
        path, size = await utils.save_file(data, ws)
    except Exception as e:
        logging.error("Could not create file: %s", e)
        raise HTTPException(
            status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
            detail="There was an error uploading the file",
        )

    return {"file_path": path}

utils.save_file()

import os
import json
import aiofiles
from fastapi import WebSocket
from typing import Tuple
from your_module import DATA_DIR, gen_random_filename

async def save_file(data: UploadFile, ws:WebSocket = None) -> Tuple[str, int]:
    """save a zip file."""
    local_filepath = os.path.join(DATA_DIR, f"data_{gen_random_filename()}.zip")
    total_size = data.size
    size = 0
    CHUNK = 64 * 1024 * 1024  # 64MB
    try:
        async with aiofiles.open(local_filepath, "wb") as f:
            while contents := await data.read(CHUNK):
                await f.write(contents)
                size += len(contents)
                print(f'{size} / {total_size}')
                progress = {"progress": int(size / total_size * 100)}
                await ws.send_text(json.dumps(data))  # 错误:应发送progress而非原始data
    except Exception as e:
        logging.error("Could not save file: %s", e)
        raise Exception("Could not write data to the local file")
    finally:
        await data.close()

    return local_filepath, size

index.js

// 错误:尝试连接HTTP POST端点,协议不匹配
var ws = new WebSocket("ws://localhost:8000/file_upload");
ws.onmessage = function(event) {
    console.log(event.data)
};

问题核心

  1. 协议冲突:/file_upload是HTTP POST端点,无法处理WebSocket的握手请求(WebSocket使用ws/wss协议,而非http/https)
  2. 参数传递矛盾:前端无法在WebSocket连接时同时传递UploadFile、认证用户等HTTP请求参数

解决方案:拆分端点+上传ID关联

将文件上传(HTTP POST)和进度推送(WebSocket)拆分为两个独立端点,通过唯一上传ID关联两者:

修改后的main.py

from fastapi import FastAPI, UploadFile, Depends, HTTPException, status, WebSocket
from typing import Annotated, Dict
import logging
import utils
import uuid
import json

app = FastAPI()

# 存储上传ID与WebSocket连接的映射(生产环境建议用Redis等分布式存储)
active_uploads: Dict[str, WebSocket] = {}

@app.websocket("/upload_progress")
async def websocket_upload_progress(websocket: WebSocket):
    await websocket.accept()
    upload_id = str(uuid.uuid4())
    # 将连接存入映射,主动向前端发送上传ID
    await websocket.send_text(json.dumps({"upload_id": upload_id}))
    active_uploads[upload_id] = websocket
    try:
        # 保持连接直到上传完成或断开
        while True:
            await websocket.receive_text()
    except:
        # 连接断开时移除映射
        if upload_id in active_uploads:
            del active_uploads[upload_id]

@app.post("/file_upload")
async def file_upload(
      data: UploadFile, name: str, user: Annotated[User, Depends(get_authenticated_user)], upload_id: str
):
    """ docs for swagger ... """
    # 验证上传ID对应的WebSocket连接存在
    ws = active_uploads.get(upload_id)
    if not ws:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Invalid or expired upload ID"
        )
    try:
        path, size = await utils.save_file(data, upload_id)
        # 上传完成通知前端
        await ws.send_text(json.dumps({"progress": 100, "status": "completed"}))
    except Exception as e:
        logging.error("Could not create file: %s", e)
        # 上传失败时通知前端
        await ws.send_text(json.dumps({"progress": -1, "error": str(e)}))
        del active_uploads[upload_id]
        raise HTTPException(
            status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
            detail="There was an error uploading the file",
        )
    # 上传完成后移除连接
    del active_uploads[upload_id]
    return {"file_path": path}

修改后的utils.save_file()

import os
import json
import aiofiles
from typing import Tuple
from your_module import DATA_DIR, gen_random_filename
from main import active_uploads

async def save_file(data: UploadFile, upload_id: str) -> Tuple[str, int]:
    """save a zip file with progress push via WebSocket."""
    local_filepath = os.path.join(DATA_DIR, f"data_{gen_random_filename()}.zip")
    total_size = data.size
    size = 0
    CHUNK = 64 * 1024 * 1024  # 64MB
    ws = active_uploads.get(upload_id)
    try:
        async with aiofiles.open(local_filepath, "wb") as f:
            while contents := await data.read(CHUNK):
                await f.write(contents)
                size += len(contents)
                progress = {"progress": int(size / total_size * 100)}
                if ws:
                    await ws.send_text(json.dumps(progress))
    except Exception as e:
        logging.error("Could not save file: %s", e)
        raise Exception("Could not write data to the local file")
    finally:
        await data.close()

    return local_filepath, size

修改后的index.js

let uploadId;
// 先连接WebSocket获取上传ID
const ws = new WebSocket("ws://localhost:8000/upload_progress");
ws.onmessage = function(event) {
    const data = JSON.parse(event.data);
    if (data.upload_id) {
        uploadId = data.upload_id;
        console.log("获取上传ID:", uploadId);
    } else if (data.progress === -1) {
        console.error("上传失败:", data.error);
    } else if (data.progress === 100) {
        console.log("上传完成");
    } else {
        console.log("上传进度:", data.progress + "%");
        // 此处可更新页面进度条DOM
    }
};

// 文件上传逻辑示例
async function uploadFile(file) {
    if (!uploadId) {
        alert("请等待WebSocket连接建立");
        return;
    }
    const formData = new FormData();
    formData.append("data", file);
    formData.append("name", file.name);
    formData.append("upload_id", uploadId);

    // 携带认证信息(示例:Bearer token)
    const response = await fetch("http://localhost:8000/file_upload", {
        method: "POST",
        headers: {
            "Authorization": "Bearer " + localStorage.getItem("auth_token")
        },
        body: formData
    });

    const result = await response.json();
    console.log("上传结果:", result);
}

关键说明

  1. 上传ID生成:在WebSocket连接建立时生成唯一ID,确保每个上传会话对应一个WebSocket连接
  2. 连接管理:使用字典存储连接,连接断开或上传完成后及时清理,避免内存泄漏
  3. 错误处理:上传失败时通过WebSocket通知前端,并清理无效连接
  4. 生产环境优化:建议用Redis等分布式存储替代全局字典,支持多进程/多实例部署

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 22:47:16