From dd48d3c1e5275a46cdb7edc8153f0abb7344d53c Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Thu, 24 Sep 2026 09:06:59 +0200 Subject: [PATCH 1/2] feat: implement the HypervisorDriver VM contract start_vm, stop_vm, reboot_vm, suspend_vm and get_vm_config existed only as declarations. netOrk called Proxmox's own power_vm and read a VM's raw config through _node_api(), so no other hypervisor could serve the same endpoints. These let netOrk talk to every hypervisor alike. The power methods accept a VM's name or vmid, wait for the Proxmox task, and raise ValueError/RuntimeError as the contract says instead of returning a result dict. A forced reboot of a container is stop + start, since LXC has no reset; suspending a container is refused. power_vm is unchanged for existing callers. get_vm_config moves the config parsing netOrk did in _parse_proxmox_hw_config into the driver and returns a VMConfigDict: disks with storage and size, NICs with model, MAC, bridge and VLAN, CPU topology, firmware, machine type and PCI/USB passthrough. get_vms reports vmid as a string ("100"), following napalm-device-types 2.0, still ordered numerically. --- CHANGELOG.md | 10 ++ napalm_proxmox/driver.py | 2 + napalm_proxmox/vm_contract_mixin.py | 184 +++++++++++++++++++++ napalm_proxmox/vm_mixin.py | 8 +- pyproject.toml | 2 +- tests/test_vm_contract.py | 239 ++++++++++++++++++++++++++++ 6 files changed, 440 insertions(+), 5 deletions(-) create mode 100644 napalm_proxmox/vm_contract_mixin.py create mode 100644 tests/test_vm_contract.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 51acf3c..a4cdc01 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added +- `HypervisorDriver` contract methods `start_vm`, `stop_vm`, `reboot_vm`, + `suspend_vm` and `get_vm_config`. They accept a VM's name or vmid, raise + `ValueError`/`RuntimeError` instead of returning a result dict, and wait + for the Proxmox task to finish. `power_vm` is unchanged. + +### Changed +- `get_vms()` reports `vmid` as a string (`"100"`), following + napalm-device-types 2.0. Ordering stays numeric. + ## [0.1.0] - 2024-01-01 ### Added diff --git a/napalm_proxmox/driver.py b/napalm_proxmox/driver.py index d419716..035b359 100644 --- a/napalm_proxmox/driver.py +++ b/napalm_proxmox/driver.py @@ -49,6 +49,7 @@ from napalm_proxmox.sdn_mixin import ProxmoxSDNMixin from napalm_proxmox.lldp_mixin import ProxmoxLLDPMixin from napalm_proxmox.config_mixin import ProxmoxConfigMixin from napalm_proxmox.vm_mixin import ProxmoxVMMixin +from napalm_proxmox.vm_contract_mixin import ProxmoxVMContractMixin from napalm_proxmox.vm_provision_mixin import ProxmoxVMProvisionMixin from napalm_proxmox.routing_mixin import ProxmoxRoutingMixin from napalm_proxmox.system_mixin import ProxmoxSystemMixin @@ -69,6 +70,7 @@ class ProxmoxDriver( ProxmoxLLDPMixin, ProxmoxConfigMixin, ProxmoxVMMixin, + ProxmoxVMContractMixin, ProxmoxVMProvisionMixin, ProxmoxRoutingMixin, ProxmoxSystemMixin, diff --git a/napalm_proxmox/vm_contract_mixin.py b/napalm_proxmox/vm_contract_mixin.py new file mode 100644 index 0000000..d71a117 --- /dev/null +++ b/napalm_proxmox/vm_contract_mixin.py @@ -0,0 +1,184 @@ +"""HypervisorDriver contract methods for Proxmox VE: power actions and VM config. + +``power_vm`` stays for callers that already use it; these are what a +hypervisor-neutral caller talks to. They raise instead of returning a +``{"success": ...}`` dict, and block until Proxmox reports the task finished. +""" + +from __future__ import annotations + +import re +from typing import Any + +from napalm_device_types.models import ( + VMConfigDict, + VMDiskDict, + VMNICDict, + VMPassthroughDict, +) + +_JsonDict = dict[str, Any] + +_POWER_TIMEOUT = 120 +_VM_DISK_KEY = re.compile(r"^(scsi|ide|virtio|sata)\d+$|^efidisk\d+$|^tpmstate\d+$") +_CT_DISK_KEY = re.compile(r"^rootfs$|^mp\d+$") +_NET_KEY = re.compile(r"^net\d+$") +_PASSTHROUGH_KEY = re.compile(r"^(hostpci|usb)\d+$") +_NIC_MODELS = {"virtio", "e1000", "e1000e", "vmxnet3", "rtl8139", "ne2k_pci"} +_SIZE = re.compile(r"^(\d+(?:\.\d+)?)([KMGT]?)$", re.I) +_GB_PER_UNIT = {"K": 1 / 1024**2, "M": 1 / 1024, "G": 1, "T": 1024, "": 1} + + +def _options(value: str) -> tuple[str, dict[str, str]]: + """Split ``"volume,key=val,..."`` into the leading bare part and its options.""" + head = "" + opts: dict[str, str] = {} + for part in str(value).split(","): + if "=" in part: + k, v = part.split("=", 1) + opts[k.strip().lower()] = v.strip() + elif not head: + head = part.strip() + return head, opts + + +def _size_gb(raw: str) -> int: + m = _SIZE.match(raw or "") + if not m: + return 0 + return int(float(m.group(1)) * _GB_PER_UNIT[m.group(2).upper()]) + + +def _boot_order(cfg: _JsonDict) -> list[str]: + boot = str(cfg.get("boot", "") or "") + if boot.startswith("order="): + return [d for d in boot[len("order=") :].split(";") if d] + bootdisk = cfg.get("bootdisk") + return [bootdisk] if bootdisk else [] + + +def _disks(cfg: _JsonDict, vm_type: str, boot_order: list[str]) -> list[VMDiskDict]: + key_re = _VM_DISK_KEY if vm_type == "vm" else _CT_DISK_KEY + disks: list[VMDiskDict] = [] + for key in sorted(cfg, key=lambda k: (k != "rootfs", k)): + if not key_re.match(key): + continue + value = str(cfg[key] or "") + head, opts = _options(value) + if opts.get("media") == "cdrom" or head in ("none", "0", ""): + continue + disks.append( + { + "device": key, + "storage": head.split(":", 1)[0], + "size": _size_gb(opts.get("size", "")), + "format": opts.get("format", ""), + "bootable": key in boot_order, + } + ) + return disks + + +def _nic(key: str, value: str) -> VMNICDict: + _, opts = _options(value) + model = next((m for m in _NIC_MODELS if m in opts), opts.get("type", "")) + mac = opts.get(model, "") if model in _NIC_MODELS else opts.get("hwaddr", "") + tag = opts.get("tag", "") + return { + "device": key, + "mac": mac.upper(), + "model": model, + "bridge": opts.get("bridge", ""), + "vlan_id": int(tag) if tag.isdigit() else 0, + } + + +def _passthrough(cfg: _JsonDict) -> list[VMPassthroughDict]: + return [ + {"slot": key, "kind": "pci" if key.startswith("hostpci") else "usb", "config": str(val)} + for key, val in sorted(cfg.items()) + if _PASSTHROUGH_KEY.match(key) + ] + + +def parse_vm_config(vmid: str, vm_type: str, cfg: _JsonDict) -> VMConfigDict: + """Turn a raw ``/qemu/{id}/config`` or ``/lxc/{id}/config`` into a VMConfigDict.""" + is_vm = vm_type == "vm" + cores = int(cfg.get("cores", 1) or 1) + sockets = int(cfg.get("sockets", 1) or 1) if is_vm else 1 + boot_order = _boot_order(cfg) + tags = str(cfg.get("tags", "") or "") + name_key = "name" if is_vm else "hostname" + result: VMConfigDict = { + "name": cfg.get(name_key) or f"{'vm' if is_vm else 'ct'}-{vmid}", + "vmid": vmid, + "vcpus": cores * sockets, + "memory": int(cfg.get("memory", 0) or 0), + "os_type": cfg.get("ostype", ""), + "boot_order": boot_order, + "disks": _disks(cfg, vm_type, boot_order), + "nics": [_nic(k, str(v)) for k, v in sorted(cfg.items()) if _NET_KEY.match(k)], + "description": cfg.get("description", ""), + "tags": [t for t in re.split(r"[;,\s]+", tags) if t], + "passthrough": _passthrough(cfg), + } + if is_vm: + result["cpu_type"] = str(cfg.get("cpu", "kvm64")).split(",")[0].removeprefix("cputype=") + result["sockets"] = sockets + result["cores_per_socket"] = cores + result["firmware"] = "efi" if cfg.get("bios") == "ovmf" else "bios" + if cfg.get("machine"): + result["machine"] = cfg["machine"] + return result + + +class ProxmoxVMContractMixin: + """HypervisorDriver's VM methods on top of the Proxmox node API.""" + + def _resolve_vm(self, name: str) -> tuple[int, str]: + """Find a guest by vmid or display name; return ``(vmid, "vm"|"container")``.""" + node = self._node_api() + for vm_type, listing in (("vm", node.qemu), ("container", node.lxc)): + for guest in listing.get() or []: + if str(guest.get("vmid")) == name or guest.get("name") == name: + return int(guest["vmid"]), vm_type + raise ValueError(f"No VM or container named or numbered {name!r}") + + def _guest_api(self, vmid: int, vm_type: str) -> Any: + node = self._node_api() + return node.qemu(vmid) if vm_type == "vm" else node.lxc(vmid) + + def _run_power(self, vmid: int, vm_type: str, action: str) -> None: + try: + upid = getattr(self._guest_api(vmid, vm_type).status, action).post() + except Exception as exc: + raise RuntimeError(f"{action} of {vm_type} {vmid} failed: {exc}") from exc + if upid: + self._wait_for_task(upid, timeout=_POWER_TIMEOUT) + + def start_vm(self, name: str) -> None: + self._run_power(*self._resolve_vm(name), "start") + + def stop_vm(self, name: str, force: bool = False) -> None: + self._run_power(*self._resolve_vm(name), "stop" if force else "shutdown") + + def reboot_vm(self, name: str, force: bool = False) -> None: + vmid, vm_type = self._resolve_vm(name) + if not force: + self._run_power(vmid, vm_type, "reboot") + elif vm_type == "vm": + self._run_power(vmid, vm_type, "reset") + else: + self._run_power(vmid, vm_type, "stop") + self._run_power(vmid, vm_type, "start") + + def suspend_vm(self, name: str) -> None: + vmid, vm_type = self._resolve_vm(name) + if vm_type != "vm": + raise RuntimeError(f"Proxmox cannot suspend container {vmid}") + self._run_power(vmid, vm_type, "suspend") + + def get_vm_config(self, name: str) -> VMConfigDict: + vmid, vm_type = self._resolve_vm(name) + cfg = self._guest_api(vmid, vm_type).config.get() or {} + return parse_vm_config(str(vmid), vm_type, cfg) diff --git a/napalm_proxmox/vm_mixin.py b/napalm_proxmox/vm_mixin.py index ee26d95..6aaf0b5 100644 --- a/napalm_proxmox/vm_mixin.py +++ b/napalm_proxmox/vm_mixin.py @@ -216,7 +216,7 @@ class ProxmoxVMMixin: """Return all VMs (QEMU) and containers (LXC) on this node. Each entry contains: - * vmid (int) - Proxmox VM/container ID + * vmid (str) - Proxmox VM/container ID, e.g. ``"100"`` * name (str) - display name * type (str) - ``"vm"`` or ``"container"`` * status (str) - ``"running"``, ``"stopped"``, etc. @@ -257,7 +257,7 @@ class ProxmoxVMMixin: disks, onboot = self._get_vm_disk_and_boot(vmid, "qemu") result.append({ - "vmid": vmid, + "vmid": str(vmid), "name": name, "type": "vm", "status": status, @@ -301,7 +301,7 @@ class ProxmoxVMMixin: disks, onboot = self._get_vm_disk_and_boot(vmid, "lxc") result.append({ - "vmid": vmid, + "vmid": str(vmid), "name": name, "type": "container", "status": status, @@ -321,7 +321,7 @@ class ProxmoxVMMixin: except Exception as exc: logger.warning("get_vms: failed to list LXC containers: %s", exc) - return sorted(result, key=lambda x: x["vmid"]) + return sorted(result, key=lambda x: int(x["vmid"])) # Disk-key prefixes for QEMU: scsi, virtio, ide, sata (exclude cdrom/none entries) _DISK_KEYS_VM = re.compile(r"^(scsi|virtio|ide|sata)\d+$") diff --git a/pyproject.toml b/pyproject.toml index 65beae4..859ae9e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,7 +25,7 @@ classifiers = [ requires-python = ">=3.9" dependencies = [ "napalm>=5.0.0", - "napalm_device_types>=0.1.0", + "napalm_device_types>=2.0.0", "paramiko>=5.0.0", # CVE-2026-44405; imported directly for SSH fallback (driver.py) "proxmoxer>=2.0.0", "netaddr>=0.9.0", diff --git a/tests/test_vm_contract.py b/tests/test_vm_contract.py new file mode 100644 index 0000000..eae97ba --- /dev/null +++ b/tests/test_vm_contract.py @@ -0,0 +1,239 @@ +"""HypervisorDriver contract methods: VM lookup, power actions, get_vm_config. + +netOrk used to call Proxmox's own ``power_vm`` and reach into ``_node_api()`` +for a VM's hardware. Both are Proxmox-only, so a second hypervisor could not +serve the same endpoints. These pin the contract methods that replace them. +""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from napalm_proxmox.vm_contract_mixin import parse_vm_config + +QEMU_LIST = [{"vmid": 100, "name": "web01", "status": "running"}] +LXC_LIST = [{"vmid": 200, "name": "dns01", "status": "running"}] + + +@pytest.fixture +def api(driver): + node = driver._node_api() + node.qemu.get.return_value = QEMU_LIST + node.lxc.get.return_value = LXC_LIST + node.qemu.return_value.status.start.post.return_value = "UPID:start" + driver._wait_for_task = MagicMock() + return node + + +class TestGetVmsReportsStringIds: + def test_vmid_is_a_string(self, driver, api): + driver.get_vm_interfaces = MagicMock(return_value=({}, False, False)) + driver._get_vm_disk_and_boot = MagicMock(return_value=([], False)) + assert [vm["vmid"] for vm in driver.get_vms()] == ["100", "200"] + + def test_ordered_numerically_not_lexically(self, driver, api): + api.qemu.get.return_value = [{"vmid": 1000, "name": "a"}, {"vmid": 99, "name": "b"}] + api.lxc.get.return_value = [] + driver.get_vm_interfaces = MagicMock(return_value=({}, False, False)) + driver._get_vm_disk_and_boot = MagicMock(return_value=([], False)) + assert [vm["vmid"] for vm in driver.get_vms()] == ["99", "1000"] + + +class TestResolveVm: + def test_by_vmid_string(self, driver, api): + assert driver._resolve_vm("100") == (100, "vm") + + def test_by_name(self, driver, api): + assert driver._resolve_vm("dns01") == (200, "container") + + def test_unknown_raises_value_error(self, driver, api): + with pytest.raises(ValueError, match="nope"): + driver._resolve_vm("nope") + + +class TestPowerActions: + def test_start_posts_and_waits_for_the_task(self, driver, api): + driver.start_vm("web01") + api.qemu.return_value.status.start.post.assert_called_once() + driver._wait_for_task.assert_called_once_with("UPID:start", timeout=120) + + @pytest.mark.parametrize(("force", "action"), [(False, "shutdown"), (True, "stop")]) + def test_stop_graceful_or_forced(self, driver, api, force, action): + driver.stop_vm("100", force=force) + getattr(api.qemu.return_value.status, action).post.assert_called_once() + + @pytest.mark.parametrize(("force", "action"), [(False, "reboot"), (True, "reset")]) + def test_reboot_graceful_or_forced(self, driver, api, force, action): + driver.reboot_vm("100", force=force) + getattr(api.qemu.return_value.status, action).post.assert_called_once() + + def test_forced_reboot_of_a_container_is_a_stop_and_start(self, driver, api): + """LXC has no reset; stop + start is the closest thing to pulling the plug.""" + driver.reboot_vm("200", force=True) + status = api.lxc.return_value.status + status.stop.post.assert_called_once() + status.start.post.assert_called_once() + + def test_suspend_vm(self, driver, api): + driver.suspend_vm("100") + api.qemu.return_value.status.suspend.post.assert_called_once() + + def test_suspend_container_is_refused(self, driver, api): + with pytest.raises(RuntimeError, match="container"): + driver.suspend_vm("200") + + def test_api_error_becomes_runtime_error(self, driver, api): + api.qemu.return_value.status.start.post.side_effect = Exception("locked") + with pytest.raises(RuntimeError, match="locked"): + driver.start_vm("100") + + +QEMU_CONFIG = { + "name": "web01", + "cores": 2, + "sockets": 2, + "memory": "8192", + "ostype": "l26", + "cpu": "host,flags=+aes", + "bios": "ovmf", + "machine": "q35", + "boot": "order=scsi0;ide2;net0", + "scsi0": "local-lvm:vm-100-disk-0,size=32G,format=raw", + "virtio1": "tank:vm-100-disk-1,size=512M", + "ide2": "local:iso/debian.iso,media=cdrom", + "efidisk0": "local-lvm:vm-100-disk-2,size=4M", + "net0": "virtio=BC:24:11:AA:BB:CC,bridge=vmbr0,tag=10,firewall=1", + "net1": "e1000=BC:24:11:AA:BB:DD,bridge=vmbr1", + "hostpci0": "0000:01:00.0,pcie=1", + "usb0": "host=1234:5678", + "description": "Production web server", + "tags": "prod;web", +} + +LXC_CONFIG = { + "hostname": "dns01", + "cores": 1, + "memory": 512, + "ostype": "debian", + "rootfs": "local-lvm:vm-200-disk-0,size=8G", + "mp0": "tank:subvol-200-disk-1,mp=/data,size=1T", + "net0": "name=eth0,bridge=vmbr0,hwaddr=BC:24:11:00:00:01,ip=dhcp,tag=20,type=veth", +} + + +class TestParseQemuConfig: + @pytest.fixture + def cfg(self): + return parse_vm_config("100", "vm", QEMU_CONFIG) + + def test_core_fields(self, cfg): + assert cfg["vmid"] == "100" + assert cfg["name"] == "web01" + assert cfg["vcpus"] == 4 + assert cfg["memory"] == 8192 + assert cfg["os_type"] == "l26" + assert cfg["description"] == "Production web server" + assert cfg["tags"] == ["prod", "web"] + + def test_boot_order(self, cfg): + assert cfg["boot_order"] == ["scsi0", "ide2", "net0"] + + def test_disks_skip_cdrom_and_normalise_size(self, cfg): + by_dev = {d["device"]: d for d in cfg["disks"]} + assert set(by_dev) == {"scsi0", "virtio1", "efidisk0"} + assert by_dev["scsi0"] == { + "device": "scsi0", + "storage": "local-lvm", + "size": 32, + "format": "raw", + "bootable": True, + } + assert by_dev["virtio1"]["size"] == 0 # 512M rounds down to 0 GB + assert by_dev["virtio1"]["bootable"] is False + + def test_nics(self, cfg): + assert cfg["nics"] == [ + { + "device": "net0", + "mac": "BC:24:11:AA:BB:CC", + "model": "virtio", + "bridge": "vmbr0", + "vlan_id": 10, + }, + { + "device": "net1", + "mac": "BC:24:11:AA:BB:DD", + "model": "e1000", + "bridge": "vmbr1", + "vlan_id": 0, + }, + ] + + def test_hardware_details(self, cfg): + assert cfg["cpu_type"] == "host" + assert cfg["sockets"] == 2 + assert cfg["cores_per_socket"] == 2 + assert cfg["firmware"] == "efi" + assert cfg["machine"] == "q35" + + def test_passthrough(self, cfg): + assert cfg["passthrough"] == [ + {"slot": "hostpci0", "kind": "pci", "config": "0000:01:00.0,pcie=1"}, + {"slot": "usb0", "kind": "usb", "config": "host=1234:5678"}, + ] + + def test_defaults_for_a_bare_config(self): + cfg = parse_vm_config("101", "vm", {}) + assert cfg["name"] == "vm-101" + assert cfg["vcpus"] == 1 + assert cfg["cpu_type"] == "kvm64" + assert cfg["firmware"] == "bios" + assert "machine" not in cfg + assert cfg["boot_order"] == [] + + def test_legacy_bootdisk(self): + cfg = parse_vm_config("101", "vm", {"boot": "cdn", "bootdisk": "scsi0"}) + assert cfg["boot_order"] == ["scsi0"] + + +class TestParseLxcConfig: + @pytest.fixture + def cfg(self): + return parse_vm_config("200", "container", LXC_CONFIG) + + def test_core_fields(self, cfg): + assert cfg["name"] == "dns01" + assert cfg["vcpus"] == 1 + assert cfg["memory"] == 512 + assert cfg["tags"] == [] + + def test_rootfs_and_mountpoint(self, cfg): + assert [(d["device"], d["storage"], d["size"]) for d in cfg["disks"]] == [ + ("rootfs", "local-lvm", 8), + ("mp0", "tank", 1024), + ] + + def test_veth_nic(self, cfg): + assert cfg["nics"] == [ + { + "device": "net0", + "mac": "BC:24:11:00:00:01", + "model": "veth", + "bridge": "vmbr0", + "vlan_id": 20, + } + ] + + def test_no_vm_only_hardware_fields(self, cfg): + assert "firmware" not in cfg + assert "sockets" not in cfg + + +class TestGetVmConfig: + def test_fetches_the_right_config(self, driver, api): + api.lxc.return_value.config.get.return_value = LXC_CONFIG + cfg = driver.get_vm_config("dns01") + assert cfg["vmid"] == "200" + assert cfg["name"] == "dns01" -- 2.54.0 From 08bfb5c1c0018f0059869b291394587f5259d6d6 Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Thu, 24 Sep 2026 10:00:12 +0200 Subject: [PATCH 2/2] feat: VM snapshots and reboot_host through the API get_vm_snapshots, create_vm_snapshot, delete_vm_snapshot and rollback_vm_snapshot for VMs and containers, so netOrk's snapshot view works on Proxmox as it does on VMware. Proxmox lists the live state as a pseudo-snapshot named "current"; it is never reported or addressable. Containers have no RAM state, so include_memory is ignored for them. reboot_host() restarts the node with POST /nodes/{node}/status command=reboot instead of /sbin/reboot over SSH. --- CHANGELOG.md | 4 ++ napalm_proxmox/driver.py | 2 + napalm_proxmox/vm_contract_mixin.py | 9 +++ napalm_proxmox/vm_snapshot_mixin.py | 74 ++++++++++++++++++++++ tests/test_reboot_host.py | 16 +++++ tests/test_vm_snapshots.py | 95 +++++++++++++++++++++++++++++ 6 files changed, 200 insertions(+) create mode 100644 napalm_proxmox/vm_snapshot_mixin.py create mode 100644 tests/test_reboot_host.py create mode 100644 tests/test_vm_snapshots.py diff --git a/CHANGELOG.md b/CHANGELOG.md index a4cdc01..fbb7a3e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 `suspend_vm` and `get_vm_config`. They accept a VM's name or vmid, raise `ValueError`/`RuntimeError` instead of returning a result dict, and wait for the Proxmox task to finish. `power_vm` is unchanged. +- Snapshot methods `get_vm_snapshots`, `create_vm_snapshot`, + `delete_vm_snapshot`, `rollback_vm_snapshot` for VMs and containers + (containers never save RAM state). +- `reboot_host()` restarts the node through the API instead of SSH. ### Changed - `get_vms()` reports `vmid` as a string (`"100"`), following diff --git a/napalm_proxmox/driver.py b/napalm_proxmox/driver.py index 035b359..515dca4 100644 --- a/napalm_proxmox/driver.py +++ b/napalm_proxmox/driver.py @@ -50,6 +50,7 @@ from napalm_proxmox.lldp_mixin import ProxmoxLLDPMixin from napalm_proxmox.config_mixin import ProxmoxConfigMixin from napalm_proxmox.vm_mixin import ProxmoxVMMixin from napalm_proxmox.vm_contract_mixin import ProxmoxVMContractMixin +from napalm_proxmox.vm_snapshot_mixin import ProxmoxVMSnapshotMixin from napalm_proxmox.vm_provision_mixin import ProxmoxVMProvisionMixin from napalm_proxmox.routing_mixin import ProxmoxRoutingMixin from napalm_proxmox.system_mixin import ProxmoxSystemMixin @@ -71,6 +72,7 @@ class ProxmoxDriver( ProxmoxConfigMixin, ProxmoxVMMixin, ProxmoxVMContractMixin, + ProxmoxVMSnapshotMixin, ProxmoxVMProvisionMixin, ProxmoxRoutingMixin, ProxmoxSystemMixin, diff --git a/napalm_proxmox/vm_contract_mixin.py b/napalm_proxmox/vm_contract_mixin.py index d71a117..4f4b44d 100644 --- a/napalm_proxmox/vm_contract_mixin.py +++ b/napalm_proxmox/vm_contract_mixin.py @@ -182,3 +182,12 @@ class ProxmoxVMContractMixin: vmid, vm_type = self._resolve_vm(name) cfg = self._guest_api(vmid, vm_type).config.get() or {} return parse_vm_config(str(vmid), vm_type, cfg) + + # -- the node itself -------------------------------------------------------- + + def reboot_host(self) -> None: + """Restart this Proxmox node through the API (no SSH involved).""" + try: + self._node_api().status.post(command="reboot") + except Exception as exc: + raise RuntimeError(f"Reboot of node {self._node_name!r} refused: {exc}") from exc diff --git a/napalm_proxmox/vm_snapshot_mixin.py b/napalm_proxmox/vm_snapshot_mixin.py new file mode 100644 index 0000000..6c6d052 --- /dev/null +++ b/napalm_proxmox/vm_snapshot_mixin.py @@ -0,0 +1,74 @@ +"""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) diff --git a/tests/test_reboot_host.py b/tests/test_reboot_host.py new file mode 100644 index 0000000..2b3d08e --- /dev/null +++ b/tests/test_reboot_host.py @@ -0,0 +1,16 @@ +"""reboot_host: restart the Proxmox node itself through the API, not over SSH.""" + +from __future__ import annotations + +import pytest + + +def test_posts_reboot_to_the_node(driver): + driver.reboot_host() + driver._node_api().status.post.assert_called_once_with(command="reboot") + + +def test_api_refusal_is_a_runtime_error(driver): + driver._node_api().status.post.side_effect = Exception("Permission check failed") + with pytest.raises(RuntimeError, match="Permission check failed"): + driver.reboot_host() diff --git a/tests/test_vm_snapshots.py b/tests/test_vm_snapshots.py new file mode 100644 index 0000000..c878d7f --- /dev/null +++ b/tests/test_vm_snapshots.py @@ -0,0 +1,95 @@ +"""HypervisorDriver snapshot methods on Proxmox (QEMU and LXC).""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +SNAPSHOTS = [ + {"name": "base", "description": "clean install", "snaptime": 1700000000, "vmstate": 0}, + {"name": "upgrade", "description": "", "snaptime": 1700000100, "parent": "base", "vmstate": 1}, + {"name": "current", "description": "You are here!", "parent": "upgrade", "running": 1}, +] + + +@pytest.fixture +def api(driver): + node = driver._node_api() + node.qemu.get.return_value = [{"vmid": 100, "name": "web01"}] + node.lxc.get.return_value = [{"vmid": 200, "name": "dns01"}] + node.qemu.return_value.snapshot.get.return_value = SNAPSHOTS + driver._wait_for_task = MagicMock() + return node + + +class TestList: + def test_flattens_and_skips_the_current_marker(self, driver, api): + assert driver.get_vm_snapshots("web01") == [ + { + "name": "base", + "vm": "web01", + "created": 1700000000.0, + "description": "clean install", + "has_memory": False, + "parent": "", + }, + { + "name": "upgrade", + "vm": "web01", + "created": 1700000100.0, + "description": "", + "has_memory": True, + "parent": "base", + }, + ] + + def test_unknown_vm(self, driver, api): + with pytest.raises(ValueError): + driver.get_vm_snapshots("nope") + + +class TestCreate: + def test_vm_with_memory(self, driver, api): + api.qemu.return_value.snapshot.post.return_value = "UPID:snap" + driver.create_vm_snapshot("100", "pre", description="d", include_memory=True) + api.qemu.return_value.snapshot.post.assert_called_once_with( + snapname="pre", description="d", vmstate=1 + ) + driver._wait_for_task.assert_called_once_with("UPID:snap", timeout=600) + + def test_container_never_saves_memory(self, driver, api): + api.lxc.return_value.snapshot.get.return_value = [] + driver.create_vm_snapshot("200", "pre", include_memory=True) + api.lxc.return_value.snapshot.post.assert_called_once_with(snapname="pre", description="") + + def test_duplicate_name(self, driver, api): + with pytest.raises(ValueError, match="already"): + driver.create_vm_snapshot("web01", "base") + + def test_api_refusal(self, driver, api): + api.qemu.return_value.snapshot.post.side_effect = Exception( + "snapshot feature is not available" + ) + with pytest.raises(RuntimeError, match="not available"): + driver.create_vm_snapshot("web01", "new") + + +class TestDeleteAndRollback: + def test_delete(self, driver, api): + driver.delete_vm_snapshot("web01", "base") + api.qemu.return_value.snapshot.assert_called_with("base") + api.qemu.return_value.snapshot.return_value.delete.assert_called_once_with() + + def test_rollback(self, driver, api): + driver.rollback_vm_snapshot("web01", "upgrade") + api.qemu.return_value.snapshot.return_value.rollback.post.assert_called_once_with() + + @pytest.mark.parametrize("method", ["delete_vm_snapshot", "rollback_vm_snapshot"]) + def test_unknown_snapshot(self, driver, api, method): + with pytest.raises(ValueError, match="no snapshot"): + getattr(driver, method)("web01", "nope") + + def test_current_is_not_a_snapshot(self, driver, api): + with pytest.raises(ValueError): + driver.rollback_vm_snapshot("web01", "current") -- 2.54.0