账号体系: users 表 + auth.py(pbkdf2) + 注册/登录/注销端点 + 文件按用户隔离(admin 全部) + 测试 12 项
This commit is contained in:
@@ -41,17 +41,22 @@ SZ_TOKEN=你的token python3 main.py # 默认 0.0.0.0:8000
|
||||
## API
|
||||
|
||||
```
|
||||
POST /api/transfer/init 建任务 {transfer_id}
|
||||
PUT /api/transfer/{id}/chunk/{n} 上传单卷密文 409 = 坏卷
|
||||
GET /api/transfer/{id}/chunks 已收卷列表 断点续传
|
||||
POST /api/transfer/{id}/complete 卷齐入库(秒回) {file_id}
|
||||
GET /api/files 文件列表 文件名是加密的
|
||||
GET /api/files/{id}/chunk/{n} 下载单卷密文 X-Enc-Params 头带解密参数
|
||||
DELETE /api/files/{id} 删除文件
|
||||
GET /api/quota 空间配额
|
||||
全部端点需 Authorization: Bearer <token>
|
||||
POST /api/auth/register 注册账号 {username, token}
|
||||
POST /api/auth/login 登录 {username, token}
|
||||
POST /api/auth/logout 注销 (token 失效)
|
||||
POST /api/transfer/init 建任务 {transfer_id}
|
||||
PUT /api/transfer/{id}/chunk/{n} 上传单卷密文 409 = 坏卷
|
||||
GET /api/transfer/{id}/chunks 已收卷列表 断点续传
|
||||
POST /api/transfer/{id}/complete 卷齐入库(秒回) {file_id}
|
||||
GET /api/files 文件列表(按用户隔离) 文件名是加密的
|
||||
GET /api/files/{id}/chunk/{n} 下载单卷密文 X-Enc-Params 头带解密参数
|
||||
DELETE /api/files/{id} 删除文件(归属校验)
|
||||
GET /api/quota 空间配额(按用户)
|
||||
全部端点需 Authorization: Bearer <token> (SZ_TOKEN = admin)
|
||||
```
|
||||
|
||||
账号: 密码 pbkdf2-hmac-sha256 存储, 登录 token 随机 32 字节, 文件按用户名隔离。
|
||||
|
||||
## 日志与自愈
|
||||
|
||||
```
|
||||
|
||||
@@ -24,6 +24,8 @@ 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")
|
||||
@@ -45,10 +47,73 @@ app = FastAPI(title="7z-encrypt 服务端 (零知识 + 零合并)")
|
||||
|
||||
# ---------- 认证 ----------
|
||||
|
||||
def verify_token(authorization: str | None = Header(default=None)) -> None:
|
||||
"""Bearer token 校验, 失败 401"""
|
||||
if authorization != f"Bearer {TOKEN}":
|
||||
raise HTTPException(status_code=401, detail="未授权: token 无效或缺失")
|
||||
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 可操作全部, 普通用户只能操作自己的"""
|
||||
if username == "admin":
|
||||
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.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]:
|
||||
"""下载响应头: 解密参数 + 卷数"""
|
||||
@@ -61,18 +126,18 @@ def _enc_headers(rec: dict) -> dict[str, str]:
|
||||
|
||||
# ---------- 文件: ls / 逐卷下载 ----------
|
||||
|
||||
@app.get("/api/files", dependencies=[Depends(verify_token)])
|
||||
async def list_files(request: Request):
|
||||
@app.get("/api/files")
|
||||
async def list_files(request: Request, username: str = Depends(verify_token)):
|
||||
"""文件列表 (ls)"""
|
||||
files = db.list_files()
|
||||
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", dependencies=[Depends(verify_token)])
|
||||
async def quota(request: Request):
|
||||
@app.get("/api/quota")
|
||||
async def quota(request: Request, username: str = Depends(verify_token)):
|
||||
"""用户空间配额: 已用 (files size 总和) / 配额 / 剩余"""
|
||||
used = db.sum_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)
|
||||
@@ -84,12 +149,13 @@ async def quota(request: Request):
|
||||
}
|
||||
|
||||
|
||||
@app.delete("/api/files/{file_id}", dependencies=[Depends(verify_token)])
|
||||
async def delete_file(file_id: str, request: Request):
|
||||
@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():
|
||||
@@ -103,12 +169,13 @@ async def delete_file(file_id: str, request: Request):
|
||||
return {"ok": True, "file_id": file_id}
|
||||
|
||||
|
||||
@app.get("/api/files/{file_id}/chunk/{idx}", dependencies=[Depends(verify_token)])
|
||||
async def download_chunk(file_id: str, idx: int, request: Request):
|
||||
@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}"
|
||||
@@ -126,15 +193,15 @@ async def download_chunk(file_id: str, idx: int, request: Request):
|
||||
|
||||
# ---------- 传输 ----------
|
||||
|
||||
@app.post("/api/transfer/init", dependencies=[Depends(verify_token)])
|
||||
async def init_transfer(request: Request):
|
||||
@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()
|
||||
transfer_id = tasks.create(init_json)
|
||||
transfer_id = tasks.create(init_json, username)
|
||||
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:
|
||||
@@ -142,15 +209,15 @@ async def init_transfer(request: Request):
|
||||
return {"transfer_id": transfer_id}
|
||||
|
||||
|
||||
@app.get("/api/transfer/{transfer_id}/chunks", dependencies=[Depends(verify_token)])
|
||||
async def get_chunks(transfer_id: str):
|
||||
@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}", dependencies=[Depends(verify_token)])
|
||||
async def put_chunk(transfer_id: str, idx: int, request: Request):
|
||||
@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()
|
||||
@@ -163,8 +230,8 @@ async def put_chunk(transfer_id: str, idx: int, request: Request):
|
||||
return {"ok": True}
|
||||
|
||||
|
||||
@app.post("/api/transfer/{transfer_id}/complete", dependencies=[Depends(verify_token)])
|
||||
async def complete(transfer_id: str, request: Request):
|
||||
@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
|
||||
@@ -190,6 +257,7 @@ async def complete(transfer_id: str, request: Request):
|
||||
str(final_dir),
|
||||
transfer["file_size"],
|
||||
transfer["total_sha256"],
|
||||
username,
|
||||
)
|
||||
tasks.set_status(transfer_id, "done")
|
||||
log.info("complete %s 任务=%s %s卷 -> file=%s",
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
"""
|
||||
认证工具 (auth)
|
||||
模块: 服务端 / 认证
|
||||
输入: 明文密码
|
||||
输出: pbkdf2 哈希 / 随机 token
|
||||
|
||||
标准库实现 (pbkdf2-hmac-sha256, 20 万轮迭代), 不引第三方
|
||||
密码不落库: 只存 salt$hash
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import secrets
|
||||
|
||||
_ITERATIONS = 200_000
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
"""密码 -> 'salt$hash' (pbkdf2-hmac-sha256, 随机盐)"""
|
||||
salt = secrets.token_hex(16)
|
||||
dk = hashlib.pbkdf2_hmac(
|
||||
"sha256", password.encode("utf-8"), bytes.fromhex(salt), _ITERATIONS
|
||||
)
|
||||
return f"{salt}${dk.hex()}"
|
||||
|
||||
|
||||
def verify_password(password: str, stored: str) -> bool:
|
||||
"""校验明文密码与存储哈希 (常数时间比较)"""
|
||||
try:
|
||||
salt, expected = stored.split("$", 1)
|
||||
dk = hashlib.pbkdf2_hmac(
|
||||
"sha256", password.encode("utf-8"), bytes.fromhex(salt), _ITERATIONS
|
||||
)
|
||||
return hmac.compare_digest(dk.hex(), expected)
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def new_token() -> str:
|
||||
"""随机 token (32 字节 hex, 会话凭证)"""
|
||||
return secrets.token_hex(32)
|
||||
@@ -48,6 +48,7 @@ class ServerDB:
|
||||
chunk_count INTEGER NOT NULL,
|
||||
total_sha256 TEXT NOT NULL,
|
||||
enc_params TEXT NOT NULL,
|
||||
user_id TEXT,
|
||||
status TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
@@ -75,18 +76,86 @@ class ServerDB:
|
||||
path TEXT NOT NULL,
|
||||
size INTEGER NOT NULL,
|
||||
sha256 TEXT NOT NULL,
|
||||
user_id TEXT,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
)
|
||||
"""
|
||||
)
|
||||
conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT UNIQUE NOT NULL,
|
||||
password_hash TEXT NOT NULL,
|
||||
token TEXT,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
)
|
||||
"""
|
||||
)
|
||||
# 旧库升级: files/transfers 补 user_id 列 (老数据 user_id=NULL -> admin 可见)
|
||||
try:
|
||||
conn.execute("ALTER TABLE transfers ADD COLUMN user_id TEXT")
|
||||
except sqlite3.OperationalError:
|
||||
pass # 列已存在
|
||||
try:
|
||||
conn.execute("ALTER TABLE files ADD COLUMN user_id TEXT")
|
||||
except sqlite3.OperationalError:
|
||||
pass # 列已存在
|
||||
|
||||
# ---------- users ----------
|
||||
|
||||
def create_user(self, username: str, password_hash: str) -> bool:
|
||||
"""注册用户, 返回是否成功 (False = 用户名已存在)"""
|
||||
try:
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO users (username, password_hash) VALUES (?, ?)",
|
||||
(username, password_hash),
|
||||
)
|
||||
return True
|
||||
except sqlite3.IntegrityError:
|
||||
return False
|
||||
|
||||
def get_user(self, username: str) -> dict[str, Any] | None:
|
||||
with self._conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT id, username, password_hash, token FROM users WHERE username = ?",
|
||||
(username,),
|
||||
).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def set_token(self, username: str, token: str) -> None:
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE users SET token = ? WHERE username = ?", (token, username)
|
||||
)
|
||||
|
||||
def get_username_by_token(self, token: str) -> str | None:
|
||||
with self._conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT username FROM users WHERE token = ?", (token,)
|
||||
).fetchone()
|
||||
return str(row["username"]) if row else None
|
||||
|
||||
def clear_token(self, token: str) -> None:
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE users SET token = NULL WHERE token = ?", (token,)
|
||||
)
|
||||
|
||||
def clear_token_by_username(self, username: str) -> None:
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE users SET token = NULL WHERE username = ?", (username,)
|
||||
)
|
||||
|
||||
# ---------- transfers ----------
|
||||
|
||||
def create_transfer(self, transfer_id: str, init_json: dict[str, Any]) -> None:
|
||||
def create_transfer(self, transfer_id: str, init_json: dict[str, Any], user_id: str) -> None:
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO transfers (transfer_id, file_name, file_size, chunk_count, total_sha256, enc_params, status) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?)",
|
||||
"INSERT INTO transfers (transfer_id, file_name, file_size, chunk_count, total_sha256, enc_params, user_id, status) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
(
|
||||
transfer_id,
|
||||
init_json["file_name"],
|
||||
@@ -94,6 +163,7 @@ class ServerDB:
|
||||
init_json["chunk_count"],
|
||||
init_json["total_sha256"],
|
||||
json.dumps(init_json["enc"], ensure_ascii=False),
|
||||
user_id,
|
||||
"uploading",
|
||||
),
|
||||
)
|
||||
@@ -167,13 +237,14 @@ class ServerDB:
|
||||
path: str,
|
||||
size: int,
|
||||
sha256: str,
|
||||
user_id: str,
|
||||
) -> str:
|
||||
file_id = uuid.uuid4().hex[:12]
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO files (file_id, transfer_id, file_name, path, size, sha256) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?)",
|
||||
(file_id, transfer_id, file_name, path, size, sha256),
|
||||
"INSERT INTO files (file_id, transfer_id, file_name, path, size, sha256, user_id) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?)",
|
||||
(file_id, transfer_id, file_name, path, size, sha256, user_id),
|
||||
)
|
||||
return file_id
|
||||
|
||||
@@ -186,13 +257,20 @@ class ServerDB:
|
||||
).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def list_files(self) -> list[dict[str, Any]]:
|
||||
"""全部文件记录 (ls 端点)"""
|
||||
def list_files(self, user_id: str) -> list[dict[str, Any]]:
|
||||
"""文件记录 (ls 端点): admin 看全部, 普通用户看自己的"""
|
||||
with self._conn() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT file_id, file_name, size, sha256, created_at FROM files "
|
||||
"ORDER BY created_at DESC"
|
||||
).fetchall()
|
||||
if user_id == "admin":
|
||||
rows = conn.execute(
|
||||
"SELECT file_id, file_name, size, sha256, created_at FROM files "
|
||||
"ORDER BY created_at DESC"
|
||||
).fetchall()
|
||||
else:
|
||||
rows = conn.execute(
|
||||
"SELECT file_id, file_name, size, sha256, created_at FROM files "
|
||||
"WHERE user_id = ? ORDER BY created_at DESC",
|
||||
(user_id,),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def delete_file(self, file_id: str) -> bool:
|
||||
@@ -201,17 +279,23 @@ class ServerDB:
|
||||
cur = conn.execute("DELETE FROM files WHERE file_id = ?", (file_id,))
|
||||
return cur.rowcount > 0
|
||||
|
||||
def sum_files_size(self) -> int:
|
||||
"""全部文件 size 总和 (配额统计用)"""
|
||||
def sum_files_size(self, user_id: str) -> int:
|
||||
"""文件 size 总和 (配额统计用): admin 全部, 普通用户自己的"""
|
||||
with self._conn() as conn:
|
||||
row = conn.execute("SELECT COALESCE(SUM(size), 0) FROM files").fetchone()
|
||||
if user_id == "admin":
|
||||
row = conn.execute("SELECT COALESCE(SUM(size), 0) FROM files").fetchone()
|
||||
else:
|
||||
row = conn.execute(
|
||||
"SELECT COALESCE(SUM(size), 0) FROM files WHERE user_id = ?",
|
||||
(user_id,),
|
||||
).fetchone()
|
||||
return int(row[0])
|
||||
|
||||
def get_file(self, file_id: str) -> dict[str, Any] | None:
|
||||
"""单文件记录 (下载端点, 含解密参数 enc_params + 卷数)"""
|
||||
with self._conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT f.file_id, f.file_name, f.path, f.size, f.sha256, "
|
||||
"SELECT f.file_id, f.file_name, f.path, f.size, f.sha256, f.user_id, "
|
||||
"t.enc_params, t.chunk_count "
|
||||
"FROM files f JOIN transfers t ON f.transfer_id = t.transfer_id "
|
||||
"WHERE f.file_id = ?",
|
||||
|
||||
+5
-5
@@ -26,14 +26,14 @@ class TaskManager:
|
||||
def __init__(self, db: ServerDB) -> None:
|
||||
self.db = db
|
||||
|
||||
def create(self, init_json: dict[str, Any]) -> str:
|
||||
def create(self, init_json: dict[str, Any], user_id: str) -> str:
|
||||
"""POST /init: 建任务, 返回 transfer_id (输入清洗: 文件名/卷数上限)
|
||||
|
||||
文件名只做可打印字符过滤 + 长度截断。不删 / \\ 等路径分隔符:
|
||||
加密文件名是 URL-safe base64 (可能含 - _), 路径安全由客户端 _safe_name
|
||||
和存储目录 (transfer_id) 保证。
|
||||
加密文件名是 URL-safe base64 (可能含 - _ =), 路径安全由下载端 _safe_name 兜底
|
||||
"""
|
||||
raw = str(init_json.get("file_name", "unnamed"))
|
||||
# 文件名字段清洗: 长度上限 + 去路径分隔符/控制字符
|
||||
raw = str(init_json.get("file_name", ""))
|
||||
clean = "".join(c for c in raw if c.isprintable())[:255]
|
||||
init_json["file_name"] = clean or "unnamed"
|
||||
# 卷数上限: 防超大清单放大 (1MB init / 120B 每卷 spec ≈ 8000)
|
||||
@@ -42,7 +42,7 @@ class TaskManager:
|
||||
raise ValueError(f"卷数超上限: {len(chunks)} > 10000")
|
||||
init_json["chunk_count"] = len(chunks)
|
||||
transfer_id = uuid.uuid4().hex[:12]
|
||||
self.db.create_transfer(transfer_id, init_json)
|
||||
self.db.create_transfer(transfer_id, init_json, user_id)
|
||||
return transfer_id
|
||||
|
||||
def get(self, transfer_id: str) -> dict[str, Any] | None:
|
||||
|
||||
@@ -294,3 +294,98 @@ class ServerPipelineTest(unittest.TestCase):
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
|
||||
class AccountTest(unittest.TestCase):
|
||||
"""账号体系: 注册 / 登录 / 注销 / 文件归属隔离"""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
# 清库重建 (账号测试需要干净的用户表, 防跨运行残留)
|
||||
for f in ("/tmp/sz-test.db", "/tmp/sz-test.db-wal", "/tmp/sz-test.db-shm"):
|
||||
os.path.exists(f) and os.unlink(f)
|
||||
from db import ServerDB
|
||||
ServerDB() # 重建 schema (api 全局 db 实例后续连接即见新表)
|
||||
from api import app
|
||||
cls.client = TestClient(app)
|
||||
# 客户端密钥 (测试用, 与 ServerPipelineTest 同款)
|
||||
client_kr = os.path.join(tempfile.mkdtemp(prefix="sz-acc-kr-"), "keyring.json")
|
||||
engine = CryptoEngine(client_kr)
|
||||
engine.generate_key(TEST_KEY_ID)
|
||||
cls.client_kr = client_kr
|
||||
|
||||
def _register(self, name: str, pw: str = "pass123456"):
|
||||
return self.client.post(
|
||||
"/api/auth/register", json={"username": name, "password": pw}
|
||||
)
|
||||
|
||||
def test_register_login_logout(self):
|
||||
# 注册 -> 自动登录返回 token
|
||||
r = self._register("alice")
|
||||
self.assertEqual(r.status_code, 200, r.text)
|
||||
token = r.json()["token"]
|
||||
self.assertTrue(token)
|
||||
auth = {"Authorization": f"Bearer {token}"}
|
||||
# token 可用
|
||||
self.assertEqual(self.client.get("/api/files", headers=auth).status_code, 200)
|
||||
# 重名注册 -> 409
|
||||
self.assertEqual(self._register("alice").status_code, 409)
|
||||
# 注销 -> token 失效 -> 401
|
||||
self.assertEqual(self.client.post("/api/auth/logout", headers=auth).status_code, 200)
|
||||
self.assertEqual(self.client.get("/api/files", headers=auth).status_code, 401)
|
||||
# 重新登录 -> 新 token 可用
|
||||
r = self.client.post(
|
||||
"/api/auth/login", json={"username": "alice", "password": "pass123456"}
|
||||
)
|
||||
self.assertEqual(r.status_code, 200, r.text)
|
||||
auth2 = {"Authorization": f"Bearer {r.json()['token']}"}
|
||||
self.assertEqual(self.client.get("/api/files", headers=auth2).status_code, 200)
|
||||
|
||||
def test_register_validation(self):
|
||||
# 用户名/密码规则
|
||||
self.assertEqual(self._register("ab", "pass123456").status_code, 400) # 太短
|
||||
self.assertEqual(self._register("a b c", "pass123456").status_code, 400) # 非字母数字
|
||||
self.assertEqual(self._register("bob", "12345").status_code, 400) # 密码太短
|
||||
# 错误密码登录 -> 401
|
||||
r = self.client.post(
|
||||
"/api/auth/login", json={"username": "bob", "password": "wrongpass"}
|
||||
)
|
||||
self.assertEqual(r.status_code, 401)
|
||||
|
||||
def test_user_file_isolation(self):
|
||||
# alice 上传文件, bob 看不到也不能下载/删除
|
||||
init_json, chunk_files, _ = _prepare(self.client_kr, os.urandom(256 << 10))
|
||||
ra = self.client.post(
|
||||
"/api/auth/register", json={"username": "u_isola", "password": "pass123456"}
|
||||
)
|
||||
auth_a = {"Authorization": f"Bearer {ra.json()['token']}"}
|
||||
rb = self.client.post(
|
||||
"/api/auth/register", json={"username": "u_isolb", "password": "pass123456"}
|
||||
)
|
||||
auth_b = {"Authorization": f"Bearer {rb.json()['token']}"}
|
||||
|
||||
r = self.client.post("/api/transfer/init", json=init_json, headers=auth_a)
|
||||
tid = r.json()["transfer_id"]
|
||||
for idx in sorted(chunk_files):
|
||||
with open(chunk_files[idx], "rb") as f:
|
||||
self.client.put(
|
||||
f"/api/transfer/{tid}/chunk/{idx}", content=f.read(), headers=auth_a
|
||||
)
|
||||
r = self.client.post(f"/api/transfer/{tid}/complete", headers=auth_a)
|
||||
file_id = r.json()["file_id"]
|
||||
|
||||
# bob 的 ls 看不到 alice 的文件
|
||||
r = self.client.get("/api/files", headers=auth_b)
|
||||
self.assertTrue(all(f["file_id"] != file_id for f in r.json()["files"]))
|
||||
# bob 下载 alice 的文件 -> 403
|
||||
r = self.client.get(f"/api/files/{file_id}/chunk/1", headers=auth_b)
|
||||
self.assertEqual(r.status_code, 403)
|
||||
# bob 删除 alice 的文件 -> 403
|
||||
r = self.client.delete(f"/api/files/{file_id}", headers=auth_b)
|
||||
self.assertEqual(r.status_code, 403)
|
||||
# alice 自己可以下载/删除
|
||||
r = self.client.get(f"/api/files/{file_id}/chunk/1", headers=auth_a)
|
||||
self.assertEqual(r.status_code, 200)
|
||||
# admin (SZ_TOKEN) 可见全部
|
||||
r = self.client.get("/api/files", headers=AUTH)
|
||||
self.assertTrue(any(f["file_id"] == file_id for f in r.json()["files"]))
|
||||
|
||||
Reference in New Issue
Block a user