diff --git a/skillopt_sleep/__main__.py b/skillopt_sleep/__main__.py index e3b7794c6..02ce16b55 100644 --- a/skillopt_sleep/__main__.py +++ b/skillopt_sleep/__main__.py @@ -40,13 +40,18 @@ from skillopt_sleep.staging import ( StagingError, adopt_skills, + adopted_skill_names, + has_adopted_legacy, has_pending_staged_managed, json_safe, + latest_adopted_staging, latest_staging, pending_staged_skills, + revert_skills, staged_skills, ) from skillopt_sleep.staging import adopt as adopt_staging +from skillopt_sleep.staging import revert as revert_staging from skillopt_sleep.state import SleepState from skillopt_sleep.tasks_file import load_tasks_file, make_tasks_payload, write_tasks_file @@ -564,12 +569,14 @@ def cmd_status(args) -> int: all_skills = [] skills = [] has_managed = False + revertable = latest_adopted_staging(project) info = { "night": state.night, "state_path": cfg.state_path, "project": project, "history_tail": state.data.get("history", [])[-5:], "latest_staging": latest, + "revertable_staging": revertable, "slow_memory_chars": len(state.slow_memory), "staged_skills": [r.get("skill_name", "") for r in skills], "adopted_skills": [ @@ -586,6 +593,8 @@ def cmd_status(args) -> int: else: print(f"[sleep] nights so far: {state.night}") print(f"[sleep] project: {project}") + if revertable: + print(f"[sleep] revertable adoption: {_display_value(revertable)}") if latest: print(f"[sleep] latest staged proposal: {_display_value(latest)}") if staging_error: @@ -779,6 +788,116 @@ def fail(code: int, kind: str, message: str, **extra: Any) -> int: return 0 +def cmd_revert(args) -> int: + cfg = _cfg_from_args(args) + project = cfg.get("invoked_project") or os.getcwd() + target = args.staging or latest_adopted_staging(project) + + def fail(code: int, kind: str, message: str, **extra: Any) -> int: + safe_message = _display_value(message) + if args.json: + payload = { + "ok": False, + "error": kind, + "message": safe_message, + "staging_dir": _display_value(target or ""), + } + payload.update(_redact_deep(extra)) + print(json.dumps(payload, ensure_ascii=False, indent=2)) + else: + print(safe_message) + return code + + raw_selected = list(getattr(args, "skills", None) or []) + if any(not str(name).strip() for name in raw_selected): + return fail( + 2, + "invalid_selection", + "[sleep] --skill names must be non-empty.", + ) + selected = [str(name).strip() for name in raw_selected] + revert_all = bool(getattr(args, "all_skills", False)) + revert_legacy = bool(getattr(args, "legacy", False)) + if sum((bool(selected), revert_all, revert_legacy)) > 1: + return fail( + 2, + "invalid_selection", + "[sleep] use exactly one of --skill, --all-skills, or --legacy.", + ) + + if not target: + return fail(1, "no_staging", "[sleep] nothing to revert.") + if not os.path.isdir(target): + return fail(1, "no_staging", f"[sleep] staging directory not found: {_display_value(target)}.") + + has_legacy = has_adopted_legacy(target) + adopted_names = adopted_skill_names(target) + + if not revert_legacy and not selected and not revert_all: + if has_legacy and adopted_names: + return fail( + 2, + "selection_required", + "[sleep] both per-skill and legacy proposals were adopted from this night; " + "pass --skill NAME or --all-skills, or use --legacy.", + ) + if not has_legacy and not adopted_names: + return fail(1, "nothing_to_revert", f"[sleep] nothing to revert from {_display_value(target)}.") + if has_legacy: + revert_legacy = True + else: + revert_all = True + + if revert_legacy: + try: + updated = revert_staging(target) + except StagingError as exc: + return fail(2, "revert_refused", f"[sleep] revert refused: {exc}") + except (OSError, ValueError, KeyError, json.JSONDecodeError) as exc: + return fail(1, "revert_failed", f"[sleep] revert failed: {exc}") + if args.json: + print(json.dumps({ + "ok": True, + "staging_dir": target, + "mode": "legacy", + "reverted_skills": [], + "updated_paths": updated, + }, ensure_ascii=False, indent=2)) + else: + print(f"[sleep] reverted managed proposal from {_display_value(target)}") + for path in updated: + print(f" <- {_display_value(path)}") + if not updated: + print("[sleep] (no adopted managed changes to revert)") + return 0 + + names = selected if not revert_all else None + try: + receipts = revert_skills(target, names) + except StagingError as exc: + return fail(2, "revert_refused", f"[sleep] revert refused: {exc}") + except OSError as exc: + return fail(1, "revert_failed", f"[sleep] revert failed: {exc}") + if args.json: + print(json.dumps(json_safe({ + "ok": True, + "staging_dir": target, + "mode": "skills", + "reverted_skills": [receipt.__dict__ for receipt in receipts], + "updated_paths": [receipt.live_skill_path for receipt in receipts], + }), ensure_ascii=False, indent=2)) + else: + print(f"[sleep] reverted from {_display_value(target)}") + for receipt in receipts: + print( + f" <- {_display_value(receipt.skill_name)}: " + f"{_display_value(receipt.live_skill_path)}" + ) + if not receipts: + print("[sleep] (no skills to revert in the selection)") + return 0 + + def cmd_harvest(args) -> int: cfg = _cfg_from_args(args) session_limit = cfg.get("max_sessions_per_night", 0) or cfg.get("max_tasks_per_night", 40) * 3 @@ -872,6 +991,21 @@ def main(argv=None) -> int: "--legacy", action="store_true", help="adopt only the staged managed skill/memory proposal", ) + p_revert = sub.add_parser("revert", help="revert an adopted proposal using its backups") + _add_common(p_revert) + p_revert.add_argument("--staging", default="", help="specific staging dir") + p_revert.add_argument( + "--skill", action="append", default=[], dest="skills", + help="revert this adopted skill (repeatable)", + ) + p_revert.add_argument( + "--all-skills", action="store_true", dest="all_skills", + help="revert every adopted per-skill proposal", + ) + p_revert.add_argument( + "--legacy", action="store_true", + help="revert only the adopted managed skill/memory proposal", + ) p_harvest = sub.add_parser("harvest", help="debug: show mined tasks") _add_common(p_harvest) p_harvest.add_argument("--output", default="", help="write mined tasks JSON for review") @@ -905,6 +1039,8 @@ def main(argv=None) -> int: return cmd_status(args) if args.cmd == "adopt": return cmd_adopt(args) + if args.cmd == "revert": + return cmd_revert(args) if args.cmd == "harvest": return cmd_harvest(args) if args.cmd == "schedule": diff --git a/skillopt_sleep/staging.py b/skillopt_sleep/staging.py index 1d2a40228..4a8d895d7 100644 --- a/skillopt_sleep/staging.py +++ b/skillopt_sleep/staging.py @@ -1184,6 +1184,17 @@ class _TransactionTarget: _WAL_FILENAME = ".adopt-transaction.json" _WAL_VERSION = 2 +_REVERT_WAL_FILENAME = ".revert-transaction.json" +_REVERT_WAL_VERSION = 1 + + +def _wal_path(staging_dir: str) -> str: + return os.path.join(staging_dir, _WAL_FILENAME) + + +def _revert_wal_path(staging_dir: str) -> str: + return os.path.join(staging_dir, _REVERT_WAL_FILENAME) + def _bytes_sha256(data: Optional[bytes]) -> str: return hashlib.sha256(data).hexdigest() if data is not None else "" @@ -1666,9 +1677,9 @@ def _read_existing_receipts( def pending_staged_skills(staging_dir: str) -> List[Dict[str, Any]]: """Return fan-out rows not already recorded in the validated adoption ledger.""" staging_dir = _canonical_staging_dir(staging_dir) - if os.path.lexists(_wal_path(staging_dir)): + if os.path.lexists(_wal_path(staging_dir)) or os.path.lexists(_revert_wal_path(staging_dir)): raise StagingError( - "an interrupted adoption must be recovered before listing pending skills" + "an interrupted adoption or rollback must be recovered before listing pending skills" ) rows = staged_skills(staging_dir) receipt_path = os.path.join(staging_dir, "adopted_skills.json") @@ -1682,9 +1693,9 @@ def pending_staged_skills(staging_dir: str) -> List[Dict[str, Any]]: def has_staged_managed(staging_dir: str) -> bool: """Return whether a validated managed skill or memory proposal is present.""" staging_dir = _canonical_staging_dir(staging_dir) - if os.path.lexists(_wal_path(staging_dir)): + if os.path.lexists(_wal_path(staging_dir)) or os.path.lexists(_revert_wal_path(staging_dir)): raise StagingError( - "an interrupted adoption must be recovered before listing managed proposals" + "an interrupted adoption or rollback must be recovered before listing managed proposals" ) return bool(_legacy_rows(_load_manifest(staging_dir))) @@ -1692,9 +1703,9 @@ def has_staged_managed(staging_dir: str) -> bool: def has_pending_staged_managed(staging_dir: str) -> bool: """Return whether a validated managed proposal target remains unadopted.""" staging_dir = _canonical_staging_dir(staging_dir) - if os.path.lexists(_wal_path(staging_dir)): + if os.path.lexists(_wal_path(staging_dir)) or os.path.lexists(_revert_wal_path(staging_dir)): raise StagingError( - "an interrupted adoption must be recovered before listing managed proposals" + "an interrupted adoption or rollback must be recovered before listing managed proposals" ) rows = _legacy_rows(_load_manifest(staging_dir)) if not rows: @@ -1774,9 +1785,6 @@ def _exclusive_create_locks(paths: Sequence[str]): pass -def _wal_path(staging_dir: str) -> str: - return os.path.join(staging_dir, _WAL_FILENAME) - def _target_wal_row(target: _TransactionTarget) -> Dict[str, Any]: created_dir_ids = target.created_dir_ids or tuple( @@ -2531,6 +2539,185 @@ def _verify_published_targets( ) +_LINEAGE_FILENAME = ".adoption_lineage.json" +_LINEAGE_LOCK_FILENAME = ".adoption_lineage.lock" + + +def _lineage_path(staging_root_dir: str) -> str: + return os.path.join(staging_root_dir, _LINEAGE_FILENAME) + + +def _lineage_lock_path(staging_root_dir: str) -> str: + return os.path.join(staging_root_dir, _LINEAGE_LOCK_FILENAME) + + +@contextmanager +def _lineage_lock(staging_root_dir: str): + lock_path = _lineage_lock_path(staging_root_dir) + with _exclusive_create_locks([lock_path]): + yield + + +def _read_lineage(staging_root_dir: str) -> List[Dict[str, Any]]: + path = _lineage_path(staging_root_dir) + if not os.path.isfile(path) or _is_link_or_junction(path): + return [] + try: + with open(path, "r", encoding="utf-8") as f: + data = json.load(f) + if isinstance(data, dict) and isinstance(data.get("entries"), list): + return data["entries"] + except Exception: + pass + return [] + + +def _write_lineage(staging_root_dir: str, entries: List[Dict[str, Any]]) -> None: + path = _lineage_path(staging_root_dir) + payload = json.dumps({"version": 1, "entries": entries}, ensure_ascii=False, indent=2) + _write_atomic(path, payload, create_parents=False) + + +def _ensure_lineage_synchronized(staging_root_dir: str) -> List[Dict[str, Any]]: + entries = _read_lineage(staging_root_dir) + recorded_adoptions: set[tuple[str, str]] = set() + active_per_target: Dict[str, List[str]] = {} + + for entry in entries: + action = entry.get("action") + sdir = entry.get("staging_dir") + if not sdir: + continue + sdir_key = _path_identity_key(sdir) + targets = entry.get("targets", []) + if action == "adopt": + for t in targets: + recorded_adoptions.add((sdir_key, t)) + active_per_target.setdefault(t, []).append(sdir) + elif action == "revert": + for t in targets: + stack = active_per_target.setdefault(t, []) + for i in range(len(stack) - 1, -1, -1): + if _path_identity_key(stack[i]) == sdir_key: + stack.pop(i) + break + + needs_update = False + if os.path.isdir(staging_root_dir) and not _is_link_or_junction(staging_root_dir): + subs = [] + for entry in os.listdir(staging_root_dir): + if not _STAGING_DIR_RE.fullmatch(entry): + continue + path = os.path.join(staging_root_dir, entry) + if _is_link_or_junction(path) or not os.path.isdir(path): + continue + manifest_path = os.path.join(path, "manifest.json") + if _is_link_or_junction(manifest_path) or not os.path.isfile(manifest_path): + continue + subs.append(path) + subs.sort(key=_staging_order) + + for sdir in subs: + sdir_key = _path_identity_key(sdir) + s_receipt = os.path.join(sdir, "adopted_skills.json") + if os.path.isfile(s_receipt) and not _is_link_or_junction(s_receipt): + try: + rows, _, _, _ = _read_existing_receipts(s_receipt, sdir) + unrecorded_skills = [] + for r in rows: + tkey = f"skill:{r['skill_name']}" + if (sdir_key, tkey) not in recorded_adoptions: + unrecorded_skills.append(tkey) + recorded_adoptions.add((sdir_key, tkey)) + active_per_target.setdefault(tkey, []).append(sdir) + if unrecorded_skills: + entries.append({ + "id": len(entries) + 1, + "action": "adopt", + "timestamp": os.path.getmtime(s_receipt), + "staging_dir": sdir, + "staging_name": os.path.basename(sdir), + "kind": "skills", + "targets": unrecorded_skills, + }) + needs_update = True + except Exception: + pass + + l_receipt = os.path.join(sdir, "adopted_legacy.json") + if os.path.isfile(l_receipt) and not _is_link_or_junction(l_receipt): + try: + l_rows, _, _, _ = _read_legacy_receipts(l_receipt, sdir) + unrecorded_legacy = [] + for r in l_rows: + tkey = f"legacy:{r['target']}" + if (sdir_key, tkey) not in recorded_adoptions: + unrecorded_legacy.append(tkey) + recorded_adoptions.add((sdir_key, tkey)) + active_per_target.setdefault(tkey, []).append(sdir) + if unrecorded_legacy: + entries.append({ + "id": len(entries) + 1, + "action": "adopt", + "timestamp": os.path.getmtime(l_receipt), + "staging_dir": sdir, + "staging_name": os.path.basename(sdir), + "kind": "legacy", + "targets": unrecorded_legacy, + }) + needs_update = True + except Exception: + pass + + if needs_update or (not os.path.exists(_lineage_path(staging_root_dir)) and entries): + _write_lineage(staging_root_dir, entries) + return entries + + +def _record_adoption_in_lineage( + staging_dir: str, *, kind: str, target_keys: Sequence[str] +) -> None: + if not target_keys: + return + staging_root_dir = os.path.dirname(staging_dir) + if not os.path.isdir(staging_root_dir) or _is_link_or_junction(staging_root_dir): + return + with _lineage_lock(staging_root_dir): + entries = _ensure_lineage_synchronized(staging_root_dir) + entries.append({ + "id": len(entries) + 1, + "action": "adopt", + "timestamp": time.time(), + "staging_dir": staging_dir, + "staging_name": os.path.basename(staging_dir), + "kind": kind, + "targets": list(target_keys), + }) + _write_lineage(staging_root_dir, entries) + + +def _record_revert_in_lineage( + staging_dir: str, *, kind: str, target_keys: Sequence[str] +) -> None: + if not target_keys: + return + staging_root_dir = os.path.dirname(staging_dir) + if not os.path.isdir(staging_root_dir) or _is_link_or_junction(staging_root_dir): + return + with _lineage_lock(staging_root_dir): + entries = _ensure_lineage_synchronized(staging_root_dir) + entries.append({ + "id": len(entries) + 1, + "action": "revert", + "timestamp": time.time(), + "staging_dir": staging_dir, + "staging_name": os.path.basename(staging_dir), + "kind": kind, + "targets": list(target_keys), + }) + _write_lineage(staging_root_dir, entries) + + def _execute_transaction_locked( staging_dir: str, *, @@ -2690,6 +2877,20 @@ def _execute_transaction_locked( _verify_published_targets(targets, published_file_ids) # WAL unlink + parent fsync is the transaction commit point. _remove_transaction_wal(staging_dir, expected_wal=wal) + try: + target_keys = [] + for t in targets: + if t.key.startswith("skill '") and t.key.endswith("'"): + target_keys.append(f"skill:{t.key[7:-1]}") + elif t.key.startswith("legacy "): + target_keys.append(f"legacy:{t.key[7:]}") + else: + target_keys.append(t.key) + _record_adoption_in_lineage( + staging_dir, kind=kind, target_keys=target_keys + ) + except Exception: + pass except BaseException as primary: recovery_errors = _recover_transaction_locked( staging_dir, @@ -2713,16 +2914,416 @@ def _execute_transaction_locked( raise +def _clean_empty_backup_dir(path: str, staging_dir: str) -> None: + """Prune empty parent directories inside staging_dir/backup after backup removal.""" + try: + backup_root = os.path.realpath(os.path.join(staging_dir, "backup")) + curr = os.path.realpath(path) + while curr != backup_root and curr.startswith(backup_root + os.sep): + if os.path.isdir(curr) and not os.listdir(curr): + os.rmdir(curr) + curr = os.path.dirname(curr) + else: + break + except Exception: + pass + + +@dataclass +class _RevertTarget: + key: str + target_key: str + live_path: str + expected_realpath: str + expected_basename: str + current_sha256: str + target_sha256: str + proposal_bytes: bytes + target_bytes: Optional[bytes] + target_mode: Optional[int] + backup_path: str + backup_sha256: str + + +def _revert_target_wal_row(target: _RevertTarget) -> Dict[str, Any]: + return { + "key": target.key, + "target_key": target.target_key, + "live_path": target.live_path, + "expected_realpath": target.expected_realpath, + "expected_basename": target.expected_basename, + "current_sha256": target.current_sha256, + "target_sha256": target.target_sha256, + "proposal_b64": _b64(target.proposal_bytes), + "target_b64": _b64(target.target_bytes), + "target_mode": target.target_mode, + "backup_path": target.backup_path, + "backup_sha256": target.backup_sha256, + } + + +def _revert_transaction_wal( + *, + kind: str, + staging_dir: str, + targets: Sequence[_RevertTarget], + receipt_path: str, + receipt_original: Optional[bytes], + receipt_mode: Optional[int], + receipt_file_id: Optional[tuple[int, int]], + receipt_after: bytes, +) -> Dict[str, Any]: + return { + "version": _REVERT_WAL_VERSION, + "kind": kind, + "staging_dir": staging_dir, + "targets": [_revert_target_wal_row(t) for t in targets], + "receipt": { + "path": receipt_path, + "original_b64": _b64(receipt_original), + "original_mode": receipt_mode, + "original_file_id": ( + list(receipt_file_id) if receipt_file_id is not None else None + ), + "proposed_b64": _b64(receipt_after), + "proposed_sha256": hashlib.sha256(receipt_after).hexdigest(), + }, + } + + +def _read_revert_wal(staging_dir: str) -> Optional[Dict[str, Any]]: + path = _revert_wal_path(staging_dir) + if not os.path.lexists(path): + return None + _remove_private_temp_aliases(path) + try: + info = os.lstat(path) + except OSError as exc: + raise StagingError(f"cannot inspect rollback recovery journal: {exc}") from exc + if ( + _is_link_or_junction(path) + or not stat.S_ISREG(info.st_mode) + or info.st_nlink != 1 + ): + raise StagingError(f"rollback recovery journal is unsafe: {path}") + try: + raw, _mode, _file_id = _file_snapshot(path) + if raw is None: + raise StagingError("rollback recovery journal disappeared") + payload = json.loads(raw.decode("utf-8")) + except Exception as exc: + raise StagingError(f"cannot read rollback recovery journal: {exc}") from exc + if ( + not isinstance(payload, dict) + or set(payload) != {"version", "kind", "staging_dir", "targets", "receipt"} + or payload.get("version") != _REVERT_WAL_VERSION + ): + raise StagingError("rollback recovery journal has an unsupported format") + return payload + + +def _write_revert_wal(staging_dir: str, wal: Dict[str, Any]) -> None: + path = _revert_wal_path(staging_dir) + if os.path.lexists(path): + raise StagingError( + f"a rollback recovery journal already exists; recover it first: {path}" + ) + payload = json.dumps(wal, ensure_ascii=False, indent=2).encode("utf-8") + _write_new_bytes(path, payload, mode=0o600) + + +def _remove_revert_wal(staging_dir: str) -> None: + path = _revert_wal_path(staging_dir) + if os.path.lexists(path): + _remove_private_temp_aliases(path) + info = os.lstat(path) + if ( + _is_link_or_junction(path) + or not stat.S_ISREG(info.st_mode) + or info.st_nlink != 1 + ): + raise StagingError(f"rollback recovery journal is unsafe: {path}") + _unlink_fsync(path) + + +def _decode_revert_wal_targets( + staging_dir: str, wal: Dict[str, Any] +) -> List[_RevertTarget]: + raw_targets = wal.get("targets") + if not isinstance(raw_targets, list) or not raw_targets: + raise StagingError("rollback recovery journal has no targets") + targets = [] + seen = set() + for row in raw_targets: + if not isinstance(row, dict): + raise StagingError("rollback recovery journal target must be an object") + key = row.get("key") + target_key = row.get("target_key") + if not key or not target_key or target_key in seen: + raise StagingError("rollback recovery journal has invalid target keys") + seen.add(target_key) + live = _safe_live_path(row.get("live_path")) + expected_realpath = _safe_live_path(row.get("expected_realpath")) + expected_basename = row.get("expected_basename") + if not live or not expected_realpath or expected_basename not in {"SKILL.md", "CLAUDE.md"}: + raise StagingError("rollback recovery journal has unsafe target paths") + current_sha256 = row.get("current_sha256") + target_sha256 = row.get("target_sha256") + if not _valid_sha256_pin(current_sha256): + raise StagingError("rollback recovery journal has invalid current hash") + if target_sha256 != "" and not _valid_sha256_pin(target_sha256): + raise StagingError("rollback recovery journal has invalid target hash") + proposal_bytes = _from_b64(row.get("proposal_b64"), field="proposal_b64") + if proposal_bytes is None or _bytes_sha256(proposal_bytes) != current_sha256: + raise StagingError("rollback recovery journal has invalid proposal bytes") + target_bytes = _from_b64(row.get("target_b64"), field="target_b64") + if _bytes_sha256(target_bytes) != target_sha256: + raise StagingError("rollback recovery journal has invalid target bytes") + target_mode = row.get("target_mode") + if target_mode is not None and (type(target_mode) is not int or target_mode < 0 or target_mode > 0o7777): + raise StagingError("rollback recovery journal has invalid file mode") + backup_path = row.get("backup_path") or "" + backup_sha256 = row.get("backup_sha256") or "" + if backup_path and not _path_is_within(backup_path, staging_dir): + raise StagingError("rollback recovery journal backup escapes staging") + targets.append(_RevertTarget( + key=key, + target_key=target_key, + live_path=live, + expected_realpath=expected_realpath, + expected_basename=expected_basename, + current_sha256=current_sha256, + target_sha256=target_sha256, + proposal_bytes=proposal_bytes, + target_bytes=target_bytes, + target_mode=target_mode, + backup_path=backup_path, + backup_sha256=backup_sha256, + )) + return targets + + +def _execute_revert_transaction_locked( + staging_dir: str, + *, + kind: str, + targets: Sequence[_RevertTarget], + receipt_path: str, + receipt_original: Optional[bytes], + receipt_mode: Optional[int], + receipt_file_id: Optional[tuple[int, int]], + receipt_after: bytes, +) -> None: + wal = _revert_transaction_wal( + kind=kind, + staging_dir=staging_dir, + targets=targets, + receipt_path=receipt_path, + receipt_original=receipt_original, + receipt_mode=receipt_mode, + receipt_file_id=receipt_file_id, + receipt_after=receipt_after, + ) + _write_revert_wal(staging_dir, wal) + try: + for target in targets: + _adopt_target_ok( + target.key, + target.live_path, + target.expected_realpath, + expected_basename=target.expected_basename, + ) + current, current_mode, _ = _file_snapshot(target.live_path) + current_sha = _bytes_sha256(current) + if current_sha == target.target_sha256: + continue + if current_sha != target.current_sha256: + raise StagingError( + f"live target for {target.key} changed during rollback" + ) + if target.target_bytes is None: + if current is not None: + _unlink_fsync(target.live_path) + else: + _write_atomic_bytes( + target.live_path, + target.target_bytes, + create_parents=False, + mode=target.target_mode, + ) + verified, verified_mode, _ = _file_snapshot(target.live_path) + if _bytes_sha256(verified) != target.target_sha256: + raise StagingError( + f"restoring {target.key} failed sha256 verification" + ) + if target.target_bytes is not None and not _modes_match(verified_mode, target.target_mode): + raise StagingError( + f"restoring {target.key} failed file mode verification" + ) + + receipt_current, receipt_current_mode, receipt_current_id = ( + _file_snapshot(receipt_path) + ) + if ( + _bytes_sha256(receipt_current) != _bytes_sha256(receipt_original) + or receipt_current_id != receipt_file_id + or not _modes_match(receipt_current_mode, receipt_mode) + ): + raise StagingError("adoption receipt changed during rollback") + + _write_atomic( + receipt_path, + receipt_after.decode("utf-8"), + create_parents=False, + ) + check_receipt, _, _ = _file_snapshot(receipt_path) + expected_receipt_sha = hashlib.sha256(receipt_after).hexdigest() + if _bytes_sha256(check_receipt) != expected_receipt_sha: + raise StagingError("adoption receipt changed during publication") + + for target in targets: + if target.backup_path and os.path.lexists(target.backup_path): + _unlink_fsync(target.backup_path) + _clean_empty_backup_dir( + os.path.dirname(target.backup_path), staging_dir + ) + + try: + _record_revert_in_lineage( + staging_dir, + kind=kind, + target_keys=[target.target_key for target in targets], + ) + except Exception: + pass + + _remove_revert_wal(staging_dir) + + except BaseException as primary: + recovery_errors = _recover_revert_transaction_locked(staging_dir, wal) + if recovery_errors: + if not os.path.lexists(_revert_wal_path(staging_dir)): + try: + _write_revert_wal(staging_dir, wal) + except BaseException as exc: + recovery_errors.append( + f"could not restore rollback recovery journal: {type(exc).__name__}: {exc}" + ) + raise StagingRecoveryError( + "rollback failed and automatic recovery was incomplete; " + f"journal retained at {staging_dir}", + primary=primary, + recovery_errors=recovery_errors, + ) from primary + raise + + +def _recover_revert_transaction_locked( + staging_dir: str, wal: Dict[str, Any] +) -> List[str]: + try: + targets = _decode_revert_wal_targets(staging_dir, wal) + except BaseException as exc: + return [f"invalid rollback recovery journal: {type(exc).__name__}: {exc}"] + receipt_info = wal.get("receipt") + if not isinstance(receipt_info, dict): + return ["invalid rollback recovery journal receipt"] + receipt_path = receipt_info.get("path") + proposed_receipt_sha = receipt_info.get("proposed_sha256") + errors: List[str] = [] + + current_receipt, _, _ = _file_snapshot(receipt_path) + current_receipt_sha = _bytes_sha256(current_receipt) + + if current_receipt_sha == proposed_receipt_sha: + for target in targets: + try: + if target.backup_path and os.path.lexists(target.backup_path): + _unlink_fsync(target.backup_path) + _clean_empty_backup_dir( + os.path.dirname(target.backup_path), staging_dir + ) + except BaseException as exc: + errors.append(f"backup removal {target.backup_path}: {type(exc).__name__}: {exc}") + if not errors: + try: + _record_revert_in_lineage( + staging_dir, + kind=wal.get("kind") or "skills", + target_keys=[target.target_key for target in targets], + ) + except Exception: + pass + try: + _remove_revert_wal(staging_dir) + except BaseException as exc: + errors.append(f"journal removal: {type(exc).__name__}: {exc}") + return errors + + for target in reversed(targets): + try: + current_live, current_mode, _ = _file_snapshot(target.live_path) + current_live_sha = _bytes_sha256(current_live) + if current_live_sha == target.target_sha256: + _write_atomic_bytes( + target.live_path, + target.proposal_bytes, + create_parents=False, + mode=target.target_mode, + ) + verified, _, _ = _file_snapshot(target.live_path) + if _bytes_sha256(verified) != target.current_sha256: + raise StagingError( + f"rollback recovery of {target.key} failed sha256 verification" + ) + elif current_live_sha == target.current_sha256: + pass + else: + raise StagingError( + f"rollback recovery conflict for {target.key}: live content changed" + ) + except BaseException as exc: + errors.append(f"target {target.key}: {type(exc).__name__}: {exc}") + + try: + orig_receipt_bytes = _from_b64( + receipt_info.get("original_b64"), field="original_b64" + ) + orig_mode = receipt_info.get("original_mode") + if orig_receipt_bytes is None: + if os.path.lexists(receipt_path): + _unlink_fsync(receipt_path) + else: + _write_atomic_bytes(receipt_path, orig_receipt_bytes, mode=orig_mode) + except BaseException as exc: + errors.append(f"receipt recovery: {type(exc).__name__}: {exc}") + + if errors: + return errors + + try: + _remove_revert_wal(staging_dir) + except BaseException as exc: + errors.append(f"journal removal: {type(exc).__name__}: {exc}") + + return errors + + @contextmanager def _adoption_locks(staging_dir: str, live_paths: Sequence[str]): staging_lock = os.path.join(staging_dir, ".adopt-skills.lock") with _exclusive_create_locks([staging_lock]): wal = _read_transaction_wal(staging_dir) + revert_wal = _read_revert_wal(staging_dir) recovery_paths: List[str] = [] if wal is not None: - recovery_paths = [ + recovery_paths.extend([ target.live_path for target in _decode_wal_targets(staging_dir, wal) - ] + ]) + if revert_wal is not None: + recovery_paths.extend([ + target.live_path for target in _decode_revert_wal_targets(staging_dir, revert_wal) + ]) all_paths = list(live_paths) + recovery_paths with _exclusive_create_locks(_target_lock_paths(all_paths)): if wal is not None: @@ -2732,28 +3333,48 @@ def _adoption_locks(staging_dir: str, live_paths: Sequence[str]): "cannot recover an interrupted adoption; journal and backups retained", recovery_errors=recovery_errors, ) + if revert_wal is not None: + recovery_errors = _recover_revert_transaction_locked(staging_dir, revert_wal) + if recovery_errors: + raise StagingRecoveryError( + "cannot recover an interrupted rollback; journal retained", + recovery_errors=recovery_errors, + ) yield def _recover_before_manifest(staging_dir: str) -> None: - """Recover a WAL without depending on a still-readable staging manifest.""" + """Recover an adoption or rollback WAL without depending on a still-readable staging manifest.""" if _is_link_or_junction(staging_dir) or not os.path.isdir(staging_dir): raise StagingError(f"staging directory is unsafe: {staging_dir}") staging_lock = os.path.join(staging_dir, ".adopt-skills.lock") with _exclusive_create_locks([staging_lock]): wal = _read_transaction_wal(staging_dir) - if wal is None: + revert_wal = _read_revert_wal(staging_dir) + if wal is None and revert_wal is None: return - targets = _decode_wal_targets(staging_dir, wal) - with _exclusive_create_locks( - _target_lock_paths([target.live_path for target in targets]) - ): - recovery_errors = _recover_transaction_locked(staging_dir, wal) - if recovery_errors: - raise StagingRecoveryError( - "cannot recover an interrupted adoption; journal and backups retained", - recovery_errors=recovery_errors, - ) + recovery_paths = [] + if wal is not None: + targets = _decode_wal_targets(staging_dir, wal) + recovery_paths.extend([t.live_path for t in targets]) + if revert_wal is not None: + revert_targets = _decode_revert_wal_targets(staging_dir, revert_wal) + recovery_paths.extend([t.live_path for t in revert_targets]) + with _exclusive_create_locks(_target_lock_paths(recovery_paths)): + if wal is not None: + recovery_errors = _recover_transaction_locked(staging_dir, wal) + if recovery_errors: + raise StagingRecoveryError( + "cannot recover an interrupted adoption; journal and backups retained", + recovery_errors=recovery_errors, + ) + if revert_wal is not None: + recovery_errors = _recover_revert_transaction_locked(staging_dir, revert_wal) + if recovery_errors: + raise StagingRecoveryError( + "cannot recover an interrupted rollback; journal retained", + recovery_errors=recovery_errors, + ) def adopt_skills( @@ -3159,3 +3780,542 @@ def adopt(staging_dir: str) -> List[str]: receipt_after=receipt_after, ) return updated + + +def adopted_skill_names(staging_dir: str) -> List[str]: + """Names of per-skill proposals adopted from this night and not yet reverted.""" + staging_dir = _canonical_staging_dir(staging_dir) + receipt_path = os.path.join(staging_dir, "adopted_skills.json") + if not os.path.exists(receipt_path) or _is_link_or_junction(receipt_path): + return [] + try: + rows, _, _, _ = _read_existing_receipts(receipt_path, staging_dir) + return [str(r["skill_name"]) for r in rows if r.get("skill_name")] + except Exception: + return [] + + +def has_adopted_legacy(staging_dir: str) -> bool: + """Whether the managed skill/memory pair is adopted and not yet reverted.""" + staging_dir = _canonical_staging_dir(staging_dir) + receipt_path = os.path.join(staging_dir, "adopted_legacy.json") + if not os.path.exists(receipt_path) or _is_link_or_junction(receipt_path): + return False + try: + rows, _, _, _ = _read_legacy_receipts(receipt_path, staging_dir) + return bool(rows) + except Exception: + return False + + +def latest_adopted_staging(project: str) -> Optional[str]: + """The newest staging directory holding an adoption that can still be undone. + + Deliberately not `latest_staging`: the newest proposal on disk may never + have been adopted, and reverting that would restore a backup for a change + the live files never received. + """ + root = staging_root(project) + if _is_link_or_junction(root) or not os.path.isdir(root): + return None + with _lineage_lock(root): + entries = _ensure_lineage_synchronized(root) + active_per_target: Dict[str, List[str]] = {} + for entry in entries: + action = entry.get("action") + sdir = entry.get("staging_dir") + if not sdir: + continue + sdir_key = _path_identity_key(sdir) + targets = entry.get("targets", []) + if action == "adopt": + for t in targets: + active_per_target.setdefault(t, []).append(sdir) + elif action == "revert": + for t in targets: + stack = active_per_target.setdefault(t, []) + for i in range(len(stack) - 1, -1, -1): + if _path_identity_key(stack[i]) == sdir_key: + stack.pop(i) + break + + for entry in reversed(entries): + if entry.get("action") != "adopt": + continue + sdir = entry.get("staging_dir") + if not sdir or not os.path.isdir(sdir) or _is_link_or_junction(sdir): + continue + is_active = any( + stack and _path_identity_key(stack[-1]) == _path_identity_key(sdir) + for stack in active_per_target.values() + ) + if is_active and (has_adopted_legacy(sdir) or adopted_skill_names(sdir)): + return sdir + + subs = [] + for entry in os.listdir(root): + if not _STAGING_DIR_RE.fullmatch(entry): + continue + path = os.path.join(root, entry) + if _is_link_or_junction(path) or not os.path.isdir(path): + continue + manifest_path = os.path.join(path, "manifest.json") + if _is_link_or_junction(manifest_path) or not os.path.isfile(manifest_path): + continue + subs.append(path) + subs.sort(key=_staging_order, reverse=True) + for p in subs: + if has_adopted_legacy(p) or adopted_skill_names(p): + return p + return None + + +def _assert_current_adopted_head( + staging_dir: str, + *, + kind: str, + skill_name: Optional[str] = None, + target_keys: Optional[Sequence[str]] = None, +) -> None: + """Enforce a strict stack model using adoption lineage: only the most + recent adopted staging night for each target can be reverted. + """ + if target_keys is None: + if kind == "skills" and skill_name: + target_keys = [f"skill:{skill_name}"] + elif kind == "legacy": + target_keys = ["legacy:skill", "legacy:memory"] + else: + target_keys = [] + if not target_keys: + return + + staging_root_dir = os.path.dirname(staging_dir) + if not os.path.isdir(staging_root_dir) or _is_link_or_junction(staging_root_dir): + return + + with _lineage_lock(staging_root_dir): + entries = _ensure_lineage_synchronized(staging_root_dir) + active_per_target: Dict[str, List[str]] = {} + for entry in entries: + action = entry.get("action") + sdir = entry.get("staging_dir") + if not sdir: + continue + sdir_key = _path_identity_key(sdir) + targets = entry.get("targets", []) + if action == "adopt": + for t in targets: + active_per_target.setdefault(t, []).append(sdir) + elif action == "revert": + for t in targets: + stack = active_per_target.setdefault(t, []) + for i in range(len(stack) - 1, -1, -1): + if _path_identity_key(stack[i]) == sdir_key: + stack.pop(i) + break + + for target_key in target_keys: + stack = active_per_target.get(target_key, []) + if not stack: + continue + current_head = stack[-1] + if _path_identity_key(current_head) != _path_identity_key(staging_dir): + head_name = os.path.basename(current_head) + this_name = os.path.basename(staging_dir) + raise StagingError( + f"cannot revert {target_key} from older adoption in {this_name!r}: " + f"newer adoption in {head_name!r} must be reverted first" + ) + + +def revert_skills( + staging_dir: str, skill_names: Optional[Sequence[str]] = None +) -> List[AdoptedSkill]: + """Revert adopted skills back to their pre-adoption state using backups.""" + staging_dir = _canonical_staging_dir(staging_dir) + _recover_before_manifest(staging_dir) + receipt_path = os.path.join(staging_dir, "adopted_skills.json") + if not os.path.exists(receipt_path): + return [] + + initial_receipts, receipt_original, receipt_mode, receipt_file_id = ( + _read_existing_receipts(receipt_path, staging_dir) + ) + if not initial_receipts: + return [] + + if skill_names is not None: + wanted_names = {str(name).strip() for name in skill_names if str(name).strip()} + rows = [r for r in initial_receipts if r["skill_name"] in wanted_names] + missing = wanted_names - {r["skill_name"] for r in rows} + if missing: + raise StagingError( + "not adopted from this night: " + ", ".join(repr(m) for m in sorted(missing)) + ) + else: + rows = list(initial_receipts) + + if not rows: + return [] + + all_staged = staged_skills(staging_dir) + manifest_by_name = {row.get("skill_name"): row for row in all_staged} + skill_roots = staged_skill_roots(staging_dir) + + initial_live_paths: List[str] = [] + for row in rows: + name = row["skill_name"] + manifest_row = manifest_by_name.get(name) + if not manifest_row: + raise StagingError(f"receipt for {name!r} not found in staging manifest") + manifest_live = _safe_live_path(manifest_row.get("live_skill_path")) + receipt_live = _safe_live_path(row.get("live_skill_path")) + if ( + not manifest_live + or not receipt_live + or _path_identity_key(manifest_live) != _path_identity_key(receipt_live) + ): + raise StagingError(f"receipt for {name!r} has untrusted live path") + expected_realpath = _safe_live_path(manifest_row.get("live_realpath")) + _adopt_live_target_ok(name, receipt_live, expected_realpath, skill_roots) + initial_live_paths.append(receipt_live) + + with _adoption_locks(staging_dir, initial_live_paths): + current_receipts, current_original, current_mode, current_file_id = ( + _read_existing_receipts(receipt_path, staging_dir) + ) + if ( + current_file_id != receipt_file_id + or current_original != receipt_original + or not _modes_match(current_mode, receipt_mode) + ): + raise StagingError("adoption receipt changed while revert was locking") + + if skill_names is not None: + active_rows = [r for r in current_receipts if r["skill_name"] in wanted_names] + else: + active_rows = list(current_receipts) + + revert_targets: List[_RevertTarget] = [] + receipts_to_return: List[AdoptedSkill] = [] + reverted_names: set[str] = set() + + for row in active_rows: + name = row["skill_name"] + manifest_row = manifest_by_name.get(name) + if not manifest_row: + raise StagingError(f"receipt for {name!r} not found in staging manifest") + + manifest_live = _safe_live_path(manifest_row.get("live_skill_path")) + receipt_live = _safe_live_path(row.get("live_skill_path")) + if ( + not manifest_live + or not receipt_live + or _path_identity_key(manifest_live) != _path_identity_key(receipt_live) + ): + raise StagingError(f"receipt for {name!r} has untrusted live path") + + expected_realpath = _safe_live_path(manifest_row.get("live_realpath")) + _adopt_live_target_ok(name, receipt_live, expected_realpath, skill_roots) + + manifest_proposal_sha = manifest_row.get("sha256") + manifest_baseline_sha = manifest_row.get("live_sha256") or "" + if not _valid_sha256_pin(manifest_proposal_sha): + raise StagingError(f"manifest proposal hash for {name!r} is invalid") + if manifest_baseline_sha != "" and not _valid_sha256_pin(manifest_baseline_sha): + raise StagingError(f"manifest baseline hash for {name!r} is invalid") + + if row.get("sha256_after") != manifest_proposal_sha: + raise StagingError( + f"receipt proposal hash for {name!r} does not match manifest" + ) + if (row.get("sha256_before") or "") != manifest_baseline_sha: + raise StagingError( + f"receipt baseline hash for {name!r} does not match manifest" + ) + + receipt_backup = row.get("backup_path") or "" + if manifest_baseline_sha == "": + if receipt_backup != "": + raise StagingError( + f"receipt for newly created skill {name!r} has unexpected backup" + ) + expected_backup = "" + else: + expected_backup = os.path.join( + staging_dir, "backup", "skills", name, "SKILL.md" + ) + if _path_identity_key(receipt_backup) != _path_identity_key(expected_backup): + raise StagingError( + f"receipt for {name!r} has invalid backup path" + ) + backup_sha = _immutable_backup_sha256(expected_backup, staging_dir) + if backup_sha is None: + raise StagingError( + f"backup for {name!r} is missing; cannot revert" + ) + if backup_sha != manifest_baseline_sha: + raise StagingError( + f"backup for {name!r} changed; cannot safely revert" + ) + + _assert_current_adopted_head( + staging_dir, kind="skills", target_keys=[f"skill:{name}"] + ) + + current, current_mode, _ = _file_snapshot(receipt_live) + current_sha = _bytes_sha256(current) + + if current_sha == manifest_proposal_sha: + if manifest_baseline_sha == "": + target_bytes = None + target_mode = None + else: + with open(expected_backup, "rb") as bf: + target_bytes = bf.read() + if _bytes_sha256(target_bytes) != manifest_baseline_sha: + raise StagingError( + f"backup content for {name!r} does not match hash" + ) + target_mode = current_mode + proposal_bytes = current or b"" + elif current_sha == manifest_baseline_sha: + target_bytes = current + target_mode = current_mode + expected_file = proposal_filename(name) + staged = os.path.join(staging_dir, expected_file) + if os.path.isfile(staged) and not _is_link_or_junction(staged): + with open(staged, "rb") as sf: + proposal_bytes = sf.read() + else: + proposal_bytes = b"" + else: + raise StagingError( + f"live skill for {name!r} changed since adoption; cannot safely revert" + ) + + revert_targets.append(_RevertTarget( + key=f"skill {name!r}", + target_key=f"skill:{name}", + live_path=receipt_live, + expected_realpath=expected_realpath, + expected_basename="SKILL.md", + current_sha256=manifest_proposal_sha, + target_sha256=manifest_baseline_sha, + proposal_bytes=proposal_bytes, + target_bytes=target_bytes, + target_mode=target_mode, + backup_path=expected_backup, + backup_sha256=manifest_baseline_sha, + )) + receipts_to_return.append(AdoptedSkill( + skill_name=name, + live_skill_path=receipt_live, + sha256_before=manifest_proposal_sha, + sha256_after=manifest_baseline_sha, + backup_path="", + )) + reverted_names.add(name) + + remaining_receipts = [ + r for r in current_receipts if r["skill_name"] not in reverted_names + ] + receipt_after = json.dumps( + remaining_receipts, ensure_ascii=False, indent=2 + ).encode("utf-8") + + _execute_revert_transaction_locked( + staging_dir, + kind="skills", + targets=revert_targets, + receipt_path=receipt_path, + receipt_original=receipt_original, + receipt_mode=receipt_mode, + receipt_file_id=receipt_file_id, + receipt_after=receipt_after, + ) + return receipts_to_return + + +def revert(staging_dir: str) -> List[str]: + """Revert adopted legacy managed skills back to their pre-adoption state.""" + staging_dir = _canonical_staging_dir(staging_dir) + _recover_before_manifest(staging_dir) + receipt_path = os.path.join(staging_dir, "adopted_legacy.json") + if not os.path.exists(receipt_path): + return [] + + initial_receipts, receipt_original, receipt_mode, receipt_file_id = ( + _read_legacy_receipts(receipt_path, staging_dir) + ) + if not initial_receipts: + return [] + + manifest = _load_manifest(staging_dir) + legacy_rows = _legacy_rows(manifest) + + initial_live_paths: List[str] = [] + for row in initial_receipts: + label = row.get("target") + manifest_row = legacy_rows.get(label) + if not manifest_row: + raise StagingError(f"legacy {label} not found in staging manifest") + manifest_live = _safe_live_path(manifest_row.get("live_path")) + receipt_live = _safe_live_path(row.get("live_path")) + if ( + not manifest_live + or not receipt_live + or _path_identity_key(manifest_live) != _path_identity_key(receipt_live) + ): + raise StagingError(f"legacy {label} receipt has untrusted live path") + expected_realpath = _safe_live_path(manifest_row.get("live_realpath")) + expected_basename = "SKILL.md" if label == "skill" else "CLAUDE.md" + _adopt_target_ok( + f"legacy {label}", + receipt_live, + expected_realpath, + expected_basename=expected_basename, + ) + initial_live_paths.append(receipt_live) + + with _adoption_locks(staging_dir, initial_live_paths): + current_receipts, current_original, current_mode, current_file_id = ( + _read_legacy_receipts(receipt_path, staging_dir) + ) + if ( + current_file_id != receipt_file_id + or current_original != receipt_original + or not _modes_match(current_mode, receipt_mode) + ): + raise StagingError("legacy adoption receipt changed while revert was locking") + + revert_targets: List[_RevertTarget] = [] + updated: List[str] = [] + + for row in current_receipts: + label = row["target"] + manifest_row = legacy_rows.get(label) + if not manifest_row: + raise StagingError(f"legacy {label} not found in staging manifest") + + manifest_live = _safe_live_path(manifest_row.get("live_path")) + receipt_live = _safe_live_path(row.get("live_path")) + if ( + not manifest_live + or not receipt_live + or _path_identity_key(manifest_live) != _path_identity_key(receipt_live) + ): + raise StagingError(f"legacy {label} receipt has untrusted live path") + + expected_realpath = _safe_live_path(manifest_row.get("live_realpath")) + expected_basename = "SKILL.md" if label == "skill" else "CLAUDE.md" + _adopt_target_ok( + f"legacy {label}", + receipt_live, + expected_realpath, + expected_basename=expected_basename, + ) + + manifest_proposal_sha = manifest_row.get("sha256") + manifest_baseline_sha = manifest_row.get("live_sha256") or "" + if not _valid_sha256_pin(manifest_proposal_sha): + raise StagingError(f"manifest proposal hash for legacy {label} is invalid") + if manifest_baseline_sha != "" and not _valid_sha256_pin(manifest_baseline_sha): + raise StagingError(f"manifest baseline hash for legacy {label} is invalid") + + if row.get("sha256_after") != manifest_proposal_sha: + raise StagingError( + f"receipt proposal hash for legacy {label} does not match manifest" + ) + if (row.get("sha256_before") or "") != manifest_baseline_sha: + raise StagingError( + f"receipt baseline hash for legacy {label} does not match manifest" + ) + + receipt_backup = row.get("backup_path") or "" + if manifest_baseline_sha == "": + if receipt_backup != "": + raise StagingError( + f"receipt for newly created legacy {label} has unexpected backup" + ) + expected_backup = "" + else: + expected_backup = os.path.join(staging_dir, "backup", expected_basename) + if _path_identity_key(receipt_backup) != _path_identity_key(expected_backup): + raise StagingError( + f"receipt for legacy {label} has invalid backup path" + ) + backup_sha = _immutable_backup_sha256(expected_backup, staging_dir) + if backup_sha is None: + raise StagingError( + f"legacy backup for {label} is missing; cannot revert" + ) + if backup_sha != manifest_baseline_sha: + raise StagingError( + f"legacy backup for {label} changed; cannot safely revert" + ) + + _assert_current_adopted_head( + staging_dir, kind="legacy", target_keys=[f"legacy:{label}"] + ) + + current, current_mode, _ = _file_snapshot(receipt_live) + current_sha = _bytes_sha256(current) + + if current_sha == manifest_proposal_sha: + if manifest_baseline_sha == "": + target_bytes = None + target_mode = None + else: + with open(expected_backup, "rb") as bf: + target_bytes = bf.read() + if _bytes_sha256(target_bytes) != manifest_baseline_sha: + raise StagingError( + f"restoring legacy {label} failed sha256 verification" + ) + target_mode = current_mode + proposal_bytes = current or b"" + elif current_sha == manifest_baseline_sha: + target_bytes = current + target_mode = current_mode + staged = os.path.join(staging_dir, expected_basename) + if os.path.isfile(staged) and not _is_link_or_junction(staged): + with open(staged, "rb") as sf: + proposal_bytes = sf.read() + else: + proposal_bytes = b"" + else: + raise StagingError( + f"legacy {label} changed since adoption; cannot safely revert" + ) + + revert_targets.append(_RevertTarget( + key=f"legacy {label}", + target_key=f"legacy:{label}", + live_path=receipt_live, + expected_realpath=expected_realpath, + expected_basename=expected_basename, + current_sha256=manifest_proposal_sha, + target_sha256=manifest_baseline_sha, + proposal_bytes=proposal_bytes, + target_bytes=target_bytes, + target_mode=target_mode, + backup_path=expected_backup, + backup_sha256=manifest_baseline_sha, + )) + updated.append(receipt_live) + + receipt_after = json.dumps([], ensure_ascii=False, indent=2).encode("utf-8") + + _execute_revert_transaction_locked( + staging_dir, + kind="legacy", + targets=revert_targets, + receipt_path=receipt_path, + receipt_original=receipt_original, + receipt_mode=receipt_mode, + receipt_file_id=receipt_file_id, + receipt_after=receipt_after, + ) + return updated diff --git a/tests/test_sleep_revert.py b/tests/test_sleep_revert.py new file mode 100644 index 000000000..d2f46685e --- /dev/null +++ b/tests/test_sleep_revert.py @@ -0,0 +1,874 @@ +"""Tests for undoing an adoption: `skillopt-sleep revert`.""" +import hashlib +import io +import json +import os +import shutil +import tempfile +import unittest +from contextlib import redirect_stdout + +from skillopt_sleep.__main__ import main +from skillopt_sleep.staging import ( + SkillProposal, + StagingError, + _recover_revert_transaction_locked, + _revert_transaction_wal, + _revert_wal_path, + _RevertTarget, + _write_revert_wal, + adopt, + adopt_skills, + adopted_skill_names, + has_adopted_legacy, + latest_adopted_staging, + revert, + revert_skills, + write_staging, +) +from skillopt_sleep.types import SleepReport + +ORIGINAL = "# original\n" +PROPOSED = "# proposed\n" +MEMORY_ORIGINAL = "# memory original\n" +MEMORY_PROPOSED = "# memory proposed\n" + + +def _sha(text: str) -> str: + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def _write(path: str, text: str) -> None: + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "w", encoding="utf-8") as f: + f.write(text) + + +def _read(path: str) -> str: + with open(path, "r", encoding="utf-8") as f: + return f.read() + + +class TestSleepRevert(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.mkdtemp() + self.project = os.path.join(self.tmp, "project") + os.makedirs(self.project, exist_ok=True) + + def tearDown(self): + shutil.rmtree(self.tmp, ignore_errors=True) + + def test_revert_skills(self): + skill_live = os.path.join(self.project, "skills", "hello", "SKILL.md") + _write(skill_live, ORIGINAL) + + proposal = SkillProposal("hello", PROPOSED, skill_live) + report = SleepReport(night=1, project=self.project) + staging = write_staging( + self.project, + report=report, + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + receipts = adopt_skills(staging) + self.assertEqual(len(receipts), 1) + self.assertEqual(_read(skill_live), PROPOSED) + + reverted = revert_skills(staging) + self.assertEqual(len(reverted), 1) + self.assertEqual(_read(skill_live), ORIGINAL) + + def test_revert_legacy(self): + skill_live = os.path.join(self.project, "SKILL.md") + _write(skill_live, ORIGINAL) + + report = SleepReport(night=1, project=self.project, accepted=True) + staging = write_staging( + self.project, + report=report, + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=skill_live, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + + updated = adopt(staging) + self.assertEqual(len(updated), 1) + self.assertEqual(_read(skill_live), PROPOSED) + + reverted = revert(staging) + self.assertEqual(len(reverted), 1) + self.assertEqual(_read(skill_live), ORIGINAL) + + def test_removes_a_file_adoption_created(self): + skill_live = os.path.join(self.project, "skills", "brand_new", "SKILL.md") + os.makedirs(os.path.dirname(skill_live), exist_ok=True) + self.assertFalse(os.path.exists(skill_live)) + + proposal = SkillProposal("brand_new", PROPOSED, skill_live) + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + adopt_skills(staging) + self.assertTrue(os.path.exists(skill_live)) + self.assertEqual(_read(skill_live), PROPOSED) + + reverted = revert_skills(staging) + self.assertEqual(len(reverted), 1) + self.assertFalse(os.path.exists(skill_live)) + + def test_restores_skill_and_memory_together(self): + skill_live = os.path.join(self.project, "SKILL.md") + memory_live = os.path.join(self.project, "CLAUDE.md") + _write(skill_live, ORIGINAL) + _write(memory_live, MEMORY_ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=MEMORY_PROPOSED, + live_skill_path=skill_live, + live_memory_path=memory_live, + skill_proposals=[], + report_md="", + ) + + adopt(staging) + self.assertEqual(_read(skill_live), PROPOSED) + self.assertEqual(_read(memory_live), MEMORY_PROPOSED) + + revert(staging) + self.assertEqual(_read(skill_live), ORIGINAL) + self.assertEqual(_read(memory_live), MEMORY_ORIGINAL) + + def test_adopt_revert_adopt_roundtrip(self): + """Regression test for Yif-Yang review: adopt -> revert -> adopt must succeed.""" + skill_live = os.path.join(self.project, "skills", "cycle", "SKILL.md") + _write(skill_live, ORIGINAL) + + proposal = SkillProposal("cycle", PROPOSED, skill_live) + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + adopt_skills(staging) + self.assertEqual(_read(skill_live), PROPOSED) + + revert_skills(staging) + self.assertEqual(_read(skill_live), ORIGINAL) + + # Adopting the same staging again must succeed without "immutable backup already exists" error + adopt_skills(staging) + self.assertEqual(_read(skill_live), PROPOSED) + + revert_skills(staging) + self.assertEqual(_read(skill_live), ORIGINAL) + + def test_revert_converges_when_run_twice(self): + skill_live = os.path.join(self.project, "skills", "idempotent", "SKILL.md") + _write(skill_live, ORIGINAL) + + proposal = SkillProposal("idempotent", PROPOSED, skill_live) + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + adopt_skills(staging) + reverted1 = revert_skills(staging) + self.assertEqual(len(reverted1), 1) + + # Second revert converges cleanly and reports nothing left to revert + reverted2 = revert_skills(staging) + self.assertEqual(len(reverted2), 0) + self.assertEqual(_read(skill_live), ORIGINAL) + + def test_never_adopted_is_a_no_op(self): + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + self.assertEqual(revert_skills(staging), []) + self.assertEqual(revert(staging), []) + + def test_refuses_when_live_file_edited_after_adoption(self): + skill_live = os.path.join(self.project, "skills", "edited", "SKILL.md") + _write(skill_live, ORIGINAL) + + proposal = SkillProposal("edited", PROPOSED, skill_live) + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + adopt_skills(staging) + _write(skill_live, "# locally modified by developer\n") + + with self.assertRaises(StagingError) as ctx: + revert_skills(staging) + self.assertIn("changed since adoption", str(ctx.exception)) + self.assertEqual(_read(skill_live), "# locally modified by developer\n") + + def test_refuses_an_unknown_skill_name(self): + skill_live = os.path.join(self.project, "skills", "alpha", "SKILL.md") + _write(skill_live, ORIGINAL) + + proposal = SkillProposal("alpha", PROPOSED, skill_live) + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[proposal], + report_md="", + ) + + adopt_skills(staging) + with self.assertRaises(StagingError) as ctx: + revert_skills(staging, ["nonexistent"]) + self.assertIn("not adopted from this night", str(ctx.exception)) + + def test_reverts_reviewed_subset_and_leaves_rest(self): + a_live = os.path.join(self.project, "skills", "a", "SKILL.md") + b_live = os.path.join(self.project, "skills", "b", "SKILL.md") + _write(a_live, ORIGINAL) + _write(b_live, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[ + SkillProposal("a", PROPOSED, a_live), + SkillProposal("b", PROPOSED, b_live), + ], + report_md="", + ) + + adopt_skills(staging, ["a", "b"]) + self.assertEqual(_read(a_live), PROPOSED) + self.assertEqual(_read(b_live), PROPOSED) + + revert_skills(staging, ["a"]) + self.assertEqual(_read(a_live), ORIGINAL) + self.assertEqual(_read(b_live), PROPOSED) + self.assertEqual(adopted_skill_names(staging), ["b"]) + + def test_per_skill_and_legacy_ledgers_independent(self): + managed = os.path.join(self.project, "SKILL.md") + skill = os.path.join(self.project, "skills", "alpha", "SKILL.md") + _write(managed, ORIGINAL) + _write(skill, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=managed, + live_memory_path=None, + skill_proposals=[SkillProposal("alpha", PROPOSED, skill)], + report_md="", + ) + + adopt(staging) + adopt_skills(staging, ["alpha"]) + + revert(staging) + self.assertEqual(_read(managed), ORIGINAL) + self.assertEqual(_read(skill), PROPOSED) + self.assertFalse(has_adopted_legacy(staging)) + self.assertEqual(adopted_skill_names(staging), ["alpha"]) + + def test_latest_adopted_staging_skips_unadopted_night(self): + """Regression test for Yif-Yang review: bare revert after newer unadopted staging.""" + live = os.path.join(self.project, "SKILL.md") + _write(live, ORIGINAL) + + night1 = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=live, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + adopt(night1) + + # Night 2 staged but never adopted + _write(os.path.join(self.project, "skills", "dummy", "SKILL.md"), ORIGINAL) + night2 = write_staging( + self.project, + report=SleepReport(night=2, project=self.project, accepted=False), + proposed_skill="# night 2 unadopted\n", + proposed_memory=None, + live_skill_path=live, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + + # latest_adopted_staging must return night1, skipping night2! + self.assertEqual(latest_adopted_staging(self.project), night1) + self.assertNotEqual(latest_adopted_staging(self.project), night2) + + revert(night1) + self.assertIsNone(latest_adopted_staging(self.project)) + + def test_history_ordering_stack_model_refuses_older_when_newer_adopted(self): + """Regression test for Yif-Yang review: ABA sequence history ordering.""" + skill_live = os.path.join(self.project, "skills", "ordered", "SKILL.md") + _write(skill_live, "state_O") + + night1 = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("ordered", "state_X", skill_live)], + report_md="", + ) + adopt_skills(night1, ["ordered"]) + self.assertEqual(_read(skill_live), "state_X") + + night2 = write_staging( + self.project, + report=SleepReport(night=2, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("ordered", "state_Y", skill_live)], + report_md="", + ) + adopt_skills(night2, ["ordered"]) + self.assertEqual(_read(skill_live), "state_Y") + + night3 = write_staging( + self.project, + report=SleepReport(night=3, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("ordered", "state_X", skill_live)], + report_md="", + ) + adopt_skills(night3, ["ordered"]) + self.assertEqual(_read(skill_live), "state_X") + + # Attempting to revert night1 when night3 is still adopted must be refused by the stack model + with self.assertRaises(StagingError) as ctx: + revert_skills(night1, ["ordered"]) + self.assertIn("newer adoption", str(ctx.exception)) + + def test_tampered_receipt_external_live_path_refused(self): + """Regression test: tampered receipt with external live_path must be rejected.""" + external_file = os.path.join(self.tmp, "outside_project.md") + _write(external_file, "# critical system file\n") + + skill_live = os.path.join(self.project, "skills", "victim", "SKILL.md") + _write(skill_live, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("victim", PROPOSED, skill_live)], + report_md="", + ) + adopt_skills(staging) + + # Tamper receipt to point to external file + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + receipts[0]["live_skill_path"] = external_file + with open(receipt_path, "w", encoding="utf-8") as f: + json.dump(receipts, f) + + with self.assertRaises(StagingError): + revert_skills(staging) + + # External file was never modified or deleted + self.assertEqual(_read(external_file), "# critical system file\n") + + def test_tampered_receipt_extra_fields_refused(self): + skill_live = os.path.join(self.project, "skills", "extra", "SKILL.md") + _write(skill_live, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("extra", PROPOSED, skill_live)], + report_md="", + ) + adopt_skills(staging) + + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + receipts[0]["injected_field"] = "bad" + with open(receipt_path, "w", encoding="utf-8") as f: + json.dump(receipts, f) + + with self.assertRaises(StagingError): + revert_skills(staging) + + def test_out_of_order_multi_night_adoptions(self): + """Lineage ledger tracks true adoption order across out-of-order adoptions.""" + live_a = os.path.join(self.project, "skills", "skill_a", "SKILL.md") + live_b = os.path.join(self.project, "skills", "skill_b", "SKILL.md") + _write(live_a, "orig_a") + _write(live_b, "orig_b") + + # Night 1 staged with skill_a + night1 = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("skill_a", "prop_a", live_a)], + report_md="", + ) + + # Night 2 staged with skill_b + night2 = write_staging( + self.project, + report=SleepReport(night=2, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("skill_b", "prop_b", live_b)], + report_md="", + ) + + # Adopt night 2 first + adopt_skills(night2, ["skill_b"]) + self.assertEqual(_read(live_b), "prop_b") + self.assertEqual(latest_adopted_staging(self.project), night2) + + # Adopt night 1 second (out of chronological / night order) + adopt_skills(night1, ["skill_a"]) + self.assertEqual(_read(live_a), "prop_a") + # Lineage ledger must identify night1 as the latest adopted staging + self.assertEqual(latest_adopted_staging(self.project), night1) + + # Reverting night1 succeeds and points latest adopted back to night2 + revert_skills(night1, ["skill_a"]) + self.assertEqual(_read(live_a), "orig_a") + self.assertEqual(latest_adopted_staging(self.project), night2) + + # Reverting night2 succeeds and leaves no adopted staging + revert_skills(night2, ["skill_b"]) + self.assertEqual(_read(live_b), "orig_b") + self.assertIsNone(latest_adopted_staging(self.project)) + + def test_whole_set_prevalidation_no_partial_mutation(self): + """Prevalidation covers the entire set before any mutations: corruption in skill 2 leaves skill 1 untouched.""" + live_a = os.path.join(self.project, "skills", "skill_a", "SKILL.md") + live_b = os.path.join(self.project, "skills", "skill_b", "SKILL.md") + _write(live_a, "orig_a") + _write(live_b, "orig_b") + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[ + SkillProposal("skill_a", "prop_a", live_a), + SkillProposal("skill_b", "prop_b", live_b), + ], + report_md="", + ) + adopt_skills(staging, ["skill_a", "skill_b"]) + self.assertEqual(_read(live_a), "prop_a") + self.assertEqual(_read(live_b), "prop_b") + + # Corrupt skill_b's backup file + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + b_receipt = next(r for r in receipts if r["skill_name"] == "skill_b") + backup_b = b_receipt["backup_path"] + _write(backup_b, "corrupted_backup_content") + + # Attempt whole-set revert + with self.assertRaises(StagingError) as ctx: + revert_skills(staging, ["skill_a", "skill_b"]) + self.assertIn("immutable backup for previously adopted skill 'skill_b' changed", str(ctx.exception)) + + # Whole-set prevalidation guarantee: skill_a was NOT reverted + self.assertEqual(_read(live_a), "prop_a") + self.assertEqual(_read(live_b), "prop_b") + + # skill_a's backup was NOT deleted + a_receipt = next(r for r in receipts if r["skill_name"] == "skill_a") + backup_a = a_receipt["backup_path"] + self.assertTrue(os.path.exists(backup_a)) + self.assertEqual(_read(backup_a), "orig_a") + + # Receipt was completely untouched + with open(receipt_path, "r", encoding="utf-8") as f: + receipts_after = json.load(f) + self.assertEqual(len(receipts_after), 2) + + def test_authoritative_manifest_hash_mismatch_refused(self): + """Receipt differing from authoritative manifest sha256_after fails closed before mutation.""" + live_a = os.path.join(self.project, "skills", "skill_a", "SKILL.md") + _write(live_a, "orig_a") + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("skill_a", "prop_a", live_a)], + report_md="", + ) + adopt_skills(staging, ["skill_a"]) + + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + receipts[0]["sha256_after"] = "0" * 64 + with open(receipt_path, "w", encoding="utf-8") as f: + json.dump(receipts, f) + + with self.assertRaises(StagingError) as ctx: + revert_skills(staging, ["skill_a"]) + self.assertIn("does not match manifest", str(ctx.exception)) + self.assertEqual(_read(live_a), "prop_a") + + def test_authoritative_manifest_created_file_unexpected_backup_refused(self): + """Newly created skill with tampered non-null backup_path is rejected by authoritative manifest check.""" + live_new = os.path.join(self.project, "skills", "newbie", "SKILL.md") + os.makedirs(os.path.dirname(live_new), exist_ok=True) + self.assertFalse(os.path.exists(live_new)) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("newbie", "created_content", live_new)], + report_md="", + ) + adopt_skills(staging, ["newbie"]) + self.assertTrue(os.path.exists(live_new)) + + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + fake_backup = os.path.join(staging, ".backups", "skills", "newbie", "SKILL.md") + _write(fake_backup, "fake") + receipts[0]["backup_path"] = fake_backup + with open(receipt_path, "w", encoding="utf-8") as f: + json.dump(receipts, f) + + with self.assertRaises(StagingError) as ctx: + revert_skills(staging, ["newbie"]) + self.assertIn("has an unexpected backup", str(ctx.exception)) + self.assertEqual(_read(live_new), "created_content") + + def test_rollback_wal_backward_recovery_on_crash_before_receipt_commit(self): + """If revert crashed before receipt commit, WAL backward recovery restores live file and preserves backups.""" + live_a = os.path.join(self.project, "skills", "skill_a", "SKILL.md") + _write(live_a, "orig_a") + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("skill_a", "prop_a", live_a)], + report_md="", + ) + adopt_skills(staging, ["skill_a"]) + + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + backup_path = receipts[0]["backup_path"] + + # Simulate crash before receipt commit: + # 1. Target was mutated to reverted state "orig_a" + _write(live_a, "orig_a") + # 2. Revert WAL was written + target = _RevertTarget( + key="skill_a", + target_key="skill:skill_a", + live_path=live_a, + expected_realpath=os.path.realpath(live_a), + expected_basename="SKILL.md", + current_sha256=_sha("prop_a"), + target_sha256=_sha("orig_a"), + proposal_bytes=b"prop_a", + target_bytes=b"orig_a", + target_mode=0o644, + backup_path=backup_path, + backup_sha256=_sha("orig_a"), + ) + receipt_bytes = json.dumps(receipts, ensure_ascii=False, indent=2).encode("utf-8") + wal = _revert_transaction_wal( + kind="skills", + staging_dir=staging, + targets=[target], + receipt_path=receipt_path, + receipt_original=receipt_bytes, + receipt_mode=0o644, + receipt_file_id=(1, 1), + receipt_after=b"", + ) + _write_revert_wal(staging, wal) + + # Trigger recovery + errors = _recover_revert_transaction_locked(staging, wal) + self.assertEqual(errors, []) + + # Backward recovery restored live target back to proposed state! + self.assertEqual(_read(live_a), "prop_a") + # Backup is preserved + self.assertTrue(os.path.exists(backup_path)) + # WAL is removed + self.assertFalse(os.path.exists(_revert_wal_path(staging))) + + # Revert can now cleanly succeed + revert_skills(staging, ["skill_a"]) + self.assertEqual(_read(live_a), "orig_a") + self.assertFalse(os.path.exists(backup_path)) + + def test_rollback_wal_forward_recovery_on_crash_after_receipt_commit(self): + """If revert crashed after receipt commit, WAL forward recovery deletes backups and cleans WAL.""" + live_a = os.path.join(self.project, "skills", "skill_a", "SKILL.md") + _write(live_a, "orig_a") + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project), + proposed_skill=None, + proposed_memory=None, + live_skill_path=None, + live_memory_path=None, + skill_proposals=[SkillProposal("skill_a", "prop_a", live_a)], + report_md="", + ) + adopt_skills(staging, ["skill_a"]) + + receipt_path = os.path.join(staging, "adopted_skills.json") + with open(receipt_path, "r", encoding="utf-8") as f: + receipts = json.load(f) + backup_path = receipts[0]["backup_path"] + + # Simulate crash after receipt commit: + # Live target was reverted + _write(live_a, "orig_a") + # Receipt was committed (empty) + receipt_after = json.dumps([], ensure_ascii=False, indent=2).encode("utf-8") + _write(receipt_path, receipt_after.decode("utf-8")) + + target = _RevertTarget( + key="skill_a", + target_key="skill:skill_a", + live_path=live_a, + expected_realpath=os.path.realpath(live_a), + expected_basename="SKILL.md", + current_sha256=_sha("prop_a"), + target_sha256=_sha("orig_a"), + proposal_bytes=b"prop_a", + target_bytes=b"orig_a", + target_mode=0o644, + backup_path=backup_path, + backup_sha256=_sha("orig_a"), + ) + receipt_bytes = json.dumps(receipts, ensure_ascii=False, indent=2).encode("utf-8") + wal = _revert_transaction_wal( + kind="skills", + staging_dir=staging, + targets=[target], + receipt_path=receipt_path, + receipt_original=receipt_bytes, + receipt_mode=0o644, + receipt_file_id=(1, 1), + receipt_after=receipt_after, + ) + _write_revert_wal(staging, wal) + + # Backup still exists + self.assertTrue(os.path.exists(backup_path)) + + # Trigger recovery + errors = _recover_revert_transaction_locked(staging, wal) + self.assertEqual(errors, []) + + # Forward recovery deleted backup and removed WAL + self.assertFalse(os.path.exists(backup_path)) + self.assertFalse(os.path.exists(_revert_wal_path(staging))) + self.assertEqual(_read(live_a), "orig_a") + + +class TestRevertCli(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.mkdtemp() + self.project = os.path.join(self.tmp, "project") + os.makedirs(self.project, exist_ok=True) + + def tearDown(self): + shutil.rmtree(self.tmp, ignore_errors=True) + + def _run(self, argv): + buf = io.StringIO() + with redirect_stdout(buf): + code = main(argv) + return code, buf.getvalue() + + def test_cli_reverts_and_reports(self): + live = os.path.join(self.project, "SKILL.md") + _write(live, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=live, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + adopt(staging) + + code, out = self._run(["revert", "--project", self.project]) + self.assertEqual(code, 0) + self.assertIn("reverted", out) + self.assertEqual(_read(live), ORIGINAL) + + def test_cli_nothing_adopted_exits_nonzero(self): + code, out = self._run(["revert", "--project", self.project]) + self.assertEqual(code, 1) + self.assertIn("nothing to revert", out) + + def test_cli_rejects_conflicting_selection_modes(self): + code, out = self._run([ + "revert", "--project", self.project, "--legacy", "--all-skills" + ]) + self.assertEqual(code, 2) + self.assertIn("exactly one of", out) + + def test_cli_requires_selection_when_both_kinds_adopted(self): + managed = os.path.join(self.project, "SKILL.md") + skill = os.path.join(self.project, "skills", "alpha", "SKILL.md") + _write(managed, ORIGINAL) + _write(skill, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=managed, + live_memory_path=None, + skill_proposals=[SkillProposal("alpha", PROPOSED, skill)], + report_md="", + ) + adopt(staging) + adopt_skills(staging, ["alpha"]) + + code, out = self._run(["revert", "--project", self.project]) + self.assertEqual(code, 2) + self.assertIn("--skill NAME", out) + self.assertEqual(_read(managed), PROPOSED) + + def test_cli_status_names_revertable_night(self): + live = os.path.join(self.project, "SKILL.md") + _write(live, ORIGINAL) + + staging = write_staging( + self.project, + report=SleepReport(night=1, project=self.project, accepted=True), + proposed_skill=PROPOSED, + proposed_memory=None, + live_skill_path=live, + live_memory_path=None, + skill_proposals=[], + report_md="", + ) + adopt(staging) + + code, out = self._run(["status", "--project", self.project, "--json"]) + self.assertEqual(code, 0) + data = json.loads(out) + self.assertEqual(data.get("revertable_staging"), staging) + + +if __name__ == "__main__": + unittest.main()