Files

453 lines
19 KiB
Python

"""
API 层 (api)
模块: 服务端 / API 层
输入: HTTP 请求 (带认证)
输出: 响应 / 路由到各模块
零知识 + 零合并设计: 服务端只存卷/传卷, 不合并不解密, 不持有密钥。
端点:
POST /api/transfer/init {init json} -> {transfer_id}
GET /api/transfer/{id}/chunks -> {received: [n...]}
PUT /api/transfer/{id}/chunk/{n} 卷密文 -> {ok} / 409
POST /api/transfer/{id}/complete -> {status, file_id} / 409 (秒回, 不合并)
GET /api/files -> {files: [...]}
GET /api/files/{id}/chunk/{n} -> 单卷密文 (头 X-Enc-Params / X-Chunk-Count)
全部端点需 Authorization: Bearer <token> (与客户端预共享)
"""
import asyncio
import base64
import json
import logging
import shutil
import subprocess
from pathlib import Path
from fastapi import Depends, FastAPI, Header, HTTPException, Request
from fastapi.responses import FileResponse, JSONResponse
from auth import hash_password, new_token, verify_password
# 业务日志 (随 uvicorn 输出到 stderr -> /tmp/sz-server.log)
log = logging.getLogger("sz")
logging.basicConfig(level=logging.INFO, format="%(asctime)s [业务] %(message)s")
from db import ServerDB
from receiver import Receiver, ReceiverError
from settings import QUOTA_BYTES, STORAGE_ROOT, TMP_ROOT, TOKEN
from storage import Storage
from task_manager import TaskManager
# ---------- 依赖装配 (单例) ----------
db = ServerDB()
tasks = TaskManager(db)
receiver = Receiver(tasks, TMP_ROOT)
storage = Storage(tasks, STORAGE_ROOT)
app = FastAPI(title="7z-encrypt 服务端 (零知识 + 零合并)")
# ---------- 认证 ----------
def verify_token(authorization: str | None = Header(default=None)) -> str:
"""Bearer token -> username; 失败 401.
SZ_TOKEN (系统 token) = admin, 兼容旧配置; 注册用户走 users 表 token"""
if authorization == f"Bearer {TOKEN}":
return "admin"
if authorization and authorization.startswith("Bearer "):
user = db.get_username_by_token(authorization[7:])
if user is not None:
return user
raise HTTPException(status_code=401, detail="未授权: token 无效或缺失")
def _check_owner(rec: dict, username: str) -> None:
"""文件归属校验: 每个账户只操作自己的空间 (admin 只认 admin/无主旧文件)"""
if username == "admin":
if rec.get("user_id") not in ("admin", None):
raise HTTPException(status_code=403, detail="无权访问他人文件")
return
if rec.get("user_id") != username:
raise HTTPException(status_code=403, detail="无权访问他人文件")
# ---------- 账号: 注册 / 登录 / 注销 ----------
@app.post("/api/auth/register")
async def register(request: Request):
"""注册 {username, password} -> 自动登录 {username, token}"""
try:
body = await request.json()
username = str(body["username"]).strip()
password = str(body["password"])
except (KeyError, TypeError, ValueError):
return JSONResponse(status_code=400, content={"error": "缺少 username/password"})
if not 3 <= len(username) <= 32 or not username.replace("_", "").isalnum():
return JSONResponse(status_code=400, content={"error": "用户名需 3-32 位字母数字或下划线"})
if not 6 <= len(password) <= 128:
return JSONResponse(status_code=400, content={"error": "密码需 6-128 位"})
if not db.create_user(username, hash_password(password)):
return JSONResponse(status_code=409, content={"error": "用户名已存在"})
token = new_token()
db.set_token(username, token)
log.info("register %s 用户=%s", request.client.host if request.client else "?", username)
return {"username": username, "token": token}
@app.post("/api/auth/login")
async def login(request: Request):
"""登录 {username, password} -> {username, token}"""
try:
body = await request.json()
username = str(body["username"]).strip()
password = str(body["password"])
except (KeyError, TypeError, ValueError):
return JSONResponse(status_code=400, content={"error": "缺少 username/password"})
user = db.get_user(username)
if user is None or not verify_password(password, user["password_hash"]):
return JSONResponse(status_code=401, content={"error": "用户名或密码错误"})
token = new_token()
db.set_token(username, token)
log.info("login %s 用户=%s", request.client.host if request.client else "?", username)
return {"username": username, "token": token}
@app.get("/api/auth/me", dependencies=[Depends(verify_token)])
async def me(username: str = Depends(verify_token)):
"""当前登录用户 (token 有效性验证)"""
return {"username": username}
@app.post("/api/auth/logout", dependencies=[Depends(verify_token)])
async def logout(username: str = Depends(verify_token)):
"""注销: 使当前 token 失效"""
if username != "admin":
db.clear_token_by_username(username)
return {"ok": True}
def _enc_headers(rec: dict) -> dict[str, str]:
"""下载响应头: 解密参数 + 卷数"""
ep = json.dumps(rec["enc_params"], ensure_ascii=False)
return {
"X-Enc-Params": base64.b64encode(ep.encode()).decode(),
"X-Chunk-Count": str(rec["chunk_count"]),
}
# ---------- 文件: ls / 逐卷下载 ----------
@app.get("/api/files")
async def list_files(request: Request, username: str = Depends(verify_token)):
"""文件列表 (ls)"""
files = db.list_files(username)
log.info("ls %s %s 个文件", request.client.host if request.client else "?", len(files))
return {"files": files}
@app.get("/api/quota")
async def quota(request: Request, username: str = Depends(verify_token)):
"""用户空间配额: 已用 (files size 总和) / 配额 / 剩余"""
used = db.sum_files_size(username)
remain = max(QUOTA_BYTES - used, 0)
log.info("quota %s 已用 %.1fMB/%.1fMB",
request.client.host if request.client else "?", used / 1048576, QUOTA_BYTES / 1048576)
return {
"used_bytes": used,
"quota_bytes": QUOTA_BYTES,
"remain_bytes": remain,
"percent": round(used * 100 / QUOTA_BYTES, 1) if QUOTA_BYTES else 0,
}
# ---------- 目录 ----------
@app.post("/api/dirs")
async def create_dir(request: Request, username: str = Depends(verify_token)):
"""创建目录 (按用户隔离, 重名 409)"""
body = await request.json()
name = str(body.get("name", "")).strip()
if not name or "/" in name or "\\" in name:
return JSONResponse(status_code=400, content={"error": "目录名非法 (不能含 / 或 \\\\)"})
if len(name) > 64:
return JSONResponse(status_code=400, content={"error": "目录名过长 (≤64)"})
dir_id = db.create_dir(name, username)
if dir_id is None:
return JSONResponse(status_code=409, content={"error": f"目录已存在: {name}"})
log.info("mkdir %s %s (dir_id=%s)", request.client.host if request.client else "?", name, dir_id)
return {"ok": True, "dir_id": dir_id, "name": name}
@app.get("/api/dirs")
async def list_dirs(request: Request, username: str = Depends(verify_token)):
"""目录列表"""
dirs = db.list_dirs(username)
return {"dirs": dirs}
@app.delete("/api/dirs/{dir_id}")
async def delete_dir(dir_id: int, request: Request, username: str = Depends(verify_token)):
"""删除目录 (归属校验, 不存在/跨账户 404)"""
rec = db.get_dir(dir_id)
if rec is None or rec["user_id"] != username:
return JSONResponse(status_code=404, content={"error": "目录不存在"})
db.delete_dir(dir_id, username)
log.info("rmdir %s dir_id=%s (%s)", request.client.host if request.client else "?", dir_id, rec["name"])
return {"ok": True, "dir_id": dir_id}
@app.delete("/api/files/{file_id}")
async def delete_file(file_id: str, request: Request, username: str = Depends(verify_token)):
"""删除文件 (卷目录 + 入库记录)"""
rec = db.get_file(file_id)
if rec is None:
return JSONResponse(status_code=404, content={"error": "文件不存在"})
_check_owner(rec, username)
try:
p = Path(rec["path"])
if p.is_dir():
shutil.rmtree(p, ignore_errors=True)
elif p.is_file():
p.unlink(missing_ok=True)
db.delete_file(file_id)
log.info("delete %s %s", request.client.host if request.client else "?", file_id)
except OSError as e:
return JSONResponse(status_code=409, content={"error": f"删除失败: {e}"})
return {"ok": True, "file_id": file_id}
def _resolve_dir(dir_id: int | None, username: str) -> int | None:
"""目录归属校验: 0/None=根目录, 否则必须属于该用户 (跨账户 403)"""
if dir_id in (None, 0):
return None
d = db.get_dir(dir_id)
if d is None or d["user_id"] != username:
raise HTTPException(status_code=403, detail="目录不存在或无权使用")
return dir_id
@app.put("/api/files/{file_id}/move")
async def move_file(file_id: str, request: Request, username: str = Depends(verify_token)):
"""移动文件到目录 (body {dir_id}; 0/缺省=根目录)"""
rec = db.get_file(file_id)
if rec is None:
return JSONResponse(status_code=404, content={"error": "文件不存在"})
_check_owner(rec, username)
body = await request.json()
try:
dir_id = _resolve_dir(body.get("dir_id"), username)
except HTTPException as e:
return JSONResponse(status_code=e.status_code, content={"error": e.detail})
db.move_file(file_id, dir_id)
log.info("move %s %s -> dir=%s", request.client.host if request.client else "?", file_id, dir_id)
return {"ok": True, "file_id": file_id, "dir_id": dir_id}
@app.post("/api/files/{file_id}/clone")
async def clone_file(file_id: str, request: Request, username: str = Depends(verify_token)):
"""克隆文件 (物理复制卷目录 + 新记录), body {dir_id} 可选"""
rec = db.get_file(file_id)
if rec is None:
return JSONResponse(status_code=404, content={"error": "文件不存在"})
_check_owner(rec, username)
body = await request.json()
try:
dir_id = _resolve_dir(body.get("dir_id"), username)
except HTTPException as e:
return JSONResponse(status_code=e.status_code, content={"error": e.detail})
src_dir = Path(rec["path"])
new_dir = src_dir.with_name(src_dir.name + "_clone")
try:
shutil.copytree(src_dir, new_dir)
except OSError as e:
return JSONResponse(status_code=409, content={"error": f"克隆失败: {e}"})
new_id = db.clone_file(file_id, dir_id, str(new_dir))
if not new_id:
shutil.rmtree(new_dir, ignore_errors=True)
return JSONResponse(status_code=404, content={"error": "源文件记录丢失"})
log.info("clone %s %s -> %s", request.client.host if request.client else "?", file_id, new_id)
return {"ok": True, "file_id": new_id}
@app.get("/api/files/{file_id}/chunk/{idx}")
async def download_chunk(file_id: str, idx: int, request: Request, username: str = Depends(verify_token)):
"""下载单卷密文 (零合并: 服务端不拼接, 客户端逐卷拉取本地合并)"""
rec = db.get_file(file_id)
if rec is None:
return JSONResponse(status_code=404, content={"error": "文件不存在"})
_check_owner(rec, username)
if not 1 <= idx <= rec["chunk_count"]:
return JSONResponse(status_code=404, content={"error": f"卷号越界: {idx}"})
chunk_path = Path(rec["path"]) / f"chunk_{idx:04d}"
if not chunk_path.exists():
return JSONResponse(status_code=404, content={"error": f"卷文件缺失: {chunk_path}"})
log.info("download %s %s%d/%d", request.client.host if request.client else "?",
file_id, idx, rec["chunk_count"])
return FileResponse(
chunk_path,
media_type="application/octet-stream",
headers=_enc_headers(rec),
)
# ---------- 传输 ----------
@app.post("/api/transfer/init")
async def init_transfer(request: Request, username: str = Depends(verify_token)):
# 防超大 json: init 元数据不该超过 1MB (chunk spec 每条约 120B, 1MB 可容纳 ~8000 卷)
length = request.headers.get("Content-Length")
if length and int(length) > 1 << 20:
return JSONResponse(status_code=413, content={"error": "init json 过大"})
try:
init_json = await request.json()
dir_id = init_json.pop("dir_id", None)
if dir_id is not None:
d = db.get_dir(int(dir_id))
if d is None or d["user_id"] != username:
return JSONResponse(status_code=403, content={"error": "目录不存在或无权使用"})
transfer_id = tasks.create(init_json, username, dir_id)
log.info("init %s %s卷 任务=%s", request.client.host if request.client else "?",
init_json.get("chunk_count"), transfer_id)
except (KeyError, TypeError, ValueError) as e:
return JSONResponse(status_code=400, content={"error": f"init json 非法: {e}"})
return {"transfer_id": transfer_id}
@app.get("/api/transfer/{transfer_id}/chunks")
async def get_chunks(transfer_id: str, username: str = Depends(verify_token)):
if tasks.get(transfer_id) is None:
return JSONResponse(status_code=404, content={"error": "任务不存在"})
return {"received": sorted(tasks.received(transfer_id))}
@app.put("/api/transfer/{transfer_id}/chunk/{idx}")
async def put_chunk(transfer_id: str, idx: int, request: Request, username: str = Depends(verify_token)):
if tasks.get(transfer_id) is None:
return JSONResponse(status_code=404, content={"error": "任务不存在"})
data = await request.body()
try:
receiver.receive(transfer_id, idx, data)
log.info("chunk %s %s%d (%dB)", request.client.host if request.client else "?",
transfer_id, idx, len(data))
except ReceiverError as e:
return JSONResponse(status_code=409, content={"error": str(e)})
return {"ok": True}
@app.post("/api/transfer/{transfer_id}/pull")
async def pull_chunks(transfer_id: str, request: Request, username: str = Depends(verify_token)):
"""aria2 反向拉卷: 客户端起临时 HTTP 服务, 服务端 aria2c 并发拉取 (打满带宽)
body: {base_url, token} —— 客户端临时服务地址 + 一次性拉取令牌
卷文件拉取后逐卷走 receiver.receive 校验 (大小+SHA-256) 落盘。
aria2 子进程走 asyncio.to_thread: 同步阻塞会卡死事件循环,
导致 pull-status/其他请求全部排队超时 (进度条 0%)。
"""
task = tasks.get(transfer_id)
if task is None:
return JSONResponse(status_code=404, content={"error": "任务不存在"})
try:
body = await request.json()
base_url = str(body["base_url"]).rstrip("/")
pull_token = str(body["token"])
except (KeyError, TypeError, ValueError):
return JSONResponse(status_code=400, content={"error": "缺少 base_url/token"})
aria2 = shutil.which("aria2c")
if aria2 is None:
return JSONResponse(status_code=409, content={"error": "服务端未安装 aria2 (apt install aria2)"})
chunk_count = task.get("chunk_count", 0)
chunk_dir = TMP_ROOT / transfer_id
chunk_dir.mkdir(parents=True, exist_ok=True)
urls_file = chunk_dir / "urls.txt"
with open(urls_file, "w", encoding="utf-8") as f:
for idx in range(1, chunk_count + 1):
f.write(f"{base_url}/pull/{transfer_id}/{idx}?token={pull_token}\n")
try:
r = await asyncio.to_thread(
subprocess.run,
[aria2, "-i", str(urls_file), "-d", str(chunk_dir),
"--max-concurrent-downloads=8", "--split=4", "--max-connection-per-server=4",
"--min-split-size=1M", "--file-allocation=none", "--allow-overwrite=true",
"--auto-file-renaming=false", "--console-log-level=notice",
"--summary-interval=1", "--quiet=false", "--no-conf=true"],
timeout=3600,
)
except OSError as e:
return JSONResponse(status_code=500, content={"error": f"aria2 启动失败: {e}"})
if r.returncode != 0:
return JSONResponse(status_code=409, content={"error": f"aria2 拉取失败 (exit {r.returncode}), 详见服务端日志"})
# 逐卷校验落盘 (大小 + SHA-256); aria2 落盘文件名 = URL 末段 (卷序号)
pulled = 0
for p in sorted(chunk_dir.iterdir()):
if p.name == "urls.txt" or p.suffix in (".aria2", ".tmp"):
continue
try:
idx = int(p.name)
except ValueError:
continue
try:
receiver.receive(transfer_id, idx, p.read_bytes())
p.unlink(missing_ok=True)
pulled += 1
except ReceiverError:
pass # 校验失败的卷忽略, 客户端会补传
urls_file.unlink(missing_ok=True)
log.info("pull %s %s %d/%d 卷 (aria2)", request.client.host if request.client else "?",
transfer_id, pulled, chunk_count)
return {"ok": True, "pulled": pulled, "total": chunk_count}
@app.get("/api/transfer/{transfer_id}/pull-status")
async def pull_status(transfer_id: str, request: Request, username: str = Depends(verify_token)):
"""aria2 拉取进度: 已落盘卷文件数 (aria2 每拉完一卷落盘, received 标记在全部拉完后才打)"""
task = tasks.get(transfer_id)
if task is None:
return JSONResponse(status_code=404, content={"error": "任务不存在"})
chunk_dir = TMP_ROOT / transfer_id
n = 0
if chunk_dir.is_dir():
for p in chunk_dir.iterdir():
if p.name != "urls.txt" and p.suffix not in (".aria2", ".tmp") and p.name.isdigit():
n += 1
return {"pulled": n, "total": task.get("chunk_count", 0)}
@app.post("/api/transfer/{transfer_id}/complete")
async def complete(transfer_id: str, request: Request, username: str = Depends(verify_token)):
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:
transfer = tasks.get(transfer_id)
assert transfer is not None
# 零合并: 只校验卷齐, 卷目录直接 move 入存储区, 秒回
got = tasks.received(transfer_id)
if len(got) != transfer["chunk_count"]:
return JSONResponse(status_code=409, content={
"status": "incomplete",
"received": sorted(got),
"error": f"缺卷: {transfer['chunk_count'] - len(got)} 卷未上传",
})
tasks.set_status(transfer_id, "storing")
final_dir = storage.store_chunks(transfer_id, TMP_ROOT / transfer_id, username)
file_id = db.insert_file(
transfer_id,
transfer["file_name"],
str(final_dir),
transfer["file_size"],
transfer["total_sha256"],
username,
transfer.get("dir_id"),
)
tasks.set_status(transfer_id, "done")
log.info("complete %s 任务=%s %s卷 -> file=%s",
request.client.host if request.client else "?", transfer_id,
transfer["chunk_count"], file_id)
return {"status": "done", "file_id": file_id}
except OSError as e:
tasks.set_status(transfer_id, "failed")
return JSONResponse(status_code=409, content={"error": str(e), "status": "failed"})