账号体系: users 表 + auth.py(pbkdf2) + 注册/登录/注销端点 + 文件按用户隔离(admin 全部) + 测试 12 项

This commit is contained in:
lou
2026-08-10 15:27:45 +08:00
parent 7e6c7d64c6
commit b79aab5c4e
6 changed files with 346 additions and 53 deletions
+14 -9
View File
@@ -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 字节, 文件按用户名隔离。
## 日志与自愈
```
+91 -23
View File
@@ -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",
+41
View File
@@ -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)
+100 -16
View File
@@ -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
View File
@@ -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:
+95
View File
@@ -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"]))