"""HypervisorDriver snapshot methods for Proxmox VE guests (QEMU and LXC).""" from __future__ import annotations from typing import Any from napalm_device_types.models import SnapshotDict #: Proxmox lists the live state as a pseudo-snapshot of this name. _CURRENT = "current" #: A RAM snapshot of a large VM takes a while to write out. _SNAPSHOT_TIMEOUT = 600 class ProxmoxVMSnapshotMixin: """Relies on ``_resolve_vm``/``_guest_api`` from ProxmoxVMContractMixin.""" _resolve_vm: Any _guest_api: Any _wait_for_task: Any def _snapshots(self, name: str) -> tuple[Any, str, str, list[dict[str, Any]]]: vmid, vm_type = self._resolve_vm(name) api = self._guest_api(vmid, vm_type) raw = [s for s in api.snapshot.get() or [] if s.get("name") != _CURRENT] return api, str(vmid), vm_type, raw def _run_task(self, call: Any, *args: Any, **kwargs: Any) -> None: try: upid = call(*args, **kwargs) except Exception as exc: raise RuntimeError(str(exc)) from exc if upid: self._wait_for_task(upid, timeout=_SNAPSHOT_TIMEOUT) @staticmethod def _require(raw: list[dict[str, Any]], snapshot: str, vm: str) -> None: if not any(s.get("name") == snapshot for s in raw): raise ValueError(f"VM {vm!r} has no snapshot named {snapshot!r}") def get_vm_snapshots(self, name: str) -> list[SnapshotDict]: _, _, _, raw = self._snapshots(name) return [ { "name": s["name"], "vm": name, "created": float(s.get("snaptime", 0)), "description": s.get("description", ""), "has_memory": bool(s.get("vmstate")), "parent": s.get("parent", ""), } for s in raw ] def create_vm_snapshot( self, name: str, snapshot: str, description: str = "", include_memory: bool = False ) -> None: api, _, vm_type, raw = self._snapshots(name) if any(s.get("name") == snapshot for s in raw): raise ValueError(f"VM {name!r} already has a snapshot named {snapshot!r}") kwargs: dict[str, Any] = {"snapname": snapshot, "description": description} if vm_type == "vm": # containers have no RAM state to save kwargs["vmstate"] = 1 if include_memory else 0 self._run_task(api.snapshot.post, **kwargs) def delete_vm_snapshot(self, name: str, snapshot: str) -> None: api, _, _, raw = self._snapshots(name) self._require(raw, snapshot, name) self._run_task(api.snapshot(snapshot).delete) def rollback_vm_snapshot(self, name: str, snapshot: str) -> None: api, _, _, raw = self._snapshots(name) self._require(raw, snapshot, name) self._run_task(api.snapshot(snapshot).rollback.post)