fix(avatar): align knowledge upload limit with production

This commit is contained in:
stefanfeng
2026-09-04 13:56:29 +08:00
parent 28553aba15
commit 6b7201e890
4 changed files with 41 additions and 9 deletions
@@ -17,7 +17,8 @@ UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "upl
os.makedirs(UPLOAD_DIR, exist_ok=True)
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 20 * 1024 * 1024
MAX_UPLOAD_BYTES = 100 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 1024 * 1024
class QAIn(BaseModel):
@@ -80,17 +81,25 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}"
path = os.path.join(avatar_dir, stored)
content = await file.read()
if len(content) > MAX_UPLOAD_BYTES:
return fail("文件不能超过 20MB", code=400)
with open(path, "wb") as f:
f.write(content)
file_size = 0
try:
# Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
with open(path, "wb") as f:
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
file_size += len(chunk)
if file_size > MAX_UPLOAD_BYTES:
raise ValueError("文件不能超过 100MB")
f.write(chunk)
except ValueError as exc:
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=len(content),
file_size=file_size,
file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing",
)
@@ -64,6 +64,29 @@ def test_upload_returns_before_background_vectorization(
db.close()
def test_upload_rejects_oversize_file_before_queuing_indexing(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.MAX_UPLOAD_BYTES", 4),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("oversize.md", b"12345", "text/markdown")},
)
payload = response.json()
assert payload["code"] == 400
assert payload["message"] == "文件不能超过 100MB"
enqueue.assert_not_called()
assert not list((tmp_path / context["avatar"].id).glob("*"))
def test_background_vectorizer_commits_ready_document_and_chunks_together(
tmp_path: Path,
authorization_context,