diff --git a/sshspray.py b/sshspray.py index 76775fb..d4466b6 100644 --- a/sshspray.py +++ b/sshspray.py @@ -10,6 +10,7 @@ from sys import exit from re import compile from os import devnull +from paramiko.ssh_exception import NoValidConnectionsError from os.path import isfile @@ -105,7 +106,7 @@ def do_key_checks(self): class Sprayer: def __init__(self, queue, username, key_file, password, passphrase, - host_key_file, port, timeout, target_list, verbosity): + host_key_file, port, timeout, retry, target_list, verbosity): self.queue = queue self.username = username self.key_file = key_file @@ -114,6 +115,7 @@ def __init__(self, queue, username, key_file, password, passphrase, self.host_key_file = host_key_file self.port = port self.timeout = timeout + self.retry = retry self.verbose = verbosity self.target_list = target_list @@ -138,17 +140,27 @@ def try_auth(self, ip): ssh = SSHClient() ssh.set_missing_host_key_policy(AutoAddPolicy()) conn_params = self.conn_params(ip) - try: - ssh.connect(**conn_params) - ssh.close() - - print(Message.success(), ip) - except Exception as e: - if not self.verbose: - pass - else: - output = [Message.fail(), ip, e] - print(*output[0:self.verbose+1]) + retry = self.retry + while retry > 0: + try: + ssh.connect(**conn_params) + ssh.close() + + print(Message.success(), ip) + except NoValidConnectionsError as e: + retry -= 1 + if retry > 0: + continue + else: + print(Message.info(), f'Failed to connect to {ip}') + return + except Exception as e: + if not self.verbose: + pass + else: + output = [Message.fail(), ip, e] + print(*output[0:self.verbose+1]) + return def run(self) -> None: print(Message.info(), "Running against {hostcount} hosts...".format(hostcount=len(self.target_list))) @@ -180,6 +192,7 @@ def arg_parse(): parser.add_argument("-v", "--verbose", action='count', default=0, help="Show failures. Use '-vv-' to show reasons for failure") parser.add_argument("-t", "--target-list", required=True, help="List of hosts to test(hostname, ip, and/or CIDR)") parser.add_argument("-w", "--wait", nargs='?', default=1, type=int, help="Timeout for each connection in seconds") + parser.add_argument("-r", "--retry", default=3, type=int, help="Times to retry in case of network failure") return parser.parse_args() @@ -241,7 +254,8 @@ def main(): passphrase=passphrase, timeout=args.wait, verbosity=verbosity, - target_list=targets) + target_list=targets, + retry=args.retry) sprayer.run()