"""Restore planning and execution with path safety.""" from __future__ import annotations import fnmatch import json import os import secrets import time from dataclasses import dataclass, field from pathlib import Path from typing import Any, Literal from .db import Database from . import dbsafe from .ingest import Repository, is_within from .store import BlobStore, sha256 PLAN_TTL_SECONDS = 15 * 60 Conflict = Literal["overwrite", "skip", "rename", "fail"] class RestoreError(Exception): def __init__(self, code: str, detail: str): super().__init__(detail) self.code = code self.detail = detail @dataclass class Criteria: as_of: float | None = None paths: list[str] = field(default_factory=list) globs: list[str] = field(default_factory=list) exclude_globs: list[str] = field(default_factory=list) extensions: list[str] = field(default_factory=list) root_id: int | None = None project_id: int | None = None changed_since: float | None = None changed_by_source: str | None = None target_dir: str | None = None conflict: Conflict = "overwrite" remove_newer_files: bool = False class Restorer: def __init__(self, db: Database, repo: Repository, blobs: BlobStore): self.db = db self.repo = repo self.blobs = blobs # Async (sha256) -> bytes; set by the service layer for on-demand # remote fetch after reindex. Defaults to local spool only. self.fetch_blob = None # path safety def allowed_bases(self) -> list[str]: return [str(Path.home())] + [r["path"] for r in self.db.roots()] @staticmethod def _protected_dirs() -> list[str]: try: from .config import Paths p = Paths.resolve() return [str(p.config_dir), str(p.data_dir), str(p.cache_dir)] except Exception: return [] def safe_target(self, target: str) -> str: """Resolve a write target; refuse traversal, symlink escapes and symlink targets.""" if not os.path.isabs(target): raise RestoreError("unsafe-path", f"{target} is not absolute") normalized = os.path.normpath(target) if normalized != target.rstrip("/") or ".." in Path(target).parts: raise RestoreError("unsafe-path", f"{target} is not a normalized path") for protected in self._protected_dirs(): if is_within(normalized, protected): raise RestoreError("unsafe-path", f"{target} is inside versiond's own data directory {protected}") parent = Path(normalized).parent existing = parent while not existing.exists() and existing != existing.parent: existing = existing.parent resolved_parent = os.path.realpath(existing) if not any(is_within(resolved_parent, base) for base in self.allowed_bases()): raise RestoreError("unsafe-path", f"{target} resolves outside the allowed roots") if os.path.islink(normalized): raise RestoreError("unsafe-path", f"{target} is a symlink") return normalized # planning def _candidate_files(self, c: Criteria) -> list[dict[str, Any]]: sql = "SELECT * FROM files WHERE 1=1" params: list[Any] = [] if c.root_id is not None: sql += " AND root_id = ?" params.append(c.root_id) if c.project_id is not None: sql += " AND project_id = ?" params.append(c.project_id) if c.paths: sql += f" AND path IN ({','.join('?' * len(c.paths))})" params.extend(c.paths) rows = self.db.all(sql, params) out = [] for row in rows: path = row["path"] if c.globs and not any(fnmatch.fnmatch(path, g) for g in c.globs): continue if any(fnmatch.fnmatch(path, g) for g in c.exclude_globs): continue if c.extensions and not any( path.endswith(e if e.startswith(".") else "." + e) for e in c.extensions ): continue if c.changed_since is not None or c.changed_by_source: sql2 = "SELECT 1 FROM versions WHERE file_id = ?" p2: list[Any] = [row["id"]] if c.changed_since is not None: sql2 += " AND captured_at > ?" p2.append(c.changed_since) if c.changed_by_source: sql2 += " AND source = ?" p2.append(c.changed_by_source) if not self.db.one(sql2 + " LIMIT 1", p2): continue out.append(row) return out def _destination(self, file: dict[str, Any], target_dir: str | None) -> str: if not target_dir: return file["path"] root = self.db.root(file["root_id"]) if file["root_id"] else None base = root["path"] if root else str(Path(file["path"]).parent) return os.path.join(target_dir, os.path.relpath(file["path"], base)) def plan(self, c: Criteria) -> dict[str, Any]: if not (c.paths or c.globs or c.root_id or c.project_id or c.changed_by_source): raise RestoreError("criteria-too-broad", "give at least one of paths, globs, root_id, project_id, changed_by_source") as_of = c.as_of if c.as_of is not None else time.time() target_dir = self.safe_target(os.path.normpath(c.target_dir)) if c.target_dir else None actions = [] for file in self._candidate_files(c): version = self.db.version_at(file["id"], as_of) destination = self._destination(file, target_dir) disk_digest = _disk_digest(destination) item: dict[str, Any] = { "file_id": file["id"], "path": file["path"], "destination": destination, "version_id": version["id"] if version else None, } if version is None: if c.remove_newer_files and disk_digest is not None and not target_dir: item["action"] = "remove" else: item["action"] = "skip-not-existing-at-as-of" elif disk_digest == version["blob_sha256"]: item["action"] = "unchanged" elif disk_digest is None: item["action"] = "create" else: item["action"] = {"overwrite": "overwrite", "skip": "skip-conflict", "rename": "write-renamed", "fail": "conflict"}[c.conflict] actions.append(item) plan_id = secrets.token_hex(8) now = time.time() summary: dict[str, int] = {} for a in actions: summary[a["action"]] = summary.get(a["action"], 0) + 1 plan = {"plan_id": plan_id, "as_of": as_of, "expires_at": now + PLAN_TTL_SECONDS, "summary": summary, "actions": actions} self.db.execute( "INSERT INTO restore_plans(id, created_at, expires_at, criteria_json, plan_json, status) " "VALUES (?, ?, ?, ?, ?, 'planned')", (plan_id, now, now + PLAN_TTL_SECONDS, json.dumps(c.__dict__), json.dumps(plan)), ) return plan def get_plan(self, plan_id: str) -> dict[str, Any]: row = self.db.one("SELECT * FROM restore_plans WHERE id = ?", (plan_id,)) if row is None: raise RestoreError("not-found", f"restore plan {plan_id} does not exist") plan = json.loads(row["plan_json"]) plan["status"] = row["status"] return plan # execution async def execute(self, plan_id: str) -> dict[str, Any]: row = self.db.one("SELECT * FROM restore_plans WHERE id = ?", (plan_id,)) if row is None: raise RestoreError("not-found", f"restore plan {plan_id} does not exist") if row["status"] != "planned": raise RestoreError("plan-used", f"restore plan {plan_id} is {row['status']}") if row["expires_at"] < time.time(): raise RestoreError("plan-expired", f"restore plan {plan_id} expired; create a new one") plan = json.loads(row["plan_json"]) if any(a["action"] == "conflict" for a in plan["actions"]): raise RestoreError("conflict", "plan contains conflicts and conflict policy is 'fail'") self.db.execute("UPDATE restore_plans SET status = 'executing' WHERE id = ?", (plan_id,)) results = [] for action in plan["actions"]: try: results.append(await self._apply(action)) except (OSError, RestoreError) as exc: results.append({**action, "result": "error", "error": str(exc)}) self.db.execute("UPDATE restore_plans SET status = 'executed' WHERE id = ?", (plan_id,)) self.db.audit("restore.execute", plan_id=plan_id, count=len(results)) return {"plan_id": plan_id, "results": results} async def _blob(self, sha256: str) -> bytes: if self.fetch_blob is not None: return await self.fetch_blob(sha256) return self.blobs.get(sha256) async def _apply(self, action: dict[str, Any]) -> dict[str, Any]: kind = action["action"] if kind in ("unchanged", "skip-conflict", "skip-not-existing-at-as-of"): return {**action, "result": "skipped"} destination = self.safe_target(action["destination"]) if kind == "remove": self._snapshot_current(destination) os.unlink(destination) self.repo.mark_deleted(destination) return {**action, "result": "removed"} version = self.db.version(action["version_id"]) try: content = await self._blob(version["blob_sha256"]) except (FileNotFoundError, ValueError) as exc: raise RestoreError("unrestorable", f"blob {version['blob_sha256']} is gone or corrupt") from exc if dbsafe.is_sqlite_image(content): # Never write an unverified database back to disk. try: dbsafe.verify_sqlite_bytes(content) except dbsafe.DatabaseUnsafe as exc: raise RestoreError("unrestorable", f"stored version fails integrity check: {exc.reason}") if kind == "write-renamed": stem, ext = os.path.splitext(destination) destination = f"{stem}.restored-{time.strftime('%Y%m%d-%H%M%S')}{ext}" result = self.write(destination, content, version["id"]) return {**action, "result": "restored", "written_to": destination, **result} def _snapshot_current(self, path: str) -> int | None: try: content, st = self.repo.read_file(path) except (OSError, Exception): return None if self.repo.filters.size_reason(len(content)): return None return self.repo.commit(path, content, "restore", "pre-restore", st)["version_id"] def write(self, destination: str, content: bytes, version_id: int) -> dict[str, Any]: """Atomically write content after snapshotting what is there now.""" destination = self.safe_target(destination) pre = self._snapshot_current(destination) if os.path.exists(destination) else None Path(destination).parent.mkdir(parents=True, exist_ok=True) mode = os.stat(destination).st_mode & 0o7777 if os.path.exists(destination) else 0o644 tmp = f"{destination}.versiond-{secrets.token_hex(4)}.tmp" fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, mode) try: with os.fdopen(fd, "wb") as fh: fh.write(content) fh.flush() os.fsync(fh.fileno()) os.replace(tmp, destination) except BaseException: if os.path.exists(tmp): os.unlink(tmp) raise st = os.stat(destination) committed = self.repo.commit( destination, content, "restore", f"restored from version {version_id}", st ) return {"pre_restore_version_id": pre, "new_version_id": committed["version_id"]} def _disk_digest(path: str) -> str | None: try: with open(path, "rb") as fh: return sha256(fh.read()) except OSError: return None