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) };
问题核心
- 协议冲突:
/file_upload是HTTP POST端点,无法处理WebSocket的握手请求(WebSocket使用ws/wss协议,而非http/https) - 参数传递矛盾:前端无法在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); }
关键说明
- 上传ID生成:在WebSocket连接建立时生成唯一ID,确保每个上传会话对应一个WebSocket连接
- 连接管理:使用字典存储连接,连接断开或上传完成后及时清理,避免内存泄漏
- 错误处理:上传失败时通过WebSocket通知前端,并清理无效连接
- 生产环境优化:建议用Redis等分布式存储替代全局字典,支持多进程/多实例部署
内容的提问来源于stack exchange,提问作者Hannon qaoud
相关产品推荐
相关产品推荐

