"""数据库 + .env 备份 / 恢复(自动落盘到 /root/eth_hedge_backups)。""" from __future__ import annotations import asyncio import json import logging import os import shutil import sqlite3 import tempfile import time import zipfile from pathlib import Path from typing import Any from .config import get_settings from .credentials import resolve_env_file_path from .models.db import Database, get_db logger = logging.getLogger(__name__) MANIFEST_NAME = "manifest.json" DB_ARCNAME = "hedge.db" ENV_ARCNAME = ".env" BACKUP_PREFIX = "eth_hedge_backup_" DEFAULT_KEEP = 14 DEFAULT_INTERVAL_HOURS = 24 def resolve_backup_dir() -> Path: """优先 BACKUP_DIR;Linux 默认 /root/eth_hedge_backups;否则用户目录。""" override = (os.environ.get("BACKUP_DIR") or "").strip() if override: p = Path(override) p.mkdir(parents=True, exist_ok=True) return p if os.name != "nt": root = Path("/root/eth_hedge_backups") try: if Path("/root").is_dir(): root.mkdir(parents=True, exist_ok=True) return root except OSError as e: logger.warning("cannot use /root/eth_hedge_backups: %s", e) p = Path.home() / "eth_hedge_backups" p.mkdir(parents=True, exist_ok=True) return p def resolve_env_path() -> Path | None: p = resolve_env_file_path() return p def _settings_int(db: Database, key: str, default: int) -> int: raw = db.get_setting(key, str(default)) try: return int(float(raw or default)) except (TypeError, ValueError): return default def _settings_bool(db: Database, key: str, default: bool) -> bool: raw = db.get_setting(key, "1" if default else "0") if raw is None or raw == "": return default return str(raw).strip().lower() in ("1", "true", "yes", "on") def backup_settings(db: Database | None = None) -> dict[str, Any]: db = db or get_db() return { "auto_enabled": _settings_bool(db, "backup_auto_enabled", True), "interval_hours": max(1, _settings_int(db, "backup_interval_hours", DEFAULT_INTERVAL_HOURS)), "keep_count": max(1, min(90, _settings_int(db, "backup_keep_count", DEFAULT_KEEP))), "last_at_ms": _settings_int(db, "backup_last_at_ms", 0) or None, "backup_dir": str(resolve_backup_dir()), } def list_backups() -> list[dict[str, Any]]: d = resolve_backup_dir() items: list[dict[str, Any]] = [] if not d.is_dir(): return items for p in sorted(d.glob(f"{BACKUP_PREFIX}*.zip"), key=lambda x: x.stat().st_mtime, reverse=True): st = p.stat() items.append( { "name": p.name, "path": str(p), "size_bytes": st.st_size, "mtime_ms": int(st.st_mtime * 1000), } ) return items def prune_backups(keep: int | None = None) -> int: keep_n = keep if keep is not None else backup_settings()["keep_count"] items = list_backups() removed = 0 for old in items[int(keep_n) :]: try: Path(old["path"]).unlink(missing_ok=True) removed += 1 except OSError as e: logger.warning("prune backup failed %s: %s", old.get("name"), e) return removed def _sqlite_snapshot(db: Database, dest: Path) -> None: dest.parent.mkdir(parents=True, exist_ok=True) if dest.exists(): dest.unlink() with db._lock: dst = sqlite3.connect(str(dest)) try: db._conn.backup(dst) dst.commit() finally: dst.close() def create_backup( *, db: Database | None = None, reason: str = "manual", include_env: bool | None = None, ) -> dict[str, Any]: db = db or get_db() backup_dir = resolve_backup_dir() ts = time.strftime("%Y%m%d_%H%M%S") out = backup_dir / f"{BACKUP_PREFIX}{ts}.zip" env_path = resolve_env_path() db_path = Path(db.path) # 机内备份默认可含 .env;BACKUP_INCLUDE_ENV=0 可关。HTTP 下载另做剥离。 if include_env is None: include_env = (os.environ.get("BACKUP_INCLUDE_ENV") or "1").strip() != "0" with tempfile.TemporaryDirectory(prefix="eth_hedge_bak_") as td: tmp_db = Path(td) / "hedge.db" _sqlite_snapshot(db, tmp_db) has_env = bool(include_env and env_path and env_path.is_file()) manifest = { "product": "比特骆驼自动化对冲系统", "version": 1, "created_at_ms": int(time.time() * 1000), "reason": reason, "db_source": str(db_path), "env_source": str(env_path) if has_env else None, "has_env": has_env, "mode": get_settings().mode, "env_name": get_settings().env_name, } with zipfile.ZipFile(out, "w", compression=zipfile.ZIP_DEFLATED) as zf: zf.write(tmp_db, DB_ARCNAME) if has_env and env_path is not None: zf.write(env_path, ENV_ARCNAME) zf.writestr( MANIFEST_NAME, json.dumps(manifest, ensure_ascii=False, indent=2), ) db.set_setting("backup_last_at_ms", str(int(time.time() * 1000))) prune_backups() st = out.stat() logger.info("backup created path=%s reason=%s size=%s has_env=%s", out, reason, st.st_size, has_env) return { "ok": True, "name": out.name, "path": str(out), "size_bytes": st.st_size, "mtime_ms": int(st.st_mtime * 1000), "backup_dir": str(backup_dir), "manifest": manifest, } def materialize_download_zip(src: Path) -> Path: """HTTP 下载用:去掉包内 .env,避免令牌失窃带走交易所密钥。""" fd, name = tempfile.mkstemp(prefix="eth_hedge_dl_", suffix=".zip") os.close(fd) dest = Path(name) with zipfile.ZipFile(src, "r") as zin, zipfile.ZipFile( dest, "w", compression=zipfile.ZIP_DEFLATED ) as zout: for info in zin.infolist(): base = Path(info.filename).name if base == ENV_ARCNAME or base.endswith(".env"): continue zout.writestr(info, zin.read(info.filename)) return dest def read_backup_file(name: str) -> Path: safe = Path(name).name if not safe.startswith(BACKUP_PREFIX) or not safe.endswith(".zip"): raise ValueError("非法备份文件名") path = resolve_backup_dir() / safe if not path.is_file(): raise FileNotFoundError(f"备份不存在: {safe}") return path def _extract_backup_zip(zip_path: Path, dest_dir: Path) -> dict[str, Path]: with zipfile.ZipFile(zip_path, "r") as zf: names = set(zf.namelist()) if DB_ARCNAME not in names: # 兼容偶发相对路径 db_candidates = [n for n in names if n.endswith("hedge.db") or n.endswith(".db")] if not db_candidates: raise ValueError("备份包缺少 hedge.db") db_name = db_candidates[0] else: db_name = DB_ARCNAME env_name = ENV_ARCNAME if ENV_ARCNAME in names else ( next((n for n in names if n.endswith(".env") or n == "env"), None) ) zf.extract(db_name, dest_dir) db_out = dest_dir / db_name if not db_out.is_file(): # zip 内可能带目录 found = list(dest_dir.rglob("*.db")) if not found: raise ValueError("解压后未找到数据库文件") db_out = found[0] env_out: Path | None = None if env_name: zf.extract(env_name, dest_dir) env_out = dest_dir / env_name if not env_out.is_file(): found_env = list(dest_dir.rglob(".env")) + list(dest_dir.rglob("*.env")) env_out = found_env[0] if found_env else None return {"db": db_out, "env": env_out} # type: ignore[dict-item] def validate_backup_zip(zip_path: Path) -> dict[str, Any]: with zipfile.ZipFile(zip_path, "r") as zf: names = zf.namelist() has_db = any(n.endswith("hedge.db") or n == DB_ARCNAME or n.endswith(".db") for n in names) has_env = any(n.endswith(".env") or n == ENV_ARCNAME for n in names) manifest = None if MANIFEST_NAME in names: try: manifest = json.loads(zf.read(MANIFEST_NAME).decode("utf-8")) except Exception: manifest = None if not has_db: raise ValueError("备份包无效:缺少数据库") return {"has_db": has_db, "has_env": has_env, "manifest": manifest, "files": names} def restore_from_zip( zip_path: Path, *, db: Database | None = None, make_safety_backup: bool = True, ) -> dict[str, Any]: """用备份包覆盖当前 hedge.db 与 .env。调用方应在此前后暂停策略并准备重启进程。""" db = db or get_db() info = validate_backup_zip(zip_path) safety = None if make_safety_backup: try: safety = create_backup(db=db, reason="pre_restore") except Exception as e: logger.warning("pre_restore backup failed: %s", e) with tempfile.TemporaryDirectory(prefix="eth_hedge_restore_") as td: extracted = _extract_backup_zip(zip_path, Path(td)) src_db = extracted["db"] src_env = extracted.get("env") db_path = Path(db.path) bak_live: Path | None = None db.close() if db_path.is_file(): bak_live = db_path.with_suffix(db_path.suffix + f".pre_restore_{int(time.time())}") shutil.copy2(db_path, bak_live) shutil.copy2(src_db, db_path) env_written = None if src_env and Path(src_env).is_file(): env_target = resolve_env_file_path() env_target.parent.mkdir(parents=True, exist_ok=True) if env_target.is_file(): shutil.copy2(env_target, env_target.with_suffix(".env.pre_restore")) shutil.copy2(src_env, env_target) try: env_target.chmod(0o600) except Exception: pass env_written = str(env_target) return { "ok": True, "db_path": str(db_path), "env_path": env_written, "had_env_in_backup": bool(info.get("has_env")), "safety_backup": safety.get("name") if safety else None, "local_db_copy": str(bak_live) if bak_live and Path(bak_live).exists() else None, "restart_required": True, "manifest": info.get("manifest"), } _auto_task: asyncio.Task[None] | None = None async def auto_backup_loop() -> None: """后台定时备份;默认开启。""" logger.info("auto backup loop started dir=%s", resolve_backup_dir()) while True: try: await asyncio.sleep(60) db = get_db() cfg = backup_settings(db) if not cfg["auto_enabled"]: continue interval_ms = int(cfg["interval_hours"]) * 3600 * 1000 last = int(cfg["last_at_ms"] or 0) now = int(time.time() * 1000) if last and now - last < interval_ms: continue await asyncio.to_thread(create_backup, db=db, reason="auto") except asyncio.CancelledError: raise except Exception: logger.exception("auto backup tick failed") await asyncio.sleep(30) def start_auto_backup_task() -> asyncio.Task[None]: global _auto_task if _auto_task is None or _auto_task.done(): _auto_task = asyncio.create_task(auto_backup_loop(), name="auto-backup") return _auto_task async def stop_auto_backup_task() -> None: global _auto_task if _auto_task and not _auto_task.done(): _auto_task.cancel() try: await _auto_task except Exception: pass _auto_task = None