diff --git a/simple_backup/simple_backup.py b/simple_backup/simple_backup.py index 6e30a9a..a595373 100755 --- a/simple_backup/simple_backup.py +++ b/simple_backup/simple_backup.py @@ -155,7 +155,6 @@ class Backup: self._remote = False self._ssh = None self._password_auth = False - self._password = None self._removed_count = 0 def check_params(self, homedir: str = '') -> int: @@ -379,6 +378,18 @@ class Backup: return None + connected = None + + try: + connected = self._ssh_login(ssh, homedir) + + return connected + finally: + # Don't leave a half-open connection (e.g. after failed authentication) + if connected is None: + ssh.close() + + def _ssh_login(self, ssh: paramiko.client.SSHClient, homedir: str) -> Optional[paramiko.client.SSHClient]: try: ssh.load_host_keys(filename=f'{homedir}/.ssh/known_hosts') except FileNotFoundError: @@ -545,6 +556,13 @@ class Backup: finally: self._remove_temp_files() + def close(self) -> None: + """Close the SSH connection, if open""" + + if self._ssh is not None: + self._ssh.close() + self._ssh = None + def _remove_temp_files(self) -> None: for path in [self._inputs_path, self._exclude_path]: if path != '': @@ -665,11 +683,6 @@ class Backup: else: logger.warning('Backup not completed successfully. Old backups will not be removed') - if self._remote: - assert self._ssh is not None - - self._ssh.close() - # Files vanishing during the transfer is normal for a live system (return code 24) if returncode not in [0, 24]: logger.error( @@ -1012,12 +1025,15 @@ def simple_backup() -> int: backup = Backup(inputs, output, exclude, keep, rsync_options, ssh_host, ssh_user, ssh_keyfile, remote_sudo, remove_before=args.remove_before_backup, verbose=args.verbose) - return_code = backup.check_params(homedir) + try: + return_code = backup.check_params(homedir) - if return_code == 0: - return backup.run() + if return_code == 0: + return backup.run() - return return_code + return return_code + finally: + backup.close() if __name__ == '__main__':