From e71267cf860250e9ca4aee1f3b974a9c287373ac Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Mon, 7 Sep 2026 17:45:22 +0800 Subject: [PATCH] fix(avatar): upload knowledge files in chunks --- .../backend/routers/knowledge.py | 211 ++++++++++++++++-- .../backend/tests/test_knowledge_storage.py | 78 +++++++ digital-avatar-app/src/api/index.ts | 68 +++++- 3 files changed, 336 insertions(+), 21 deletions(-) diff --git a/digital-avatar-app/backend/routers/knowledge.py b/digital-avatar-app/backend/routers/knowledge.py index 02dc68d..6388d06 100644 --- a/digital-avatar-app/backend/routers/knowledge.py +++ b/digital-avatar-app/backend/routers/knowledge.py @@ -1,4 +1,7 @@ import os +import json +import shutil +import time import uuid from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException @@ -19,6 +22,9 @@ os.makedirs(UPLOAD_DIR, exist_ok=True) ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"} MAX_UPLOAD_BYTES = 50 * 1024 * 1024 UPLOAD_CHUNK_BYTES = 1024 * 1024 +MULTIPART_CHUNK_BYTES = 5 * 1024 * 1024 +MULTIPART_ROOT = ".multipart" +MULTIPART_TTL_SECONDS = 24 * 60 * 60 class QAIn(BaseModel): @@ -31,6 +37,74 @@ class EnabledIn(BaseModel): enabled: bool = True +class MultipartUploadIn(BaseModel): + filename: str + fileSize: int + totalChunks: int + + +def _validate_document(filename: str, file_size: int): + ext = os.path.splitext(filename or "")[1].lower() + if ext not in ALLOWED_EXT: + return None, f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx" + if file_size <= 0: + return None, "文件内容不能为空" + if file_size > MAX_UPLOAD_BYTES: + return None, "文件不能超过 50MB" + return ext, "" + + +def _multipart_dir(avatar_id: str, upload_id: str) -> str: + safe_avatar_id = os.path.basename(avatar_id) + safe_upload_id = os.path.basename(upload_id) + if ( + safe_avatar_id != avatar_id + or safe_upload_id != upload_id + or len(upload_id) != 32 + or any(character not in "0123456789abcdef" for character in upload_id) + ): + raise HTTPException(status_code=400, detail="上传标识无效") + return os.path.join(UPLOAD_DIR, MULTIPART_ROOT, safe_avatar_id, safe_upload_id) + + +def _purge_stale_multipart_uploads(avatar_id: str): + avatar_upload_root = os.path.join(UPLOAD_DIR, MULTIPART_ROOT, os.path.basename(avatar_id)) + if not os.path.isdir(avatar_upload_root): + return + cutoff = time.time() - MULTIPART_TTL_SECONDS + for entry in os.scandir(avatar_upload_root): + if entry.is_dir(follow_symlinks=False) and entry.stat(follow_symlinks=False).st_mtime < cutoff: + shutil.rmtree(entry.path, ignore_errors=True) + + +def _read_multipart_metadata(avatar_id: str, upload_id: str) -> tuple[str, dict]: + upload_dir = _multipart_dir(avatar_id, upload_id) + metadata_path = os.path.join(upload_dir, "metadata.json") + if not os.path.isfile(metadata_path): + raise HTTPException(status_code=404, detail="上传任务不存在或已过期") + with open(metadata_path, "r", encoding="utf-8") as stream: + return upload_dir, json.load(stream) + + +def _create_knowledge_doc(db: Session, avatar_id: str, filename: str, ext: str, file_size: int, stored: str): + doc = KnowledgeDoc( + id=uuid.uuid4().hex, + avatar_id=avatar_id, + filename=filename, + file_type=ext.lstrip("."), + file_size=file_size, + file_url=f"/api/files/{avatar_id}/{stored}", + status="parsing", + index_stage="queued", + index_progress=0, + ) + db.add(doc) + db.commit() + db.refresh(doc) + knowledge_vectorizer.enqueue(doc.id) + return doc + + def _doc_payload(doc: KnowledgeDoc) -> dict: payload = doc.to_dict() stored_name = os.path.basename(doc.file_url or "") @@ -74,9 +148,9 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D @router.post("/avatar/{avatar_id}/knowledge/docs") async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)): _require_owned_avatar(db, avatar_id, authorization) - ext = os.path.splitext(file.filename or "")[1].lower() - if ext not in ALLOWED_EXT: - return fail(f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx", code=400) + ext, validation_error = _validate_document(file.filename or "", 1) + if validation_error: + return fail(validation_error, code=400) avatar_dir = os.path.join(UPLOAD_DIR, avatar_id) os.makedirs(avatar_dir, exist_ok=True) stored = f"{uuid.uuid4().hex}{ext}" @@ -94,25 +168,124 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization if os.path.exists(path): os.remove(path) return fail(str(exc), code=400) - doc = KnowledgeDoc( - id=uuid.uuid4().hex, - avatar_id=avatar_id, - filename=file.filename, - file_type=ext.lstrip("."), - file_size=file_size, - file_url=f"/api/files/{avatar_id}/{stored}", - status="parsing", - index_stage="queued", - index_progress=0, + if file_size == 0: + if os.path.exists(path): + os.remove(path) + return fail("文件内容不能为空", code=400) + + doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored) + return ok(_doc_payload(doc)) + + +@router.post("/avatar/{avatar_id}/knowledge/uploads") +def create_multipart_upload( + avatar_id: str, + body: MultipartUploadIn, + authorization: str = Header(None), + db: Session = Depends(get_db), +): + _require_owned_avatar(db, avatar_id, authorization) + ext, validation_error = _validate_document(body.filename, body.fileSize) + if validation_error: + return fail(validation_error, code=400) + expected_chunks = (body.fileSize + MULTIPART_CHUNK_BYTES - 1) // MULTIPART_CHUNK_BYTES + if body.totalChunks != expected_chunks: + return fail("文件分片数量不正确", code=400) + + _purge_stale_multipart_uploads(avatar_id) + upload_id = uuid.uuid4().hex + upload_dir = _multipart_dir(avatar_id, upload_id) + os.makedirs(upload_dir, exist_ok=False) + metadata = { + "filename": body.filename, + "fileSize": body.fileSize, + "totalChunks": body.totalChunks, + "extension": ext, + } + with open(os.path.join(upload_dir, "metadata.json"), "w", encoding="utf-8") as stream: + json.dump(metadata, stream, ensure_ascii=False) + return ok({"uploadId": upload_id, "chunkSize": MULTIPART_CHUNK_BYTES}) + + +@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/chunks/{chunk_index}") +async def upload_multipart_chunk( + avatar_id: str, + upload_id: str, + chunk_index: int, + file: UploadFile = File(...), + authorization: str = Header(None), + db: Session = Depends(get_db), +): + _require_owned_avatar(db, avatar_id, authorization) + upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id) + total_chunks = int(metadata["totalChunks"]) + if chunk_index < 0 or chunk_index >= total_chunks: + return fail("文件分片序号不正确", code=400) + + expected_size = min( + MULTIPART_CHUNK_BYTES, + int(metadata["fileSize"]) - chunk_index * MULTIPART_CHUNK_BYTES, ) + part_path = os.path.join(upload_dir, f"{chunk_index}.part") + temporary_path = f"{part_path}.uploading" + received = 0 + try: + with open(temporary_path, "wb") as stream: + while chunk := await file.read(UPLOAD_CHUNK_BYTES): + received += len(chunk) + if received > expected_size: + raise ValueError("文件分片大小不正确") + stream.write(chunk) + if received != expected_size: + raise ValueError("文件分片大小不正确") + os.replace(temporary_path, part_path) + except ValueError as exc: + if os.path.exists(temporary_path): + os.remove(temporary_path) + return fail(str(exc), code=400) + return ok({"chunkIndex": chunk_index, "uploadedBytes": received}) - # Persist and acknowledge the upload first. Extraction and embeddings may take - # minutes for a PDF and must never consume the browser request timeout. - db.add(doc) - db.commit() - db.refresh(doc) - knowledge_vectorizer.enqueue(doc.id) +@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/complete") +def complete_multipart_upload( + avatar_id: str, + upload_id: str, + authorization: str = Header(None), + db: Session = Depends(get_db), +): + _require_owned_avatar(db, avatar_id, authorization) + upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id) + total_chunks = int(metadata["totalChunks"]) + part_paths = [os.path.join(upload_dir, f"{index}.part") for index in range(total_chunks)] + if not all(os.path.isfile(path) for path in part_paths): + return fail("文件分片尚未上传完整", code=400) + if sum(os.path.getsize(path) for path in part_paths) != int(metadata["fileSize"]): + return fail("文件分片总大小不正确", code=400) + + avatar_dir = os.path.join(UPLOAD_DIR, avatar_id) + os.makedirs(avatar_dir, exist_ok=True) + stored = f"{uuid.uuid4().hex}{metadata['extension']}" + final_path = os.path.join(avatar_dir, stored) + temporary_path = f"{final_path}.assembling" + try: + with open(temporary_path, "wb") as output: + for part_path in part_paths: + with open(part_path, "rb") as source: + shutil.copyfileobj(source, output, UPLOAD_CHUNK_BYTES) + os.replace(temporary_path, final_path) + doc = _create_knowledge_doc( + db, + avatar_id, + metadata["filename"], + metadata["extension"], + int(metadata["fileSize"]), + stored, + ) + except Exception: + if os.path.exists(temporary_path): + os.remove(temporary_path) + raise + shutil.rmtree(upload_dir, ignore_errors=True) return ok(_doc_payload(doc)) diff --git a/digital-avatar-app/backend/tests/test_knowledge_storage.py b/digital-avatar-app/backend/tests/test_knowledge_storage.py index ea361f0..63a155f 100644 --- a/digital-avatar-app/backend/tests/test_knowledge_storage.py +++ b/digital-avatar-app/backend/tests/test_knowledge_storage.py @@ -87,6 +87,84 @@ def test_upload_rejects_oversize_file_before_queuing_indexing( assert not list((tmp_path / context["avatar"].id).glob("*")) +def test_multipart_upload_reassembles_file_before_queuing_indexing( + tmp_path: Path, + authorization_context, +): + context = authorization_context + avatar_id = context["avatar"].id + content = b"0123456789" + with ( + patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)), + patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4), + patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue, + ): + created = client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads", + headers=context["owner_headers"], + json={"filename": "large.pdf", "fileSize": len(content), "totalChunks": 3}, + ).json()["data"] + + for index, chunk in enumerate((content[:4], content[4:8], content[8:])): + response = client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/{index}", + headers=context["owner_headers"], + files={"file": (f"chunk-{index}", chunk, "application/octet-stream")}, + ) + assert response.json()["code"] == 200 + + completed = client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete", + headers=context["owner_headers"], + ).json()["data"] + + assert completed["status"] == "parsing" + assert completed["fileSize"] == len(content) + enqueue.assert_called_once_with(completed["id"]) + stored_path = tmp_path / avatar_id / Path(completed["fileUrl"]).name + assert stored_path.read_bytes() == content + assert not (tmp_path / ".multipart" / avatar_id / created["uploadId"]).exists() + + db = SessionLocal() + try: + stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == completed["id"]).one() + db.delete(stored) + db.commit() + finally: + db.close() + + +def test_multipart_upload_rejects_incomplete_parts( + tmp_path: Path, + authorization_context, +): + context = authorization_context + avatar_id = context["avatar"].id + with ( + patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)), + patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4), + patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue, + ): + created = client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads", + headers=context["owner_headers"], + json={"filename": "large.pdf", "fileSize": 6, "totalChunks": 2}, + ).json()["data"] + client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/0", + headers=context["owner_headers"], + files={"file": ("chunk-0", b"0123", "application/octet-stream")}, + ) + response = client.post( + f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete", + headers=context["owner_headers"], + ) + + assert response.json()["code"] == 400 + assert response.json()["message"] == "文件分片尚未上传完整" + enqueue.assert_not_called() + + def test_background_vectorizer_commits_ready_document_and_chunks_together( tmp_path: Path, authorization_context, diff --git a/digital-avatar-app/src/api/index.ts b/digital-avatar-app/src/api/index.ts index 8467333..03708c8 100644 --- a/digital-avatar-app/src/api/index.ts +++ b/digital-avatar-app/src/api/index.ts @@ -333,12 +333,76 @@ export interface SearchResult { export const getKnowledgeDocs = (avatarId: string) => request.get(`/avatar/${avatarId}/knowledge/docs`) -// 上传文档(支持 md/txt/pdf/doc/docx/xlsx) -export const uploadKnowledgeDoc = ( +const KNOWLEDGE_UPLOAD_CHUNK_SIZE = 5 * 1024 * 1024 + +const uploadKnowledgeChunk = async ( + avatarId: string, + uploadId: string, + chunkIndex: number, + chunk: Blob, + onProgress?: (loaded: number) => void +) => { + const form = new FormData() + form.append('file', chunk, `chunk-${chunkIndex}`) + let reportedLoaded = 0 + for (let attempt = 1; attempt <= 3; attempt += 1) { + try { + await request.post( + `/avatar/${avatarId}/knowledge/uploads/${uploadId}/chunks/${chunkIndex}`, + form, + { + headers: { 'Content-Type': 'multipart/form-data' }, + timeout: 2 * 60 * 1000, + onUploadProgress: (event) => { + reportedLoaded = Math.max(reportedLoaded, Math.min(event.loaded, chunk.size)) + onProgress?.(reportedLoaded) + } + } + ) + return + } catch (error: any) { + const status = Number(error?.response?.status || 0) + const retryable = !status || status === 408 || status === 429 || status >= 500 + if (!retryable || attempt === 3) throw error + await new Promise((resolve) => window.setTimeout(resolve, attempt * 800)) + } + } +} + +// 大文件拆成 5MB 分片,避免生产代理的请求体限制拦截整个文件。 +export const uploadKnowledgeDoc = async ( avatarId: string, file: File, onUploadProgress?: (loaded: number, total: number) => void ) => { + if (file.size > KNOWLEDGE_UPLOAD_CHUNK_SIZE) { + const totalChunks = Math.ceil(file.size / KNOWLEDGE_UPLOAD_CHUNK_SIZE) + const upload: any = await request.post(`/avatar/${avatarId}/knowledge/uploads`, { + filename: file.name, + fileSize: file.size, + totalChunks + }) + let uploadedBytes = 0 + for (let index = 0; index < totalChunks; index += 1) { + const start = index * KNOWLEDGE_UPLOAD_CHUNK_SIZE + const chunk = file.slice(start, Math.min(start + KNOWLEDGE_UPLOAD_CHUNK_SIZE, file.size)) + await uploadKnowledgeChunk( + avatarId, + upload.uploadId, + index, + chunk, + (chunkLoaded) => onUploadProgress?.(uploadedBytes + chunkLoaded, file.size) + ) + uploadedBytes += chunk.size + onUploadProgress?.(uploadedBytes, file.size) + } + return request.post( + `/avatar/${avatarId}/knowledge/uploads/${upload.uploadId}/complete`, + undefined, + { timeout: 2 * 60 * 1000 } + ) + } + const form = new FormData() form.append('file', file) return request.post(`/avatar/${avatarId}/knowledge/docs`, form, { -- 2.54.0