nex_math/backend/dbtool.py

291 lines
9.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""数据库备份 / 恢复 / 快照工具。
``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())