From ee018d0bc078378b5586d61ba2a4d3717d9748eb Mon Sep 17 00:00:00 2001 From: Aleksei Slaikovskii Date: Dec 05 2017 11:41:28 +0000 Subject: Complete fix for OpenSSH transport and freeze fix for Paramiko one. --- diff --git a/pytest_multihost/host.py b/pytest_multihost/host.py index a095a43..ddabb71 100644 --- a/pytest_multihost/host.py +++ b/pytest_multihost/host.py @@ -27,9 +27,11 @@ class BaseHost(object): transport_class = transport.SSHTransport command_prelude = '' - def __init__(self, domain, hostname, role, ip=None, - external_hostname=None, username=None, password=None, - test_dir=None, host_type=None): + def __init__( + self, domain, hostname, role, + ip=None, external_hostname=None, username=None, password=None, + test_dir=None, host_type=None + ): self.host_type = host_type self.domain = domain self.role = str(role) @@ -51,16 +53,21 @@ class BaseHost(object): shortname, dot, ext_domain = hostname.partition('.') self.shortname = shortname - self.hostname = (hostname[:-1] - if hostname.endswith('.') - else shortname + '.' + self.domain.name) + self.hostname = ( + hostname[:-1] + if hostname.endswith('.') + else shortname + '.' + self.domain.name + ) self.external_hostname = str(external_hostname or hostname) self.netbios = self.domain.name.split('.')[0].upper() - self.logger_name = '%s.%s.%s' % ( - self.__module__, type(self).__name__, shortname) + self.logger_name = '{module}.{name}.{shortname}'.format( + module=self.__module__, + name=type(self).__name__, + shortname=shortname + ) self.log = self.config.get_logger(self.logger_name) if ip: @@ -69,7 +76,8 @@ class BaseHost(object): if self.config.ipv6: # $(dig +short $M $rrtype|tail -1) dig = subprocess.Popen( - ['dig', '+short', self.external_hostname, 'AAAA']) + ['dig', '+short', self.external_hostname, 'AAAA'] + ) stdout, stderr = dig.communicate() self.ip = stdout.splitlines()[-1].strip() else: @@ -79,8 +87,10 @@ class BaseHost(object): self.ip = None if not self.ip: - raise RuntimeError('Could not determine IP address of %s' % - self.external_hostname) + raise RuntimeError( + 'Could not determine IP address of {}'.format( + self.external_hostname) + ) self.host_key = None self.ssh_port = 22 @@ -94,8 +104,10 @@ class BaseHost(object): return template.format(s=self) def __repr__(self): - template = ('<{s.__module__}.{s.__class__.__name__} ' - '{s.hostname} ({s.role})>') + template = ( + '<{s.__module__}.{s.__class__.__name__} ' + '{s.hostname} ({s.role})>' + ) return template.format(s=self) def add_log_collector(self, collector): @@ -127,14 +139,16 @@ class BaseHost(object): password = dct.pop('password', None) host_type = dct.pop('host_type', 'default') - check_config_dict_empty(dct, 'host %s' % hostname) + check_config_dict_empty(dct, 'host {}'.format(hostname)) - return cls(domain, hostname, role, - ip=ip, - external_hostname=external_hostname, - username=username, - password=password, - host_type=host_type) + return cls( + domain, hostname, role, + ip=ip, + external_hostname=external_hostname, + username=username, + password=password, + host_type=host_type + ) def to_dict(self): """Export info about this Host to a dict""" @@ -199,9 +213,11 @@ class BaseHost(object): for collector in self.log_collectors: collector(self, filename) - def run_command(self, argv, set_env=True, stdin_text=None, - log_stdout=True, raiseonerr=True, - cwd=None, bg=False): + def run_command( + self, argv, set_env=True, + stdin_text=None, log_stdout=True, raiseonerr=True, + cwd=None, bg=False + ): """Run the given command on this host Returns a Command instance. The command will have already run in the @@ -224,15 +240,20 @@ class BaseHost(object): # Set working directory if cwd is None: cwd = self.test_dir - command.stdin.write('cd %s\n' % shell_quote(cwd)) + command.stdin.write('cd {}\n'.format(shell_quote(cwd))) # Set the environment if set_env: - command.stdin.write('. %s\n' % shell_quote(self.env_sh_path)) + command.stdin.write('. {}\n'.format(shell_quote(self.env_sh_path))) if self.command_prelude: command.stdin.write(self.command_prelude) + if stdin_text: + command.stdin.write('echo \'') + command.stdin.write(stdin_text) + command.stdin.write('\' | ') + if isinstance(argv, basestring): # Run a shell command given as a string command.stdin.write('(') @@ -244,16 +265,13 @@ class BaseHost(object): command.stdin.write(shell_quote(arg)) command.stdin.write(' ') - command.stdin.write(';exit\n') - if stdin_text: - command.stdin.write(stdin_text) + command.stdin.write('\nexit\n') command.stdin.flush() if not bg: command.wait(raiseonerr=raiseonerr) return command - class Host(BaseHost): """A Unix host""" command_prelude = 'set -e\n' diff --git a/pytest_multihost/transport.py b/pytest_multihost/transport.py index 8a36f02..8c1215e 100644 --- a/pytest_multihost/transport.py +++ b/pytest_multihost/transport.py @@ -11,6 +11,7 @@ OpenSSHTransport (if Paramiko is not importable, or the PYTESTMULTIHOST_SSH_TRANSPORT environment variable is set to "openssh"). """ +import base64 import os import socket import threading @@ -39,7 +40,9 @@ class Transport(object): """ def __init__(self, host): self.host = host - self.logger_name = '%s.%s' % (host.logger_name, type(self).__name__) + self.logger_name = '{}.{}'.format( + host.logger_name, type(self).__name__ + ) self.log = host.config.get_logger(self.logger_name) self._command_index = 0 @@ -96,7 +99,10 @@ class Transport(object): def get_next_command_logger_name(self): self._command_index += 1 - return '%s.cmd%s' % (self.host.logger_name, self._command_index) + return '{logger}.cmd{index}'.format( + logger=self.host.logger_name, + index=self._command_index + ) def rmdir(self, path): """Remove directory""" @@ -126,8 +132,12 @@ class Command(object): be strings containing the output, and ``returncode`` will contain the exit code. """ - def __init__(self, argv, logger_name=None, log_stdout=True, - get_logger=None): + def __init__( + self, argv, + logger_name=None, + log_stdout=True, + get_logger=None + ): self.returncode = None self.argv = argv self._done = False @@ -135,7 +145,9 @@ class Command(object): if logger_name: self.logger_name = logger_name else: - self.logger_name = '%s.%s' % (self.__module__, type(self).__name__) + self.logger_name = '{}.{}'.format( + self.__module__, type(self).__name__ + ) if get_logger is None: get_logger = logging.getLogger self.get_logger = get_logger @@ -173,22 +185,28 @@ class ParamikoTransport(Transport): """Transport that uses the Paramiko SSH2 library""" def __init__(self, host): super(ParamikoTransport, self).__init__(host) - sock = socket.create_connection((host.external_hostname, - host.ssh_port)) + sock = socket.create_connection( + (host.external_hostname, host.ssh_port) + ) self._transport = transport = paramiko.Transport(sock) transport.connect(hostkey=host.host_key) if host.ssh_key_filename: filename = os.path.expanduser(host.ssh_key_filename) key = paramiko.RSAKey.from_private_key_file(filename) self.log.debug( - 'Authenticating with private RSA key using user %s' % - host.ssh_username) + 'Authenticating with private RSA key using user %s', + host.ssh_username + ) transport.auth_publickey(username=host.ssh_username, key=key) elif host.ssh_password: - self.log.debug('Authenticating with password using user %s' % - host.ssh_username) - transport.auth_password(username=host.ssh_username, - password=host.ssh_password) + self.log.debug( + 'Authenticating with password using user %s', + host.ssh_username + ) + transport.auth_password( + username=host.ssh_username, + password=host.ssh_password + ) else: self.log.critical('No SSH credentials configured') raise RuntimeError('No SSH credentials configured') @@ -252,9 +270,12 @@ class ParamikoTransport(Transport): logger_name = self.get_next_command_logger_name() ssh = self._transport.open_channel('session') self.log.info('RUN %s', argv) - return SSHCommand(ssh, argv, logger_name=logger_name, - log_stdout=log_stdout, - get_logger=self.host.config.get_logger) + return SSHCommand( + ssh, argv, + logger_name=logger_name, + log_stdout=log_stdout, + get_logger=self.host.config.get_logger + ) def get_file(self, remotepath, localpath): self.log.debug('GET %s', remotepath) @@ -301,11 +322,13 @@ class OpenSSHTransport(Transport): control_file = os.path.join(self.control_dir.path, 'control') known_hosts_file = os.path.join(self.control_dir.path, 'known_hosts') - argv = ['ssh', - '-l', self.host.ssh_username, - '-o', 'ControlPath=%s' % control_file, - '-o', 'StrictHostKeyChecking=no', - '-o', 'UserKnownHostsFile=%s' % known_hosts_file] + argv = [ + 'ssh', + '-l', self.host.ssh_username, + '-o', 'ControlPath={}'.format(control_file), + '-o', 'StrictHostKeyChecking=no', + '-o', 'UserKnownHostsFile={}'.format(known_hosts_file) + ] if self.host.ssh_key_filename: key_filename = os.path.expanduser(self.host.ssh_key_filename) @@ -339,9 +362,12 @@ class OpenSSHTransport(Transport): argv = command logger_name = self.get_next_command_logger_name() ssh = SSHCallWrapper(self.ssh_argv + list(command)) - return SSHCommand(ssh, argv, logger_name, log_stdout=log_stdout, - collect_output=collect_output, - get_logger=self.host.config.get_logger) + return SSHCommand( + ssh, argv, logger_name, + log_stdout=log_stdout, + collect_output=collect_output, + get_logger=self.host.config.get_logger + ) def file_exists(self, path): self.log.info('STAT %s', path) @@ -364,15 +390,15 @@ class OpenSSHTransport(Transport): def get_file_contents(self, filename, encoding=None): self.log.info('GET %s', filename) - cmd = self._run(['cat', filename], log_stdout=False) + cmd = self._run(['base64', filename], log_stdout=False) cmd.wait(raiseonerr=False) if cmd.returncode == 0: - result = cmd.stdout_text + result = base64.b64decode(cmd.stdout_text) if encoding: result = result.decode(encoding) return result else: - raise IOError('File %r could not be read' % filename) + raise IOError('File {!r} could not be read'.format(filename)) def rmdir(self, path): self.log.info('RMDIR %s', path) @@ -384,15 +410,20 @@ class OpenSSHTransport(Transport): cmd = self._run(['rm', filepath]) cmd.wait() if cmd.returncode != 0: - raise IOError('File %r could not be deleted' % filepath) + raise IOError( + 'File {!r} could not be deleted'.foramat(filepath) + ) def rename_file(self, oldpath, newpath): self.log.info('RENAME %s TO %s', oldpath, newpath) cmd = self._run(['mv', oldpath, newpath]) cmd.wait() if cmd.returncode != 0: - raise IOError('File %r could not be renamed to %r ' - % (oldpath, newpath)) + raise IOError( + 'File {old!r} could not be renamed to {new!r} '.format( + old=oldpath, new=newpath + ) + ) class SSHCallWrapper(object): @@ -419,9 +450,6 @@ class SSHCallWrapper(object): assert mode == 'rb' return self.command.stderr - def shutdown_write(self): - self.command.stdin.close() - def recv_exit_status(self): return self.command.wait() @@ -433,9 +461,11 @@ class SSHCommand(Command): """Command implementation for ParamikoTransport and OpenSSHTranspport""" def __init__(self, ssh, argv, logger_name, log_stdout=True, collect_output=True, encoding='utf-8', get_logger=None): - super(SSHCommand, self).__init__(argv, logger_name, - log_stdout=log_stdout, - get_logger=get_logger) + super(SSHCommand, self).__init__( + argv, logger_name, + log_stdout=log_stdout, + get_logger=get_logger + ) self._stdout_lines = [] self._stderr_lines = [] self.running_threads = set() @@ -445,22 +475,26 @@ class SSHCommand(Command): self.log.debug('RUN %s', argv) self._ssh.invoke_shell() + def wrap_file(file, encoding): if encoding is None or sys.version_info < (3, 0): return file else: return io.TextIOWrapper(file, encoding=encoding) - stdin = self.stdin = wrap_file(self._ssh.makefile('wb'), 'utf-8') + + self.stdin = wrap_file(self._ssh.makefile('wb'), 'utf-8') stdout = wrap_file(self._ssh.makefile('rb'), encoding) stderr = wrap_file(self._ssh.makefile_stderr('rb'), encoding) if collect_output: - self._start_pipe_thread(self._stdout_lines, stdout, 'out', - log_stdout) + self._start_pipe_thread( + self._stdout_lines, stdout, + 'out', log_stdout + ) self._start_pipe_thread(self._stderr_lines, stderr, 'err', True) def _end_process(self): - self._ssh.shutdown_write() + self.stdin.close() while self.running_threads: self.running_threads.pop().join() @@ -491,7 +525,10 @@ class SSHCommand(Command): return thread -if not have_paramiko or os.environ.get('PYTESTMULTIHOST_SSH_TRANSPORT') == 'openssh': +if ( + not have_paramiko or + os.environ.get('PYTESTMULTIHOST_SSH_TRANSPORT') == 'openssh' +): SSHTransport = OpenSSHTransport else: SSHTransport = ParamikoTransport