""" 入库层 (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 # 列已存在 try: conn.execute("ALTER TABLE transfers ADD COLUMN dir_id INTEGER") except sqlite3.OperationalError: pass # 列已存在 try: conn.execute("ALTER TABLE files ADD COLUMN dir_id INTEGER") 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, dir_id: int | None = None) -> 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, dir_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, dir_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, dir_id: int | None = None, ) -> 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, dir_id) " "VALUES (?, ?, ?, ?, ?, ?, ?, ?)", (file_id, transfer_id, file_name, path, size, sha256, user_id, dir_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 f.file_id, f.file_name, f.size, f.sha256, f.created_at, f.dir_id, d.name AS dir_name " "FROM files f LEFT JOIN dirs d ON d.id = f.dir_id " "WHERE f.user_id = ? OR f.user_id IS NULL ORDER BY f.created_at DESC", ("admin",), ).fetchall() else: rows = conn.execute( "SELECT f.file_id, f.file_name, f.size, f.sha256, f.created_at, f.dir_id, d.name AS dir_name " "FROM files f LEFT JOIN dirs d ON d.id = f.dir_id " "WHERE f.user_id = ? ORDER BY f.created_at DESC", (user_id,), ).fetchall() return [dict(r) for r in rows] def move_file(self, file_id: str, dir_id: int | None) -> bool: """移动文件到目录 (dir_id=None=根目录), 返回是否更新""" with self._conn() as conn: cur = conn.execute( "UPDATE files SET dir_id = ? WHERE file_id = ?", (dir_id, file_id) ) return cur.rowcount > 0 def clone_file(self, file_id: str, dir_id: int | None, new_path: str) -> str: """克隆文件记录 (卷目录已物理复制到 new_path), 返回新 file_id""" new_id = uuid.uuid4().hex[:12] with self._conn() as conn: row = conn.execute( "SELECT transfer_id, file_name, size, sha256, user_id FROM files WHERE file_id = ?", (file_id,), ).fetchone() if row is None: return "" conn.execute( "INSERT INTO files (file_id, transfer_id, file_name, path, size, sha256, user_id, dir_id) " "VALUES (?, ?, ?, ?, ?, ?, ?, ?)", (new_id, row["transfer_id"], row["file_name"], new_path, row["size"], row["sha256"], row["user_id"], dir_id), ) return new_id 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