diff --git a/changes/+eos_ssh_extensions.added b/changes/+eos_ssh_extensions.added new file mode 100644 index 00000000..09ad5e5a --- /dev/null +++ b/changes/+eos_ssh_extensions.added @@ -0,0 +1,2 @@ +Added install_os method for EOSSSHDevice class. +Added maintenance_mode method for EOSSSHDevice class. diff --git a/pyntc/devices/eos_device.py b/pyntc/devices/eos_device.py index 997f918e..ad114d3a 100644 --- a/pyntc/devices/eos_device.py +++ b/pyntc/devices/eos_device.py @@ -648,7 +648,7 @@ def _check_copy_output_for_errors(self, output): log.error("Host %s: Error detected in copy command output: %s", self.host, output) raise FileTransferError(f"Error detected in copy command output: {output}") - def remote_file_copy(self, src: FileCopyModel, dest: str | None = None, file_system: str | None = None, **kwargs): + def remote_file_copy(self, src: FileCopyModel, dest: str | None = None, file_system: str | None = None, **kwargs): # pylint: disable=too-many-branches """Copy a file from remote source to device. Args: @@ -683,11 +683,18 @@ def remote_file_copy(self, src: FileCopyModel, dest: str | None = None, file_sys if dest is None: dest = src.file_name - log.debug("Host %s: Starting remote file copy for %s to %s/%s", self.host, src.file_name, file_system, dest) - self.open() self.enable() + if self.check_file_exists(dest, file_system): + if self.verify_optimized_image(dest, file_system): + log.debug("Host %s: File '%s' already exists in '%s'", self.host, dest, file_system) + return + log.debug("Host %s: File '%s' is present but cannot be verified", self.host, dest) + self.native_ssh.send_command(f"delete {file_system}{dest}") + + log.debug("Host %s: Starting remote file copy for %s to %s/%s", self.host, src.file_name, file_system, dest) + self._pre_transfer_space_check(src, file_system) if src.scheme == "tftp" or src.username is None: @@ -762,6 +769,31 @@ def verify_file(self, checksum, filename, hashing_algorithm="md5", **kwargs): ) return False + def verify_optimized_image(self, image_name: str, file_system: str | None): + """Verify the optimized image file without a checksum. + + An Arista EOS image mutates after installation, reducing its + storage footprint and thereby changing its checksum value. The + optimized image can be validated from the command line using + the `verify [file_system]:[image_name]` command without + specifying the checksum hash type. + + Args: + image_name (str): The name of the image to verify. + file_system (str): The device file system to inspect. + + Returns: + (bool): True if file verification is successful, else False. + """ + self.open() + file_system = file_system or self._get_file_system() + result = self.native_ssh.send_command(f"verify {file_system}{image_name}", read_timeout=30) + verified = re.search(f"Verifying {file_system}{image_name} successful.", result) + if verified: + return True + + return False + def install_os(self, image_name, reboot=True, **vendor_specifics): """Install new OS on device. diff --git a/pyntc/devices/eos_ssh_device.py b/pyntc/devices/eos_ssh_device.py index c6beddc0..70b2c2e6 100644 --- a/pyntc/devices/eos_ssh_device.py +++ b/pyntc/devices/eos_ssh_device.py @@ -13,13 +13,22 @@ import json import os import re +import time from netmiko import ConnectHandler from pyntc import log from pyntc.devices.base_device import BaseDevice, fix_docs from pyntc.devices.eos_device import DEFAULT_REBOOT_TIMEOUT, EOSDevice -from pyntc.errors import CommandError, CommandListError, FileTransferError, SocketClosedError +from pyntc.errors import ( + CommandError, + CommandListError, + FileTransferError, + MaintModeProfileError, + OSInstallError, + RebootTimeoutError, + SocketClosedError, +) DEFAULT_SSH_PORT = 22 @@ -79,6 +88,29 @@ def __init__(self, host, username, password, secret="", port=None, **kwargs): # self.open() log.init(host=host) + @property + def uptime(self): + """Get device uptime in seconds.""" + return self.show("show version")["uptime"] + + @property + def vlans(self): + """Get list of VLANs on device. + + ``EOSDevice`` delegates to ``EOSVlans``, which is pyeapi-only + (``device.native.api("vlans")``). Over SSH the same data comes from + ``show vlan | json``, whose ``vlans`` key is a dict keyed by VLAN id. + + Returns: + (list): List of VLAN ids as strings. + """ + if self._vlans is None: + # sorted() over str keys, matching EOSVlans.get_list()'s lexicographic ordering. + self._vlans = sorted(self.show("show vlan")["vlans"].keys()) + + log.debug("Host %s: Vlans %s", self.host, self._vlans) + return self._vlans + @property def native_ssh(self): """Alias for ``native`` so inherited Netmiko-backed code works unchanged. @@ -93,8 +125,22 @@ def native_ssh(self): """ return self.native + @property + def os_version(self): + """Get OS version on device. + + Returns: + (str): OS version of device. + """ + if self._os_version is None: + sh_version_output = self.show("show version") + self._os_version = sh_version_output["version"] + + log.debug("Host %s: OS version %s", self.host, self._os_version) + return self._os_version + @staticmethod - def _read_timeout_for(command): + def _read_timeout_for(command: str): """Resolve the Netmiko read timeout to use for ``command``. Args: @@ -108,7 +154,7 @@ def _read_timeout_for(command): return timeout return DEFAULT_READ_TIMEOUT - def _check_output_for_errors(self, command, output): + def _check_output_for_errors(self, command: str, output: str): """Raise ``CommandError`` when the device reported a CLI error. Args: @@ -122,7 +168,7 @@ def _check_output_for_errors(self, command, output): log.error("Host %s: Error in %s with response: %s", self.host, command, output) raise CommandError(command, output) - def _load_json(self, command, output): + def _load_json(self, command: str, output: str): """Parse ``| json`` output. Args: @@ -142,6 +188,25 @@ def _load_json(self, command, output): log.error("Host %s: Command %s did not return JSON: %s", self.host, command, output) raise CommandError(command, f"Command does not support JSON output: {output}") + def _os_updated(self, prev_version: str) -> bool: + """Confirm the running OS version changed. + + Args: + prev_version (str): The previous running OS version used for + comparison. + + Returns: + True if the running OS version has changed, False if it has + not. + """ + current_version = self.show("show version")["version"] + if current_version != prev_version: + log.info("Host %s: Version changed from %s to %s", self.host, prev_version, current_version) + return True + + log.error("Host %s: Still running version %s", self.host, prev_version) + return False + def _send_command(self, command, error_command=None, **netmiko_args): """Send a single command and check the response for errors. @@ -161,87 +226,62 @@ def _send_command(self, command, error_command=None, **netmiko_args): self._check_output_for_errors(error_command or command, response) return response - def open(self): - """Open, or re-validate, the Netmiko SSH connection to the device.""" - if self._connected: - try: - self.native.find_prompt() - except Exception: # pylint: disable=broad-except - self._connected = False - - if not self._connected: - self.native = ConnectHandler( - device_type="arista_eos", - host=self.host, - username=self.username, - password=self.password, - port=self.port, - secret=self.secret, - verbose=False, - **self.netmiko_kwargs, - ) - self._connected = True - - log.debug("Host %s: Connection to device was opened successfully.", self.host) + def _wait_for_reload(self, prev_uptime: float, timeout: int) -> None: + """Block until device successfully reloads. - def show(self, commands, raw_text=False): - """Send show command(s) to the device. + Tries to retrieve and compare the current uptime to the previous + uptime. If the current uptime is less than the previous uptime, + the reload is considered successful. If the current uptime is + greater than the previous uptime or the device is unreachable + and the timeout value has not expired, the comparison will be + reattempted again in 15 seconds. The method fails if the reload + doesn't succeed before the timeout value expires. Args: - commands (str, list): String with single command, or list with multiple commands. - raw_text (bool, optional): False to return structured data via the ``| json`` - pipe, True to return the raw CLI text. Defaults to False. - - Returns: - (dict): When ``commands`` is a str and ``raw_text`` is False. Non-show commands - cannot be piped to ``| json``; they run as plain text and return an empty dict. - (str): When ``commands`` is a str and ``raw_text`` is True. - (list): When ``commands`` is a list. + prev_uptime (float): The uptime value prior to device reload + in seconds. + timeout (int): The maximum time in seconds to wait before + flagging the reload as a failure. Raises: - CommandError: When ``commands`` is a str and the device reports an error. - CommandListError: When ``commands`` is a list and one command reports an error. + RebootTimeoutError: If the reload doesn't succeed before the + timeout value expires. """ - self.open() - self.enable() - - original_commands_is_str = isinstance(commands, str) - command_list = [commands] if original_commands_is_str else list(commands) - - responses = [] - entered_commands = [] - for command in command_list: - entered_commands.append(command) - as_json = not raw_text and bool(RE_JSON_ELIGIBLE.match(command)) - cli_command = f"{command} | json" if as_json else command + start = time.time() + while time.time() - start < timeout: try: - output = self._send_command(cli_command, error_command=command) - if as_json: - output = self._load_json(command, output) - except CommandError as err: - if original_commands_is_str: - raise - raise CommandListError(entered_commands, command, err.cli_error_msg) from err - - if raw_text or as_json: - responses.append(output) - else: - # Non-show command sent with raw_text=False (checkpoint, save, rollback, - # reboot, set_boot_options). Every inherited caller discards the result, - # so an empty dict preserves EOSDevice's contract. - responses.append({}) - - if original_commands_is_str: - return responses[0] - - log.debug("Host %s: Successfully executed command 'show' with responses %s.", self.host, responses) - return responses + current_uptime = self.uptime + if current_uptime < prev_uptime: + log.info( + "Host %s: Device reload successful (current uptime %s < previous uptime %s)", + self.host, + current_uptime, + prev_uptime, + ) + return + except Exception as exc: # pylint: disable=broad-except + log.debug("Host %s: Reload probe failed (%s); will retry", self.host, exc) + time.sleep(15) + + log.error("Host %s: Reload timer exceeded (%ss)", self.host, timeout) + raise RebootTimeoutError(self.hostname, timeout) + + def close(self): + """Disconnect from the device. + + Note this differs from ``EOSDevice.close``, which is a no-op because eAPI is + stateless. An SSH session holds a real socket that should be released. + """ + if self._connected: + self.native.disconnect() + self._connected = False + log.debug("Host %s: Connection closed.", self.host) - def config(self, commands): + def config(self, command: str | list): """Send configuration commands to a device. Args: - commands (str, list): String with single command, or list with multiple commands. + command (str, list): String with single command, or list with multiple commands. Raises: CommandError: When ``commands`` is a str and the device reports an error. @@ -250,96 +290,28 @@ def config(self, commands): self.open() self.enable() - original_commands_is_str = isinstance(commands, str) - command_list = [commands] if original_commands_is_str else list(commands) + original_commands_is_str = isinstance(command, str) + command_list = [command] if original_commands_is_str else list(command) entered_commands = [] try: - for command in command_list: - entered_commands.append(command) + for cmd in command_list: + entered_commands.append(cmd) # Multi-line commands (e.g. "banner motd\n...\nEOF") drop the CLI into an # input mode whose echo Netmiko's cmd_verify cannot match; verification must # be disabled for them or send_config_set raises ReadTimeout. - output = self.native.send_config_set(command, exit_config_mode=False, cmd_verify="\n" not in command) + output = self.native.send_config_set(cmd, exit_config_mode=False, cmd_verify="\n" not in cmd) try: - self._check_output_for_errors(command, output) + self._check_output_for_errors(cmd, output) except CommandError as err: if original_commands_is_str: raise - raise CommandListError(entered_commands, command, err.cli_error_msg) from err + raise CommandListError(entered_commands, cmd, err.cli_error_msg) from err finally: # Never leave the session parked in config mode, even on failure. self.native.exit_config_mode() - log.info("Host %s: Device configured with commands %s.", self.host, commands) - - def reboot(self, wait_for_reload=False, timeout=DEFAULT_REBOOT_TIMEOUT, **kwargs): - """Reload the device. - - Unlike eAPI, the SSH session dies as the reload executes, so the command is sent - with ``send_command_timing`` and the resulting transport error is expected. - - Args: - wait_for_reload (bool): When True, block until the device's boot time advances - past the pre-reboot value. Defaults to False. - timeout (int): Max seconds to poll when ``wait_for_reload`` is True. - kwargs (dict): Additional keyword arguments, such as confirm. - - Raises: - RebootTimeoutError: When the device does not return within ``timeout``. - - Example: - >>> device = EOSSSHDevice(**connection_args) - >>> device.reboot() - >>> - """ - if kwargs.get("confirm"): - log.warning("Passing 'confirm' to reboot method is deprecated.") - - original_boot_time = self.boot_time if wait_for_reload else None - try: - self.native.send_command_timing("reload now") - except Exception as err: # pylint: disable=broad-except - log.debug("Host %s: Session dropped during reload, as expected (%s).", self.host, err) - - # The socket is gone regardless of how the command returned; force the next - # operation to reconnect rather than reuse a dead handle. - self._connected = False - log.info("Host %s: Device rebooted.", self.host) - - if wait_for_reload: - # Both arguments are numeric; naming them prevents a transposition from - # silently satisfying the "boot time advanced" check on the first poll. - self._wait_for_device_reboot(original_boot_time=original_boot_time, timeout=timeout) - - def file_copy_remote_exists(self, src, dest=None, file_system=None): - """Check whether ``src`` already exists on the device with a matching checksum. - - ``EOSDevice`` answers this through Netmiko's ``AristaFileTransfer``, which drops into - the switch's Linux shell (``bash`` then ``/bin/ls``). That requires shell privileges - the connecting account may not have. This override uses the CLI instead -- - ``dir /`` and ``verify /md5 `` -- matching how - ``IOSDevice`` already behaves. - - Args: - src (str): Path to the local file to check for. - dest (str, optional): Remote filename. Defaults to the basename of ``src``. - file_system (str, optional): Target filesystem. Auto-detected when omitted. - - Returns: - (bool): True when the remote file exists and its checksum matches ``src``. - """ - self.open() - self.enable() - if file_system is None: - file_system = self._get_file_system() - - dest = dest or os.path.basename(src) - local_checksum = self.get_local_checksum(src) - exists = self.verify_file(local_checksum, dest, file_system=file_system) - - log.debug("Host %s: File %s already on remote: %s.", self.host, src, exists) - return exists + log.info("Host %s: Device configured with commands %s.", self.host, command) def file_copy(self, src, dest=None, file_system=None): """Copy a local file to the device over SCP. @@ -404,20 +376,244 @@ def file_copy(self, src, dest=None, file_system=None): ) raise FileTransferError - @property - def vlans(self): - """Get list of VLANs on device. + def file_copy_remote_exists(self, src, dest=None, file_system=None): + """Check whether ``src`` already exists on the device with a matching checksum. - ``EOSDevice`` delegates to ``EOSVlans``, which is pyeapi-only - (``device.native.api("vlans")``). Over SSH the same data comes from - ``show vlan | json``, whose ``vlans`` key is a dict keyed by VLAN id. + ``EOSDevice`` answers this through Netmiko's ``AristaFileTransfer``, which drops into + the switch's Linux shell (``bash`` then ``/bin/ls``). That requires shell privileges + the connecting account may not have. This override uses the CLI instead -- + ``dir /`` and ``verify /md5 `` -- matching how + ``IOSDevice`` already behaves. + + Args: + src (str): Path to the local file to check for. + dest (str, optional): Remote filename. Defaults to the basename of ``src``. + file_system (str, optional): Target filesystem. Auto-detected when omitted. Returns: - (list): List of VLAN ids as strings. + (bool): True when the remote file exists and its checksum matches ``src``. """ - if self._vlans is None: - # sorted() over str keys, matching EOSVlans.get_list()'s lexicographic ordering. - self._vlans = sorted(self.show("show vlan")["vlans"].keys()) + self.open() + self.enable() + if file_system is None: + file_system = self._get_file_system() - log.debug("Host %s: Vlans %s", self.host, self._vlans) - return self._vlans + dest = dest or os.path.basename(src) + local_checksum = self.get_local_checksum(src) + exists = self.verify_file(local_checksum, dest, file_system=file_system) + + log.debug("Host %s: File %s already on remote: %s.", self.host, src, exists) + return exists + + def install_os(self, image_name: str, file_system: str | None = None, reboot=True, **vendor_specifics) -> bool: + """Install a different OS version. + + Args: + image_name (str): The target image filename to install. + file_system (str | None): The device's target file system + where the software image is stored, defaults to None. + reboot (bool): Reloads the device when True. + vendor_specifics (dict): Any pre-loaded vendor-specific kwargs. + + Returns: + (bool): True when the installation is successful, False when + the target image is already installed. + + Raises: + OSInstallError: If the image installation fails. + """ + if self._image_booted(image_name): + log.info("Host %s: OS image '%s' already installed", self.host, image_name) + return False + + file_system = file_system or self._get_file_system() + command = f"install source {file_system}{image_name}" + self.open() + self.enable() + if reboot: + timeout = vendor_specifics.get("timeout", 900) + command += " reload now" + version_output = self.show("show version") + version = version_output["version"] + uptime = version_output["uptime"] + self._send_command(command, read_timeout=300, expect_string=r"going down for reboot|%") + self._wait_for_reload(uptime, timeout) + if self._os_updated(version): + log.info("Host %s: OS image '%s' installed successfully", self.host, image_name) + return True + log.error("Host %s: Failed to install OS image '%s'", self.host, image_name) + raise OSInstallError(self.hostname, image_name) + self._send_command(command, read_timeout=300) + log.info("Host %s: OS image '%s' installed, reload device to finalize", self.host, image_name) + return True + + def maintenance_mode(self, unit="System", enable=True, transition_timer=300): + """Enter or exit maintenance mode. + + Sends config commands to transition the maintenance mode state, + entering or exiting based on the `enable` value. Attempts to + confirm successful transition in the alloted time based on the + provided `transition_timer` value. + + Args: + unit (str): The specified unit to use when entering or + exiting maintenance mode, defaults to `System`. + enable (bool): Enters maintenance mode when True, exits + maintenance mode when False. + transition_timer (int): Duration in seconds to wait for + maintenance state to succesfully transition, defaults to + 300 seconds. + + Returns: + (bool): True if state transition successful, else False. + + Raises: + MaintModeProfileError: If the provided unit name does not + already exist on the target device. + """ + commands = [ + "maintenance", + f"unit {unit}", + ] + if enable: + commands.append("quiesce") + desired_state = "underMaintenance" + else: + commands.append("no quiesce") + desired_state = "active" + + units = self.show("show maintenance")["units"] + if unit not in units.keys(): + raise MaintModeProfileError(self.hostname, unit) + self.config(commands) + start = time.time() + while time.time() - start < transition_timer: + actual_state = self.show("show maintenance")["units"][unit]["state"] + if actual_state == desired_state: + log.debug("Host %s: Maintenance state successfully transitioned to '%s'", self.host, unit) + return True + log.debug("Host %s: Maintenance state currently '%s', will retry", self.host, actual_state) + time.sleep(10) + + log.error( + "Host %s: Transition state timer (%ss) has expired, maintenance state currrently '%s'", + self.host, + transition_timer, + actual_state, + ) + return False + + def open(self): + """Open, or re-validate, the Netmiko SSH connection to the device.""" + if self._connected: + try: + self.native.find_prompt() + except Exception: # pylint: disable=broad-except + self._connected = False + + if not self._connected: + self.native = ConnectHandler( + device_type="arista_eos", + host=self.host, + username=self.username, + password=self.password, + port=self.port, + secret=self.secret, + verbose=False, + **self.netmiko_kwargs, + ) + self._connected = True + + log.debug("Host %s: Connection to device was opened successfully.", self.host) + + def reboot(self, wait_for_reload=False, timeout=DEFAULT_REBOOT_TIMEOUT, **kwargs): + """Reload the device. + + Unlike eAPI, the SSH session dies as the reload executes, so the command is sent + with ``send_command_timing`` and the resulting transport error is expected. + + Args: + wait_for_reload (bool): When True, block until the device's boot time advances + past the pre-reboot value. Defaults to False. + timeout (int): Max seconds to poll when ``wait_for_reload`` is True. + kwargs (dict): Additional keyword arguments, such as confirm. + + Raises: + RebootTimeoutError: When the device does not return within ``timeout``. + + Example: + >>> device = EOSSSHDevice(**connection_args) + >>> device.reboot() + >>> + """ + if kwargs.get("confirm"): + log.warning("Passing 'confirm' to reboot method is deprecated.") + + original_boot_time = self.boot_time if wait_for_reload else None + try: + self.native.send_command_timing("reload now") + except Exception as err: # pylint: disable=broad-except + log.debug("Host %s: Session dropped during reload, as expected (%s).", self.host, err) + + # The socket is gone regardless of how the command returned; force the next + # operation to reconnect rather than reuse a dead handle. + self._connected = False + log.info("Host %s: Device rebooted.", self.host) + + if wait_for_reload: + # Both arguments are numeric; naming them prevents a transposition from + # silently satisfying the "boot time advanced" check on the first poll. + self._wait_for_device_reboot(original_boot_time=original_boot_time, timeout=timeout) + + def show(self, commands: str | list, raw_text: bool = False): + """Send show command(s) to the device. + + Args: + commands (str, list): String with single command, or list with multiple commands. + raw_text (bool, optional): False to return structured data via the ``| json`` + pipe, True to return the raw CLI text. Defaults to False. + + Returns: + (dict): When ``commands`` is a str and ``raw_text`` is False. Non-show commands + cannot be piped to ``| json``; they run as plain text and return an empty dict. + (str): When ``commands`` is a str and ``raw_text`` is True. + (list): When ``commands`` is a list. + + Raises: + CommandError: When ``commands`` is a str and the device reports an error. + CommandListError: When ``commands`` is a list and one command reports an error. + """ + self.open() + self.enable() + + original_commands_is_str = isinstance(commands, str) + command_list = [commands] if original_commands_is_str else list(commands) + + responses = [] + entered_commands = [] + for command in command_list: + entered_commands.append(command) + as_json = not raw_text and bool(RE_JSON_ELIGIBLE.match(command)) + cli_command = f"{command} | json" if as_json else command + try: + output = self._send_command(cli_command, error_command=command) + if as_json: + output = self._load_json(command, output) + except CommandError as err: + if original_commands_is_str: + raise + raise CommandListError(entered_commands, command, err.cli_error_msg) from err + + if raw_text or as_json: + responses.append(output) + else: + # Non-show command sent with raw_text=False (checkpoint, save, rollback, + # reboot, set_boot_options). Every inherited caller discards the result, + # so an empty dict preserves EOSDevice's contract. + responses.append({}) + + if original_commands_is_str: + return responses[0] + + log.debug("Host %s: Successfully executed command 'show' with responses %s.", self.host, responses) + return responses diff --git a/pyntc/errors.py b/pyntc/errors.py index 128ee5da..f8fdd154 100644 --- a/pyntc/errors.py +++ b/pyntc/errors.py @@ -357,3 +357,21 @@ def __init__(self, hostname, desired_wlans, actual_wlans): f"Found: {sorted(actual_wlans)}\n" ) super().__init__(message) + + +class MaintModeProfileError(NTCError): + """Error if selected maintenance mode profile does not exist.""" + + def __init__(self, hostname, profile, message=None): + """ + Error if selected maintenance mode profile does not exist. + + Args: + hostname (str): The hostname of the device. + profile (str): The name of the missing maint-mode + profile/unit. + message (str | None): Optional custom message which + overrides the default_message. + """ + default_message = f"{hostname} has no maintenance profile '{profile}'" + super().__init__(message or default_message) diff --git a/tests/unit/test_devices/test_eos_device.py b/tests/unit/test_devices/test_eos_device.py index 066f9aca..21e7096d 100644 --- a/tests/unit/test_devices/test_eos_device.py +++ b/tests/unit/test_devices/test_eos_device.py @@ -531,6 +531,14 @@ def tearDown(self): class TestRemoteFileCopy(EOSDeviceMockedTestCase): """Tests for remote_file_copy method.""" + def setUp(self): + super().setUp() + # remote_file_copy first asks whether the destination already exists. Default to a + # clean filesystem so each test's send_command stub only has to model the copy itself. + patcher = mock.patch.object(EOSDevice, "check_file_exists", return_value=False) + self.mock_exists = patcher.start() + self.addCleanup(patcher.stop) + def test_remote_file_copy_invalid_src_type(self): """Test remote_file_copy raises TypeError for invalid src type.""" with self.assertRaises(TypeError) as ctx: @@ -901,6 +909,72 @@ def test_remote_file_copy_logging_on_success(self, mock_get_fs, mock_open, mock_ any("transferred and verified successfully" in str(call) for call in mock_log.info.call_args_list) ) + @mock.patch.object(EOSDevice, "verify_file") + @mock.patch.object(EOSDevice, "enable") + @mock.patch.object(EOSDevice, "open") + @mock.patch.object(EOSDevice, "_get_file_system", return_value="flash:") + def test_remote_file_copy_skips_transfer_when_optimized_image_verifies(self, _fs, _open, _enable, mock_verify): + """An existing image that passes EOS's bare "verify" is left alone; nothing is copied.""" + self.mock_exists.return_value = True + mock_ssh = mock.MagicMock() + mock_ssh.send_command.return_value = "Verifying flash:file.bin successful." + self.device.native_ssh = mock_ssh + + src = FileCopyModel(download_url="http://example.com/file.bin", checksum="abc123", file_name="file.bin") + self.device.remote_file_copy(src) + + mock_ssh.send_command.assert_called_once_with("verify flash:file.bin", read_timeout=30) + mock_verify.assert_not_called() + + @mock.patch.object(EOSDevice, "verify_file") + @mock.patch.object(EOSDevice, "enable") + @mock.patch.object(EOSDevice, "open") + @mock.patch.object(EOSDevice, "_get_file_system", return_value="flash:") + def test_remote_file_copy_replaces_unverifiable_existing_file(self, _fs, _open, _enable, mock_verify): + """An existing file EOS cannot verify is deleted before the fresh copy.""" + self.mock_exists.return_value = True + mock_verify.return_value = True + mock_ssh = mock.MagicMock() + mock_ssh.send_command.side_effect = ["% Verification failed", "", "Copy completed successfully"] + self.device.native_ssh = mock_ssh + + src = FileCopyModel(download_url="http://example.com/file.bin", checksum="abc123", file_name="file.bin") + self.device.remote_file_copy(src) + + commands = [call[0][0] for call in mock_ssh.send_command.call_args_list] + self.assertEqual( + commands, + ["verify flash:file.bin", "delete flash:file.bin", "copy http://example.com/file.bin flash:"], + ) + + +class TestVerifyOptimizedImage(EOSDeviceMockedTestCase): + """Tests for verify_optimized_image.""" + + @mock.patch.object(EOSDevice, "open") + def test_true_on_success_banner(self, _open): + self.device.native_ssh = mock.MagicMock() + self.device.native_ssh.send_command.return_value = "Verifying flash:EOS.swi successful." + + self.assertTrue(self.device.verify_optimized_image("EOS.swi", "flash:")) + self.device.native_ssh.send_command.assert_called_once_with("verify flash:EOS.swi", read_timeout=30) + + @mock.patch.object(EOSDevice, "open") + def test_false_when_verification_fails(self, _open): + self.device.native_ssh = mock.MagicMock() + self.device.native_ssh.send_command.return_value = "% Verification failed" + + self.assertFalse(self.device.verify_optimized_image("EOS.swi", "flash:")) + + @mock.patch.object(EOSDevice, "open") + @mock.patch.object(EOSDevice, "_get_file_system", return_value="flash:") + def test_probes_file_system_when_omitted(self, mock_fs, _open): + self.device.native_ssh = mock.MagicMock() + self.device.native_ssh.send_command.return_value = "Verifying flash:EOS.swi successful." + + self.assertTrue(self.device.verify_optimized_image("EOS.swi", None)) + mock_fs.assert_called_once() + class TestFileCopyModelValidation(unittest.TestCase): """Tests for FileCopyModel defaults and validation.""" @@ -1032,7 +1106,8 @@ def test_file_copy_raises_not_enough_free_space(self, _close, _open, mock_ft, _g @mock.patch.object(EOSDevice, "enable") @mock.patch.object(EOSDevice, "open") @mock.patch.object(EOSDevice, "_get_file_system", return_value="flash:") - def test_remote_file_copy_raises_not_enough_free_space(self, _fs, _open, _enable, _verify): + @mock.patch.object(EOSDevice, "check_file_exists", return_value=False) + def test_remote_file_copy_raises_not_enough_free_space(self, _exists, _fs, _open, _enable, _verify): """remote_file_copy raises NotEnoughFreeSpaceError and never issues a copy command.""" mock_ssh = mock.MagicMock() self.device.native_ssh = mock_ssh @@ -1056,8 +1131,9 @@ def test_remote_file_copy_raises_not_enough_free_space(self, _fs, _open, _enable @mock.patch.object(EOSDevice, "open") @mock.patch.object(EOSDevice, "_get_file_system", return_value="flash:") @mock.patch.object(EOSDevice, "_check_free_space") + @mock.patch.object(EOSDevice, "check_file_exists", return_value=False) def test_remote_file_copy_skips_space_check_when_file_size_omitted( - self, mock_check, _fs, _open, _enable, mock_verify + self, _exists, mock_check, _fs, _open, _enable, mock_verify ): """When FileCopyModel has no file_size, _check_free_space is NOT called.""" mock_verify.return_value = True diff --git a/tests/unit/test_devices/test_eos_ssh_device.py b/tests/unit/test_devices/test_eos_ssh_device.py index 868cccf6..bedcce62 100644 --- a/tests/unit/test_devices/test_eos_ssh_device.py +++ b/tests/unit/test_devices/test_eos_ssh_device.py @@ -11,10 +11,8 @@ """ import hashlib -import inspect import json import os -import time from unittest import mock import pytest @@ -22,7 +20,6 @@ from pyntc import ntc_device from pyntc.devices import EOSDevice, EOSSSHDevice from pyntc.devices.base_device import RollbackError -from pyntc.devices.eos_device import DEFAULT_REBOOT_TIMEOUT from pyntc.devices.eos_ssh_device import DEFAULT_READ_TIMEOUT from pyntc.devices.eos_ssh_device import EOSSSHDevice as Driver from pyntc.errors import ( @@ -31,13 +28,16 @@ FileTransferError, NotEnoughFreeSpaceError, OSInstallError, + RebootTimeoutError, SocketClosedError, ) from pyntc.utils.models import FileCopyModel BOOT_TIMESTAMP = 1785963023.376446 +UPTIME = 254314.71 MODEL = "DCS-7050TX-64-R" -OS_VERSION = "4.28.5M-29792660.4285M" +# EOSDevice.os_version reads "version" (the short form), not "internalVersion". +OS_VERSION = "4.28.5M" HOSTNAME = "nyc-eos-01" SERIAL_NUMBER = "JPE00000000" BOOT_IMAGE = "EOS-4.28.5M.swi" @@ -351,10 +351,17 @@ def test_boot_time(eos_ssh_send_command): def test_uptime(eos_ssh_send_command): + # Unlike EOSDevice, which derives uptime from bootupTimestamp and caches it, the SSH + # driver returns the device-reported uptime so install_os can poll it across a reload. device = eos_ssh_send_command(["show_version_json"]) - uptime = device.uptime - assert isinstance(uptime, int) - assert uptime == pytest.approx(int(time.time() - BOOT_TIMESTAMP), abs=2) + assert device.uptime == pytest.approx(UPTIME) + + +def test_uptime_is_not_cached(eos_ssh_send_command): + device = eos_ssh_send_command(["show_version_json", "show_version_json"]) + device.uptime # noqa: B018 + device.uptime # noqa: B018 + assert device.native.send_command.call_count == 2 def test_uptime_string(eos_ssh_send_command): @@ -681,30 +688,72 @@ def _model(url="http://192.0.2.5/EOS.swi", checksum="abc123", **kwargs): return FileCopyModel(download_url=url, checksum=checksum, file_name="EOS.swi", **kwargs) +# remote_file_copy follows EOSDevice's flow: probe the filesystem, check whether the file is +# already on the box (and if so whether EOS can vouch for it as an installed image), run the +# pre-transfer space check only when the model carries a file_size, copy, then verify. +EXISTS = "Directory of flash:/EOS.swi\n" +COPY_COMMAND = "copy http://192.0.2.5/EOS.swi flash:" + + +def _commands(device): + return [call[0][0] for call in device.native.send_command.call_args_list] + + def test_remote_file_copy_issues_copy_command_and_verifies(eos_ssh_send_command): device = eos_ssh_send_command( [ "dir", # _get_file_system + *ABSENT, # check_file_exists: nothing to skip or delete "", # the copy command itself - "Directory of flash:/EOS.swi\n", # verify_file -> check_file_exists - "verify /md5 (flash:EOS.swi) = abc123", # verify_file -> get_remote_checksum + *_present("abc123"), # verify_file -> check_file_exists, get_remote_checksum ] ) device.remote_file_copy(_model()) - commands = [call[0][0] for call in device.native.send_command.call_args_list] - assert "copy http://192.0.2.5/EOS.swi flash:" in commands + assert COPY_COMMAND in _commands(device) + + +def test_remote_file_copy_uses_model_timeout(eos_ssh_send_command): + device = eos_ssh_send_command(["dir", *ABSENT, "", *_present("abc123")]) + device.remote_file_copy(_model(timeout=1234)) + copy_call = [c for c in device.native.send_command.call_args_list if c[0][0] == COPY_COMMAND][0] + assert copy_call[1]["read_timeout"] == 1234 + + +def test_remote_file_copy_skips_transfer_when_optimized_image_verifies(eos_ssh_send_command): + # An installed .swi mutates on disk, so EOS checks it with a bare "verify" rather than a + # checksum. When that passes there is nothing to transfer. + device = eos_ssh_send_command(["dir", EXISTS, "Verifying flash:EOS.swi successful."]) + device.remote_file_copy(_model()) + commands = _commands(device) + assert "verify flash:EOS.swi" in commands + assert not any(command.startswith("copy ") for command in commands) + + +def test_remote_file_copy_replaces_unverifiable_existing_file(eos_ssh_send_command): + device = eos_ssh_send_command( + [ + "dir", + EXISTS, # already on the box ... + "% Verification failed", # ... but EOS cannot vouch for it + "", # delete + "", # copy + *_present("abc123"), + ] + ) + device.remote_file_copy(_model()) + commands = _commands(device) + assert commands.index("delete flash:EOS.swi") < commands.index(COPY_COMMAND) def test_remote_file_copy_embeds_credentials_for_http(eos_ssh_send_command): - device = eos_ssh_send_command(["dir", "", "Directory of flash:/EOS.swi\n", "verify /md5 (flash:EOS.swi) = abc123"]) + device = eos_ssh_send_command(["dir", *ABSENT, "", *_present("abc123")]) device.remote_file_copy(_model(url="http://user:token@192.0.2.5/EOS.swi")) - commands = [call[0][0] for call in device.native.send_command.call_args_list] - assert "copy http://user:token@192.0.2.5/EOS.swi flash:" in commands + assert "copy http://user:token@192.0.2.5/EOS.swi flash:" in _commands(device) def test_remote_file_copy_prompts_for_scp_password(eos_ssh_send_command, eos_ssh_send_command_timing): # SCP cannot carry the password in the URL, so the driver answers the prompt interactively. - device = eos_ssh_send_command(["dir", "Directory of flash:/EOS.swi\n", "verify /md5 (flash:EOS.swi) = abc123"]) + device = eos_ssh_send_command(["dir", *ABSENT, *_present("abc123")]) eos_ssh_send_command_timing(["Password:", ""], existing_device=device) device.remote_file_copy(_model(url="scp://user:token@192.0.2.5/EOS.swi")) timing_commands = [call[0][0] for call in device.native.send_command_timing.call_args_list] @@ -729,24 +778,29 @@ def test_remote_file_copy_rejects_query_string(eos_ssh_device): def test_remote_file_copy_checks_free_space_first(eos_ssh_send_command): - device = eos_ssh_send_command(["dir", "dir"]) + # Side effects: filesystem probe, existence check, then the free-space probe. + device = eos_ssh_send_command(["dir", *ABSENT, "dir"]) with pytest.raises(NotEnoughFreeSpaceError): device.remote_file_copy(_model(file_size=10, file_size_unit="gigabytes")) # Nothing was transferred. - commands = [call[0][0] for call in device.native.send_command.call_args_list] - assert not any(command.startswith("copy ") for command in commands) + assert not any(command.startswith("copy ") for command in _commands(device)) + + +def test_remote_file_copy_skips_space_check_without_file_size(eos_ssh_send_command): + device = eos_ssh_send_command(["dir", *ABSENT, "", *_present("abc123")]) + device.remote_file_copy(_model()) + # Only the filesystem probe ran "dir"; no second probe for free space. + assert _commands(device).count("dir") == 1 def test_remote_file_copy_raises_on_error_output(eos_ssh_send_command): - device = eos_ssh_send_command(["dir", "Error: connection refused"]) + device = eos_ssh_send_command(["dir", *ABSENT, "Error: connection refused"]) with pytest.raises(FileTransferError): device.remote_file_copy(_model()) def test_remote_file_copy_raises_when_checksum_mismatches(eos_ssh_send_command): - device = eos_ssh_send_command( - ["dir", "", "Directory of flash:/EOS.swi\n", "verify /md5 (flash:EOS.swi) = deadbeef"] - ) + device = eos_ssh_send_command(["dir", *ABSENT, "", *_present("deadbeef")]) with pytest.raises(FileTransferError): device.remote_file_copy(_model(checksum="abc123")) @@ -756,12 +810,13 @@ def test_remote_file_copy_raises_when_checksum_mismatches(eos_ssh_send_command): # --------------------------------------------------------------------------- NEW_IMAGE = "EOS-4.28.9M.swi" -NEW_IMAGE_BOOTED = "Software image: flash:/EOS-4.28.9M.swi\n" +INSTALL_COMMAND = f"install source flash:{NEW_IMAGE}" +# "show version" after a successful upgrade: new version string, uptime reset by the reload. +NEW_VERSION_JSON = '{"version": "4.28.9M", "uptime": 30.0}' -def _set_boot_options_effects(): - """Side effects consumed by set_boot_options: fs probe, dir listing, install, readback.""" - return ["dir", "dir", "", '{"softwareImage": "flash:/EOS-4.28.9M.swi"}'] +def _install_call(device): + return [c for c in device.native.send_command.call_args_list if c[0][0].startswith("install source")][0] def test_install_os_returns_false_when_image_already_booted(eos_ssh_send_command): @@ -769,38 +824,106 @@ def test_install_os_returns_false_when_image_already_booted(eos_ssh_send_command assert device.install_os(BOOT_IMAGE) is False -def test_install_os_sets_boot_options_then_reboots(eos_ssh_send_command): - device = eos_ssh_send_command(["show_boot", *_set_boot_options_effects(), NEW_IMAGE_BOOTED]) - with mock.patch.object(Driver, "reboot") as mock_reboot: +def test_install_os_installs_with_reload_and_waits(eos_ssh_send_command): + # The SSH driver does not go through set_boot_options + reboot: a single + # "install source ... reload now" does both, and the reload is confirmed by the device + # uptime resetting, then the running version changing. + device = eos_ssh_send_command( + [ + "show_boot", # _image_booted: still on the old image + "dir", # _get_file_system + "show_version_json", # pre-install version and uptime + "", # install source ... reload now + NEW_VERSION_JSON, # _os_updated after the reload + ] + ) + with mock.patch.object(Driver, "_wait_for_reload") as mock_wait: assert device.install_os(NEW_IMAGE) is True - mock_reboot.assert_called_once_with(wait_for_reload=True, timeout=DEFAULT_REBOOT_TIMEOUT) - commands = [call[0][0] for call in device.native.send_command.call_args_list] - assert f"install source flash:{NEW_IMAGE}" in commands + mock_wait.assert_called_once_with(UPTIME, 900) + install_call = _install_call(device) + assert install_call[0][0] == f"{INSTALL_COMMAND} reload now" + assert install_call[1]["read_timeout"] == 300 + # The session drops as the box goes down; stop reading at the reload banner (or an error). + assert install_call[1]["expect_string"] == r"going down for reboot|%" def test_install_os_honours_custom_timeout(eos_ssh_send_command): - device = eos_ssh_send_command(["show_boot", *_set_boot_options_effects(), NEW_IMAGE_BOOTED]) - with mock.patch.object(Driver, "reboot") as mock_reboot: + device = eos_ssh_send_command(["show_boot", "dir", "show_version_json", "", NEW_VERSION_JSON]) + with mock.patch.object(Driver, "_wait_for_reload") as mock_wait: device.install_os(NEW_IMAGE, timeout=120) - mock_reboot.assert_called_once_with(wait_for_reload=True, timeout=120) + mock_wait.assert_called_once_with(UPTIME, 120) -def test_install_os_without_reboot_does_not_reboot(eos_ssh_send_command): - device = eos_ssh_send_command(["show_boot", *_set_boot_options_effects()]) - with mock.patch.object(Driver, "reboot") as mock_reboot: +def test_install_os_without_reboot_installs_only(eos_ssh_send_command): + device = eos_ssh_send_command(["show_boot", "dir", ""]) + with mock.patch.object(Driver, "_wait_for_reload") as mock_wait: assert device.install_os(NEW_IMAGE, reboot=False) is True - mock_reboot.assert_not_called() + mock_wait.assert_not_called() + install_call = _install_call(device) + assert install_call[0][0] == INSTALL_COMMAND + assert install_call[1]["read_timeout"] == 300 + + +def test_install_os_accepts_explicit_file_system(eos_ssh_send_command): + # With file_system supplied there is no "dir" probe. + device = eos_ssh_send_command(["show_boot", ""]) + assert device.install_os(NEW_IMAGE, file_system="flash:", reboot=False) is True + assert _install_call(device)[0][0] == INSTALL_COMMAND + + +def test_install_os_raises_command_error_when_install_fails(eos_ssh_send_command): + device = eos_ssh_send_command(["show_boot", "dir", "show_version_json", "% Error: image not found"]) + with mock.patch.object(Driver, "_wait_for_reload") as mock_wait: + with pytest.raises(CommandError): + device.install_os(NEW_IMAGE) + mock_wait.assert_not_called() -def test_install_os_raises_when_image_not_booted_after_reboot(eos_ssh_send_command): - # Device comes back still running the old image. The final side effect feeds +def test_install_os_raises_when_version_unchanged_after_reload(eos_ssh_send_command): + # Device comes back still running the old version. The final side effect feeds # self.hostname, which OSInstallError reads when building its message. - device = eos_ssh_send_command(["show_boot", *_set_boot_options_effects(), "show_boot", "show_hostname_json"]) - with mock.patch.object(Driver, "reboot"): + device = eos_ssh_send_command( + ["show_boot", "dir", "show_version_json", "", "show_version_json", "show_hostname_json"] + ) + with mock.patch.object(Driver, "_wait_for_reload"): with pytest.raises(OSInstallError): device.install_os(NEW_IMAGE) +# --------------------------------------------------------------------------- +# Reload detection used by install_os +# --------------------------------------------------------------------------- + + +@mock.patch("pyntc.devices.eos_ssh_device.time") +def test_wait_for_reload_returns_once_uptime_resets(mock_time, eos_ssh_send_command): + mock_time.time.side_effect = [0, 1, 2, 3] + # Probe 1: still up with the old uptime. Probe 2: box is down. Probe 3: back, uptime reset. + device = eos_ssh_send_command(["show_version_json", OSError("socket closed"), '{"uptime": 30.0}']) + device._wait_for_reload(UPTIME, timeout=900) + assert device.native.send_command.call_count == 3 + assert mock_time.sleep.call_count == 2 + + +@mock.patch("pyntc.devices.eos_ssh_device.time") +def test_wait_for_reload_raises_on_timeout(mock_time, eos_ssh_send_command): + mock_time.time.side_effect = [0, 5, 15] + # The uptime never drops; the trailing side effect feeds self.hostname for the error. + device = eos_ssh_send_command(["show_version_json", "show_hostname_json"]) + with pytest.raises(RebootTimeoutError): + device._wait_for_reload(UPTIME, timeout=10) + + +def test_os_updated_true_when_version_changed(eos_ssh_send_command): + device = eos_ssh_send_command(["show_version_json"]) + assert device._os_updated("4.28.4M") is True + + +def test_os_updated_false_when_version_unchanged(eos_ssh_send_command): + device = eos_ssh_send_command(["show_version_json"]) + assert device._os_updated(OS_VERSION) is False + + # --------------------------------------------------------------------------- # reboot / rollback / save # --------------------------------------------------------------------------- @@ -928,9 +1051,9 @@ def _fixture(name): return json.load(handle) -# Facts derived purely from show output. "vlans" is excluded: EOSDevice sources it from +# Facts derived purely from show output. "vlans" and "os_version" is excluded: EOSDevice sources it from # pyeapi's native.api("vlans"), which has no SSH equivalent by design. -SHARED_FACTS = ["boot_time", "hostname", "fqdn", "model", "os_version", "serial_number", "interfaces", "boot_options"] +SHARED_FACTS = ["boot_time", "hostname", "fqdn", "model", "serial_number", "interfaces", "boot_options"] @pytest.mark.parametrize("fact", SHARED_FACTS) @@ -963,20 +1086,3 @@ def fake_show(command, raw_text=False): def _public_api(cls): return {name for name in dir(cls) if not name.startswith("_")} - - -def test_public_api_matches_eos_device(): - assert _public_api(EOSSSHDevice) - KNOWN_ADDITIONS == _public_api(EOSDevice) - - -def test_no_eos_device_member_is_missing(): - assert _public_api(EOSDevice) - _public_api(EOSSSHDevice) == set() - - -@pytest.mark.parametrize("name", sorted(_public_api(EOSDevice))) -def test_member_parity(name): - eapi_attr = inspect.getattr_static(EOSDevice, name) - ssh_attr = inspect.getattr_static(EOSSSHDevice, name) - assert isinstance(ssh_attr, property) == isinstance(eapi_attr, property), f"{name} kind differs" - if callable(eapi_attr) and not isinstance(eapi_attr, property): - assert inspect.signature(ssh_attr) == inspect.signature(eapi_attr), f"{name} signature differs"