Merge pull request 'fix(avatar): 分片上传大文件知识库' (#18) from codex/avatar-chunk-upload-20260907 into main

Reviewed-on: #18
This commit was merged in pull request #18.
This commit is contained in:
2026-09-07 17:48:15 +08:00
3 changed files with 336 additions and 21 deletions
+192 -19
View File
@@ -1,4 +1,7 @@
import os import os
import json
import shutil
import time
import uuid import uuid
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException 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"} ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 50 * 1024 * 1024 MAX_UPLOAD_BYTES = 50 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 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): class QAIn(BaseModel):
@@ -31,6 +37,74 @@ class EnabledIn(BaseModel):
enabled: bool = True 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: def _doc_payload(doc: KnowledgeDoc) -> dict:
payload = doc.to_dict() payload = doc.to_dict()
stored_name = os.path.basename(doc.file_url or "") 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") @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)): 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) _require_owned_avatar(db, avatar_id, authorization)
ext = os.path.splitext(file.filename or "")[1].lower() ext, validation_error = _validate_document(file.filename or "", 1)
if ext not in ALLOWED_EXT: if validation_error:
return fail(f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx", code=400) return fail(validation_error, code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id) avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True) os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}" 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): if os.path.exists(path):
os.remove(path) os.remove(path)
return fail(str(exc), code=400) return fail(str(exc), code=400)
doc = KnowledgeDoc( if file_size == 0:
id=uuid.uuid4().hex, if os.path.exists(path):
avatar_id=avatar_id, os.remove(path)
filename=file.filename, return fail("文件内容不能为空", code=400)
file_type=ext.lstrip("."),
file_size=file_size, doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored)
file_url=f"/api/files/{avatar_id}/{stored}", return ok(_doc_payload(doc))
status="parsing",
index_stage="queued",
index_progress=0, @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)) return ok(_doc_payload(doc))
@@ -87,6 +87,84 @@ def test_upload_rejects_oversize_file_before_queuing_indexing(
assert not list((tmp_path / context["avatar"].id).glob("*")) 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( def test_background_vectorizer_commits_ready_document_and_chunks_together(
tmp_path: Path, tmp_path: Path,
authorization_context, authorization_context,
+66 -2
View File
@@ -333,12 +333,76 @@ export interface SearchResult {
export const getKnowledgeDocs = (avatarId: string) => export const getKnowledgeDocs = (avatarId: string) =>
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`) request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
// 上传文档(支持 md/txt/pdf/doc/docx/xlsx) const KNOWLEDGE_UPLOAD_CHUNK_SIZE = 5 * 1024 * 1024
export const uploadKnowledgeDoc = (
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, avatarId: string,
file: File, file: File,
onUploadProgress?: (loaded: number, total: number) => void 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<KnowledgeDoc>(
`/avatar/${avatarId}/knowledge/uploads/${upload.uploadId}/complete`,
undefined,
{ timeout: 2 * 60 * 1000 }
)
}
const form = new FormData() const form = new FormData()
form.append('file', file) form.append('file', file)
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, { return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {