291 lines
9.7 KiB
Python
291 lines
9.7 KiB
Python
"""数据库备份 / 恢复 / 快照工具。
|
||
|
||
``data/*.db`` 在 .gitignore 里,仓库只带代码与电子书,不带数据库文件:
|
||
库一旦重置,章节、题目这些"内容"就跟着一起没了。这个工具补上两件事:
|
||
|
||
* **backup / restore**:本机时间戳备份,seed --force 前自动执行;
|
||
* **dump / load**:把目录类内容(教材、章节、知识点、课程、题库)导成
|
||
可直接提交的 SQL 文本,放进 git 就等于远端存了一份数据库快照,
|
||
而且不含账号、答题记录等个人数据。
|
||
|
||
用法见 scripts/db.sh(backend/ 目录下也可以 python -m dbtool)。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import sqlite3
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
|
||
from database import DATABASE_URL, DEFAULT_DB_PATH, REPO_ROOT
|
||
|
||
BACKUP_DIR = REPO_ROOT / "data" / "backups"
|
||
DUMP_DIR = REPO_ROOT / "data" / "dump"
|
||
DEFAULT_DUMP = DUMP_DIR / "catalog.sql"
|
||
|
||
# 目录类内容:可以公开提交的部分(不含用户、答题记录、错题本)
|
||
CATALOG_TABLES = [
|
||
"textbooks",
|
||
"chapters",
|
||
"knowledge",
|
||
"chapter_knowledge",
|
||
"courses",
|
||
"knowledge_relations",
|
||
"knowledge_resources",
|
||
"questions",
|
||
"question_knowledge",
|
||
]
|
||
|
||
# 个人学习数据:只进本机备份,不进 git
|
||
PRIVATE_TABLES = [
|
||
"users",
|
||
"roles",
|
||
"permissions",
|
||
"role_permissions",
|
||
"user_roles",
|
||
"user_knowledge",
|
||
"questions",
|
||
"question_knowledge",
|
||
"attempt_sessions",
|
||
"attempt_items",
|
||
"error_entries",
|
||
"daily_completions",
|
||
"ebook_progress",
|
||
"llm_settings",
|
||
]
|
||
|
||
|
||
def db_path() -> Path:
|
||
"""当前 SQLite 文件路径(非 sqlite 后端时退回默认路径)。"""
|
||
prefix = "sqlite:///"
|
||
if DATABASE_URL.startswith(prefix):
|
||
return Path(DATABASE_URL[len(prefix):])
|
||
return DEFAULT_DB_PATH
|
||
|
||
|
||
def backup(label: str = "", keep: int = 20) -> Path:
|
||
"""在线备份当前库(用 SQLite backup API,服务正在写也安全)。"""
|
||
source = db_path()
|
||
if not source.is_file():
|
||
raise FileNotFoundError(f"数据库不存在:{source}")
|
||
BACKUP_DIR.mkdir(parents=True, exist_ok=True)
|
||
stamp = datetime.now().strftime("%Y%m%d-%H%M%S")
|
||
suffix = f"-{_slug(label)}" if label else ""
|
||
target = BACKUP_DIR / f"{source.stem}-{stamp}{suffix}.db"
|
||
src = sqlite3.connect(source)
|
||
try:
|
||
dst = sqlite3.connect(target)
|
||
try:
|
||
src.backup(dst)
|
||
finally:
|
||
dst.close()
|
||
finally:
|
||
src.close()
|
||
prune(keep)
|
||
return target
|
||
|
||
|
||
def restore(source_file: Path, protect: bool = True) -> Path | None:
|
||
"""用备份覆盖当前库;覆盖前先把当前库存一份。"""
|
||
if not source_file.is_file():
|
||
raise FileNotFoundError(f"备份不存在:{source_file}")
|
||
saved = backup(label="pre-restore") if protect else None
|
||
src = sqlite3.connect(source_file)
|
||
try:
|
||
dst = sqlite3.connect(db_path())
|
||
try:
|
||
src.backup(dst)
|
||
finally:
|
||
dst.close()
|
||
finally:
|
||
src.close()
|
||
return saved
|
||
|
||
|
||
def prune(keep: int = 20) -> None:
|
||
"""只保留最近 keep 份备份。"""
|
||
if keep <= 0 or not BACKUP_DIR.is_dir():
|
||
return
|
||
files = sorted(BACKUP_DIR.glob("*.db"), key=lambda p: p.stat().st_mtime)
|
||
for old in files[:-keep]:
|
||
old.unlink(missing_ok=True)
|
||
|
||
|
||
def list_backups() -> list[Path]:
|
||
if not BACKUP_DIR.is_dir():
|
||
return []
|
||
return sorted(BACKUP_DIR.glob("*.db"), key=lambda p: p.name, reverse=True)
|
||
|
||
|
||
def _literal(value: object) -> str:
|
||
if value is None:
|
||
return "NULL"
|
||
if isinstance(value, (int, float)):
|
||
return repr(value)
|
||
if isinstance(value, (bytes, bytearray)):
|
||
return f"X'{bytes(value).hex()}'"
|
||
return "'" + str(value).replace("'", "''") + "'"
|
||
|
||
|
||
def dump_sql(
|
||
out: Path = DEFAULT_DUMP, tables: list[str] | None = None
|
||
) -> tuple[Path, int]:
|
||
"""把内容表导成 SQL 文本(先清表再按原 id 插回,可重复执行)。"""
|
||
wanted = tables or CATALOG_TABLES
|
||
conn = sqlite3.connect(db_path())
|
||
conn.row_factory = sqlite3.Row
|
||
try:
|
||
existing = {
|
||
row[0]
|
||
for row in conn.execute(
|
||
"SELECT name FROM sqlite_master WHERE type='table'"
|
||
)
|
||
}
|
||
lines = [
|
||
"-- nex_math 内容快照(教材 / 章节 / 知识点 / 课程 / 题库)",
|
||
f"-- 生成时间:{datetime.now().isoformat(timespec='seconds')}",
|
||
"-- 恢复:./scripts/db.sh load data/dump/catalog.sql",
|
||
"PRAGMA foreign_keys=OFF;",
|
||
"BEGIN TRANSACTION;",
|
||
]
|
||
rows_total = 0
|
||
for table in reversed(wanted):
|
||
if table in existing:
|
||
lines.append(f"DELETE FROM {table};")
|
||
for table in wanted:
|
||
if table not in existing:
|
||
continue
|
||
count = conn.execute(f"SELECT COUNT(*) FROM {table}").fetchone()[0]
|
||
if not count:
|
||
continue
|
||
lines.append(f"-- {table}: {count} 行")
|
||
for row in conn.execute(f"SELECT * FROM {table}"):
|
||
columns = row.keys()
|
||
values = ", ".join(_literal(row[key]) for key in columns)
|
||
lines.append(
|
||
f"INSERT INTO {table} ({', '.join(columns)}) "
|
||
f"VALUES ({values});"
|
||
)
|
||
rows_total += 1
|
||
lines += ["COMMIT;", "VACUUM;"]
|
||
finally:
|
||
conn.close()
|
||
out.parent.mkdir(parents=True, exist_ok=True)
|
||
out.write_text("\n".join(lines) + "\n", encoding="utf-8")
|
||
return out, rows_total
|
||
|
||
|
||
def _has_tables(conn: sqlite3.Connection) -> bool:
|
||
return (
|
||
conn.execute(
|
||
"SELECT 1 FROM sqlite_master WHERE type='table' LIMIT 1"
|
||
).fetchone()
|
||
is not None
|
||
)
|
||
|
||
|
||
def load_sql(script: Path) -> bool:
|
||
"""把 dump 出来的 SQL 文本灌回当前库(内容表按快照里的状态覆盖)。
|
||
|
||
库文件不存在或还没有表时,先按 ORM 模型建好表结构再灌数据,
|
||
这样新克隆的仓库也能一条命令导入远端快照。返回值表示本次是否新建了库。
|
||
"""
|
||
if not script.is_file():
|
||
raise FileNotFoundError(f"快照不存在:{script}")
|
||
target = db_path()
|
||
target.parent.mkdir(parents=True, exist_ok=True)
|
||
created = False
|
||
conn = sqlite3.connect(target)
|
||
try:
|
||
if not _has_tables(conn):
|
||
conn.close()
|
||
# models 必须先导入,Base.metadata 才知道要建哪些表
|
||
import models # noqa: F401
|
||
from database import ensure_schema # 建表(含旧库补列)
|
||
|
||
ensure_schema()
|
||
created = True
|
||
conn = sqlite3.connect(target)
|
||
conn.executescript(script.read_text(encoding="utf-8"))
|
||
conn.commit()
|
||
finally:
|
||
conn.close()
|
||
return created
|
||
|
||
|
||
def _slug(text: str) -> str:
|
||
cleaned = "".join(ch if ch.isalnum() or ch in "-_" else "-" for ch in text)
|
||
return cleaned.strip("-")[:32] or "backup"
|
||
|
||
|
||
def main(argv: list[str] | None = None) -> int:
|
||
parser = argparse.ArgumentParser(description="nex_math 数据库备份工具")
|
||
sub = parser.add_subparsers(dest="cmd", required=True)
|
||
|
||
p_backup = sub.add_parser("backup", help="备份当前库")
|
||
p_backup.add_argument("label", nargs="?", default="", help="备份标签")
|
||
|
||
p_restore = sub.add_parser("restore", help="用备份覆盖当前库")
|
||
p_restore.add_argument("file", nargs="?", default="latest", help="备份文件")
|
||
|
||
sub.add_parser("list", help="列出本机备份")
|
||
|
||
p_dump = sub.add_parser("dump", help="导出可提交的内容快照")
|
||
p_dump.add_argument("out", nargs="?", default=str(DEFAULT_DUMP))
|
||
p_dump.add_argument("--full", action="store_true", help="含用户与学习数据")
|
||
|
||
p_load = sub.add_parser("load", help="把内容快照灌回当前库")
|
||
p_load.add_argument("file", nargs="?", default=str(DEFAULT_DUMP))
|
||
|
||
sub.add_parser("path", help="打印数据库文件路径")
|
||
|
||
args = parser.parse_args(argv)
|
||
|
||
if args.cmd == "path":
|
||
print(db_path())
|
||
elif args.cmd == "backup":
|
||
print(backup(label=args.label))
|
||
elif args.cmd == "list":
|
||
items = list_backups()
|
||
if not items:
|
||
print(f"({BACKUP_DIR} 下暂无备份)")
|
||
for item in items:
|
||
print(f"{item} {item.stat().st_size // 1024} KB")
|
||
elif args.cmd == "restore":
|
||
target = (
|
||
list_backups()[0]
|
||
if args.file == "latest"
|
||
else Path(args.file)
|
||
)
|
||
if not target.is_absolute():
|
||
target = (BACKUP_DIR / args.file).resolve()
|
||
saved = restore(target)
|
||
print(f"已恢复:{target}")
|
||
if saved:
|
||
print(f"覆盖前的库已备份到:{saved}")
|
||
print("提示:重启后端后生效。")
|
||
elif args.cmd == "dump":
|
||
if args.full:
|
||
# 含账号与学习数据:只写进本机备份目录,不进 git
|
||
tables = CATALOG_TABLES + [
|
||
t for t in PRIVATE_TABLES if t not in CATALOG_TABLES
|
||
]
|
||
out = BACKUP_DIR / "math.full.sql"
|
||
else:
|
||
tables = CATALOG_TABLES
|
||
out = Path(args.out)
|
||
written, rows = dump_sql(out, tables)
|
||
print(f"{written} {rows} 行")
|
||
elif args.cmd == "load":
|
||
created = load_sql(Path(args.file))
|
||
print(f"已导入:{args.file}")
|
||
if created:
|
||
print("提示:这是一个新建的库。快照只含内容(教材/章节/知识点/题目),"
|
||
"账号与学习数据请用 ./scripts/init.sh 或 python -m seed 生成。")
|
||
print("提示:重启后端后生效。")
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
raise SystemExit(main())
|