diff --git a/api.py b/api.py index 949be41..1e2cfe9 100644 --- a/api.py +++ b/api.py @@ -103,6 +103,10 @@ async def put_chunk(transfer_id: str, idx: int, request: Request): async def complete(transfer_id: str): if tasks.get(transfer_id) is None: return JSONResponse(status_code=404, content={"error": "任务不存在"}) + # 幂等: 已完成的任务直接返回已有 file_id (客户端超时重试/断线重连场景) + done_rec = db.get_file_by_transfer(transfer_id) + if done_rec is not None: + return {"status": "done", "file_id": done_rec["file_id"]} try: tasks.set_status(transfer_id, "assembling") merged = assembler.assemble(transfer_id) diff --git a/db.py b/db.py index 05ac569..3dc94c2 100644 --- a/db.py +++ b/db.py @@ -177,6 +177,15 @@ class ServerDB: ) return file_id + def get_file_by_transfer(self, transfer_id: str) -> dict[str, Any] | None: + """按任务查已入库文件 (complete 幂等: 已完成直接返回)""" + with self._conn() as conn: + row = conn.execute( + "SELECT file_id, file_name, path, size FROM files WHERE transfer_id = ?", + (transfer_id,), + ).fetchone() + return dict(row) if row else None + def list_files(self) -> list[dict[str, Any]]: """全部文件记录 (ls 端点)""" with self._conn() as conn: diff --git a/tests/test_api.py b/tests/test_api.py index 0159f7b..436c143 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -175,6 +175,13 @@ class ServerPipelineTest(unittest.TestCase): # 未知 file_id -> 404 self.assertEqual(self.client.get("/api/files/nope").status_code, 404) + # complete 幂等: 重复 complete 返回同一个 file_id, 不新增入库记录 + before = len(self.client.get("/api/files").json()["files"]) + r = self.client.post(f"/api/transfer/{tid}/complete") + self.assertEqual(r.status_code, 200) + self.assertEqual(r.json()["file_id"], file_id) + self.assertEqual(len(self.client.get("/api/files").json()["files"]), before) + if __name__ == "__main__": unittest.main()