"""What the ESXi and vCenter drivers share: session, VM lookup, common getters.""" from __future__ import annotations from pathlib import Path from typing import Any, ClassVar from napalm.base.exceptions import ConnectionException from napalm_device_types import HypervisorDriver from napalm_device_types.models import ( SnapshotDict, StorageVolumeDict, VirtualNetworkDict, VMConfigDict, ) from napalm_vmware import _session, paths from napalm_vmware._inventory import Inventory from napalm_vmware.parse.environment import environment from napalm_vmware.parse.networks import NetworkIndex, network_index, virtual_networks from napalm_vmware.parse.snapshots import snapshot_list from napalm_vmware.parse.storage import storage_pools from napalm_vmware.parse.vm_config import vm_config from napalm_vmware.parse.vms import vm_list from napalm_vmware.parse.warnings import host_warnings from napalm_vmware.provision.image import DEFAULT_CACHE_DIR _DEFAULT_PORT = 443 #: ``about.apiType`` -> the driver that handles it, for a helpful refusal. _DRIVER_FOR_API = {"HostAgent": "vmware_esxi", "VirtualCenter": "vmware_vcenter"} class VmwareBaseDriver(HypervisorDriver): """Shared implementation. Concrete drivers set ``API_TYPE`` and ``STANDALONE``.""" VENDOR = "VMware" USES_SSH = False platform = "vmware" #: ``about.apiType`` this driver accepts. API_TYPE: ClassVar[str] = "" #: True when the driver talks to one host directly (ESXi), False for vCenter. STANDALONE: ClassVar[bool] = True def __init__( self, hostname: str, username: str, password: str, timeout: int = 60, optional_args: dict[str, Any] | None = None, ) -> None: self.hostname = hostname self.username = username self.password = password self.timeout = timeout args = optional_args or {} self._port = int(args.get("port") or _DEFAULT_PORT) self._verify_ssl = bool(args.get("verify_ssl", args.get("ssl_verify", True))) # Where converted cloud images are kept between provisioning jobs. self._image_cache_dir = Path(args.get("image_cache_dir") or DEFAULT_CACHE_DIR) self._si: Any = None self._inventory: Any = None # -- session ------------------------------------------------------------- def open(self) -> None: try: si = _session.connect( self.hostname, self._port, self.username, self.password, verify_ssl=self._verify_ssl, timeout=self.timeout, ) except Exception as exc: raise ConnectionException(f"Cannot connect to {self.hostname}: {exc}") from exc inventory = Inventory(si.RetrieveContent()) about = inventory.about() if about.get("apiType") != self.API_TYPE: _session.disconnect(si) other = _DRIVER_FOR_API.get(about.get("apiType", ""), "another driver") raise ConnectionException( f"{self.hostname} is {about.get('fullName', 'not a supported VMware endpoint')}" f"; use the {other} driver for it" ) self._si, self._inventory = si, inventory def close(self) -> None: if self._si is not None: _session.disconnect(self._si) self._si = self._inventory = None def is_alive(self) -> dict[str, bool]: return {"is_alive": self._si is not None and _session.alive(self._si)} def _mo(self, vim_type: Any, moref: str) -> Any: """A live managed-object reference for a MoRef id from a plain row.""" return vim_type(moref, self._si._stub) # -- inventory reads ------------------------------------------------------- def _hosts(self) -> list[dict[str, Any]]: return self._inventory.collect("HostSystem", paths.HOST) def _vm_rows(self) -> list[dict[str, Any]]: return self._inventory.collect("VirtualMachine", paths.VM) def _index(self, hosts: list[dict[str, Any]]) -> NetworkIndex: dv_portgroups = self._inventory.collect("DistributedVirtualPortgroup", paths.DV_PORTGROUP) return network_index(hosts, dv_portgroups) def _find_vm(self, name: str) -> dict[str, Any]: """The VM whose instance UUID, MoRef or name is ``name``.""" vms = [ vm for vm in self._vm_rows() if vm.get("config.instanceUuid") and not vm.get("config.template") ] for key in ("config.instanceUuid", "_moref"): match = [vm for vm in vms if vm.get(key) == name] if match: return match[0] match = [vm for vm in vms if vm.get("name") == name] if len(match) > 1: raise ValueError(f"{len(match)} VMs are named {name!r}; address one by its vmid") if not match: raise ValueError(f"There is no VM named or identified by {name!r}") return match[0] # -- HypervisorDriver ------------------------------------------------------ def get_vms(self) -> list[dict[str, Any]]: hosts = self._hosts() return vm_list(self._vm_rows(), hosts, self._index(hosts)) def get_vm_config(self, name: str) -> VMConfigDict: return vm_config(self._find_vm(name), self._index(self._hosts())) def get_vm_snapshots(self, name: str) -> list[SnapshotDict]: return snapshot_list(self._find_vm(name)) def get_vm_storage_pools(self) -> dict[str, StorageVolumeDict]: return storage_pools(self._inventory.collect("Datastore", paths.DATASTORE)) def get_virtual_networks(self) -> dict[str, VirtualNetworkDict]: return virtual_networks( # type: ignore[return-value] self._hosts(), self._inventory.collect("DistributedVirtualPortgroup", paths.DV_PORTGROUP), self._inventory.collect("DistributedVirtualSwitch", paths.DV_SWITCH), ) def get_environment(self) -> dict[str, Any]: return environment(self._hosts()) def get_device_warnings(self) -> list[dict[str, Any]]: return host_warnings(self._hosts(), self._inventory.licenses(), standalone=self.STANDALONE)