225 lines
8.3 KiB
Python
225 lines
8.3 KiB
Python
"""
|
|
入库层 (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,
|
|
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,
|
|
created_at TEXT NOT NULL DEFAULT (datetime('now'))
|
|
)
|
|
"""
|
|
)
|
|
|
|
# ---------- transfers ----------
|
|
|
|
def create_transfer(self, transfer_id: str, init_json: dict[str, Any]) -> 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 (?, ?, ?, ?, ?, ?, ?)",
|
|
(
|
|
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),
|
|
"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,
|
|
) -> 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),
|
|
)
|
|
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) -> list[dict[str, Any]]:
|
|
"""全部文件记录 (ls 端点)"""
|
|
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()
|
|
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) -> int:
|
|
"""全部文件 size 总和 (配额统计用)"""
|
|
with self._conn() as conn:
|
|
row = conn.execute("SELECT COALESCE(SUM(size), 0) FROM files").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, "
|
|
"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
|