Files
7z-encrypt-server/db.py
T

361 lines
14 KiB
Python
Raw Normal View History

"""
入库层 (db)
模块: 服务端 / 入库层
输入: 文件路径 + 元数据
输出: file_id
SQLite 实现 (老板拍板: PRoot 容器装 PG/MySQL 服务进程都会 fsync 卡死,
轻量文件库最稳)。表结构与 PostgreSQL 版一致, 未来可平滑切回 PG。
表:
transfers 任务主表 (状态机 uploading->assembling->decrypting->done/failed)
transfer_chunks 每卷信息 + 接收时间 (null = 未收, 断点续传依据)
files 明文文件记录
"""
import json
import sqlite3
import uuid
from pathlib import Path
from typing import Any
from settings import DB_PATH
class ServerDB:
"""服务端数据库访问 (sqlite3, 每次操作独立连接, WAL 并发安全)"""
def __init__(self, db_path: str | Path = DB_PATH) -> None:
self.db_path = Path(db_path)
self._init_schema()
def _conn(self) -> sqlite3.Connection:
conn = sqlite3.connect(self.db_path)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA journal_mode=WAL")
conn.execute("PRAGMA busy_timeout=5000")
return conn
def _init_schema(self) -> None:
self.db_path.parent.mkdir(parents=True, exist_ok=True)
with self._conn() as conn:
conn.execute(
"""
CREATE TABLE IF NOT EXISTS transfers (
transfer_id TEXT PRIMARY KEY,
file_name TEXT NOT NULL,
file_size INTEGER NOT NULL,
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'))
)
"""
)
conn.execute(
"""
CREATE TABLE IF NOT EXISTS transfer_chunks (
transfer_id TEXT NOT NULL REFERENCES transfers(transfer_id),
idx INTEGER NOT NULL,
size INTEGER NOT NULL,
sha256 TEXT NOT NULL,
received_at TEXT,
PRIMARY KEY (transfer_id, idx)
)
"""
)
conn.execute(
"""
CREATE TABLE IF NOT EXISTS files (
file_id TEXT PRIMARY KEY,
transfer_id TEXT NOT NULL REFERENCES transfers(transfer_id),
file_name TEXT NOT NULL,
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'))
)
"""
)
conn.execute(
"""
CREATE TABLE IF NOT EXISTS dirs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
user_id TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
UNIQUE (name, user_id)
)
"""
)
# 旧库升级: 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,)
)
# ---------- dirs ----------
def create_dir(self, name: str, user_id: str) -> int | None:
"""创建目录, 返回 dir id; None = 重名"""
try:
with self._conn() as conn:
cur = conn.execute(
"INSERT INTO dirs (name, user_id) VALUES (?, ?)", (name, user_id)
)
rid = cur.lastrowid
return int(rid) if rid is not None else None
except sqlite3.IntegrityError:
return None
def get_dir(self, dir_id: int) -> dict[str, Any] | None:
with self._conn() as conn:
row = conn.execute(
"SELECT id, name, user_id, created_at FROM dirs WHERE id = ?", (dir_id,)
).fetchone()
return dict(row) if row else None
def list_dirs(self, user_id: str) -> list[dict[str, Any]]:
with self._conn() as conn:
rows = conn.execute(
"SELECT id, name, created_at FROM dirs WHERE user_id = ? ORDER BY id",
(user_id,),
).fetchall()
return [dict(r) for r in rows]
def delete_dir(self, dir_id: int, user_id: str) -> bool:
"""删除目录 (归属校验), 返回是否删除"""
with self._conn() as conn:
cur = conn.execute(
"DELETE FROM dirs WHERE id = ? AND user_id = ?", (dir_id, user_id)
)
return cur.rowcount > 0
# ---------- transfers ----------
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, user_id, status) "
"VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
(
transfer_id,
init_json["file_name"],
init_json["file_size"],
init_json["chunk_count"],
init_json["total_sha256"],
json.dumps(init_json["enc"], ensure_ascii=False),
user_id,
"uploading",
),
)
conn.executemany(
"INSERT INTO transfer_chunks (transfer_id, idx, size, sha256) VALUES (?, ?, ?, ?)",
[
(transfer_id, ch["index"], ch["size"], ch["sha256"])
for ch in init_json["chunks"]
],
)
def get_transfer(self, transfer_id: str) -> dict[str, Any] | None:
with self._conn() as conn:
row = conn.execute(
"SELECT * FROM transfers WHERE transfer_id = ?", (transfer_id,)
).fetchone()
if not row:
return None
d = dict(row)
d["enc_params"] = json.loads(d["enc_params"])
return d
def set_status(self, transfer_id: str, status: str) -> None:
with self._conn() as conn:
conn.execute(
"UPDATE transfers SET status = ?, updated_at = datetime('now') WHERE transfer_id = ?",
(status, transfer_id),
)
# ---------- chunks ----------
def get_chunk_spec(self, transfer_id: str, idx: int) -> dict[str, Any] | None:
"""卷规格 (size/sha256), 接收校验用"""
with self._conn() as conn:
row = conn.execute(
"SELECT size, sha256 FROM transfer_chunks WHERE transfer_id = ? AND idx = ?",
(transfer_id, idx),
).fetchone()
return dict(row) if row else None
def is_chunk_received(self, transfer_id: str, idx: int) -> bool:
with self._conn() as conn:
row = conn.execute(
"SELECT 1 FROM transfer_chunks WHERE transfer_id = ? AND idx = ? AND received_at IS NOT NULL",
(transfer_id, idx),
).fetchone()
return row is not None
def mark_chunk_received(self, transfer_id: str, idx: int) -> None:
with self._conn() as conn:
conn.execute(
"UPDATE transfer_chunks SET received_at = datetime('now') "
"WHERE transfer_id = ? AND idx = ?",
(transfer_id, idx),
)
def get_received_chunks(self, transfer_id: str) -> set[int]:
with self._conn() as conn:
rows = conn.execute(
"SELECT idx FROM transfer_chunks WHERE transfer_id = ? AND received_at IS NOT NULL",
(transfer_id,),
).fetchall()
return {r["idx"] for r in rows}
# ---------- files ----------
def insert_file(
self,
transfer_id: str,
file_name: str,
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, user_id) "
"VALUES (?, ?, ?, ?, ?, ?, ?)",
(file_id, transfer_id, file_name, path, size, sha256, user_id),
)
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, user_id: str) -> list[dict[str, Any]]:
"""文件记录 (ls 端点): 每个账户只看自己的空间 (admin 含无主旧文件)"""
with self._conn() as conn:
if user_id == "admin":
rows = conn.execute(
"SELECT file_id, file_name, size, sha256, created_at FROM files "
"WHERE user_id = ? OR user_id IS NULL ORDER BY created_at DESC",
("admin",),
).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:
"""删除文件记录, 返回是否删除 (del 端点)"""
with self._conn() as conn:
cur = conn.execute("DELETE FROM files WHERE file_id = ?", (file_id,))
return cur.rowcount > 0
def sum_files_size(self, user_id: str) -> int:
"""文件 size 总和 (配额统计用): 每个账户只算自己的空间"""
with self._conn() as conn:
if user_id == "admin":
row = conn.execute(
"SELECT COALESCE(SUM(size), 0) FROM files WHERE user_id = ? OR user_id IS NULL",
("admin",),
).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, 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 = ?",
(file_id,),
).fetchone()
if not row:
return None
d = dict(row)
d["enc_params"] = json.loads(d["enc_params"])
return d