#!/usr/bin/python ''' Usage: python vsphere8_upgrade_certificate_checks.py This script checks for the presence of SHA-1 cert in vSphere. If server is not localhost, user will be prompted to enter administrative user credentials. Credentials must always be specified on Windows, but this script must run on the target Windows machine. ''' import atexit import logging import getpass import os import re import six import ssl import sys import subprocess import OpenSSL from threading import Timer try: # Not available on SDDC Manager. PYTHONPATH must be used. sys.path.append(os.environ['VMWARE_PYTHON_PATH']) except KeyError: pass try: # Not available on CSWIN. Login will be required. sys.path.append('/usr/lib/vmware-updatemgr/python/hcl') except KeyError: pass from pyVim.connect import SmartConnect, Disconnect from pyVmomi import vim, vmodl, SoapStubAdapter, VmomiSupport __version__ = "1.8" PY2 = sys.version_info[0] == 2 PY3 = sys.version_info[0] == 3 server = '' port_number = 443 user = '' password = '' ALLOWED_SIGNATURE_ALGORITHMS = set(["sha512WithRSAEncryption", "sha384WithRSAEncryption", "sha256WithRSAEncryption", "sha224WithRSAEncryption", "ecdsa-with-SHA512", "ecdsa-with-SHA384", "ecdsa-with-SHA256", "ecdsa-with-SHA224", "dsa_with_SHA256", "dsa_with_SHA224"]) FETCH_CERTIFICATES_TIMEOUT_SECONDS = 120 PROP_NAME = "name" PROP_PARENT = "parent" PROP_CONFIG_PRODUCT_VERSION = "config.product.version" PROP_CONFIG_PRODUCT_BUILD = "config.product.build" PROP_RUNTIME_CONNECTIONSTATE = "runtime.connectionState" HOST_PROPS = [PROP_NAME, PROP_PARENT, PROP_CONFIG_PRODUCT_VERSION, PROP_CONFIG_PRODUCT_BUILD, PROP_RUNTIME_CONNECTIONSTATE] FAULT_KEY = "fault" def encode(obj, encodingFunction): ''' @param obj: Object on which to execute the provided encoding function, if it is a dict or list, executes it on the elements inside @param encodingFunction: Function to execute on the provided object @type function ''' if isinstance(obj, dict): result = {} for k, v in list(obj.items()): result[k] = encode(v, encodingFunction) return result elif isinstance(obj, list): return [encode(element, encodingFunction) for element in obj] else: return encodingFunction(obj) def toUnicode(noneUniCodeStr): '''Method to encode string to unicode. If it cannot achieve that it will return the same string @param noneUniCodeStr: String to encode to unicode @type str ''' if noneUniCodeStr is None or \ (PY2 and not isinstance(noneUniCodeStr, str)) or \ (PY3 and isinstance(noneUniCodeStr, str)): return noneUniCodeStr # Try with file system encoding, if it is not that, try with couple more # and after that give up try: return noneUniCodeStr.decode(sys.getfilesystemencoding()) except (UnicodeDecodeError, LookupError): pass # if it is windows its more likely to succeed with mbcs, however if # that fails we will try with utf-8 and then give up if os.name == 'nt': try: return noneUniCodeStr.decode(WIN_ENCODING) except (UnicodeDecodeError, LookupError): pass # valid ascii is valid utf-8 so no point of trying ascii try: return noneUniCodeStr.decode(DEF_ENCODING) except (UnicodeDecodeError, LookupError): pass # cannot decode, return same string and leave the system to fail logging.error('Tried to decode a string to unicode but it wasn\'t successful.' 'Expecting system failures') return noneUniCodeStr def to_file_system_encoding(obj): '''Method to encode unicode to file system encoding. If it cannot achieve that it will return the same string @param obj: Unicode object to encode using file system encoding @type unicode ''' # This is wrapping unicode in file system encoding so that Popen can read it if obj is not None and ((PY2 and isinstance(obj, unicode)) \ or (PY3 and isinstance(obj, str))): obj = obj.encode(sys.getfilesystemencoding()) return obj def run_command(cmd, stdin=None, timeout=FETCH_CERTIFICATES_TIMEOUT_SECONDS): ''' execute a command with the given input and return the return code and output ''' if PY3 and isinstance(stdin, str): # Need to be bytes the stdin stdin = stdin.encode(sys.getdefaultencoding()) cmd = encode(cmd, toUnicode) if PY2: cmd = encode(cmd, to_file_system_encoding) process = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, stdin=subprocess.PIPE) timer = Timer(timeout, process.kill) try: timer.start() stdout, stderr = process.communicate(stdin) rc = process.wait() finally: timer.cancel() if PY3: # Return type is bytes for Python 3+ thus need to cast to str stdout = toUnicode(stdout) stderr = toUnicode(stderr) if rc == 0: stdout = stdout.rstrip() if stderr is not None: stderr = stderr.rstrip() return rc, stdout, stderr def vecs_list_stores(): ''' get store names present in vecs ''' if os.name == 'nt': bin_path = os.path.join(os.environ['VMWARE_CIS_HOME'], 'vmafdd', 'vecs-cli.exe') else: bin_path = '/usr/lib/vmware-vmafd/bin/vecs-cli' if server != "localhost": command = [bin_path, 'store', 'list', '--server', server, '--upn', user] else: command = [bin_path, 'store', 'list'] (rc, stdout, stderr) = run_command(command, password) if rc != 0: msg = "Failed to list certificates from store" logging.error(stderr) raise Exception("%s\n Error: %s" % (msg, stderr)) #remove BACKUP_STORE from store list store_list = stdout.splitlines() for name in store_list: if "BACKUP_STORE" == name.upper() or "ENTER PASSWORD" in name.upper() or "SMS" == name.upper(): store_list.remove(name) return store_list def vecs_list_certificates(store_name): ''' get certificates in store ''' if os.name == 'nt': bin_path = os.path.join(os.environ['VMWARE_CIS_HOME'], 'vmafdd', 'vecs-cli.exe') else: bin_path = '/usr/lib/vmware-vmafd/bin/vecs-cli' if server != "localhost": command = [bin_path, 'entry', 'list', '--store', store_name, '--server', server, '--upn', user] else: command = [bin_path, 'entry', 'list', '--store', store_name] (rc, stdout, stderr) = run_command(command, password) if rc != 0: msg = "Failed to get certificate from store" logging.error(stderr) raise Exception("%s\n Error: %s" % (msg, stderr)) pem_format = re.compile(r'-----BEGIN CERTIFICATE-----[\s\S]+?-----END CERTIFICATE-----') certs = pem_format.findall(stdout) return certs def validate_vecs(): ''' check the existence of SHA-1 cert in vecs stores and add the weak certs to list ''' vecs_refresh() error_spec_list = [] vecs_stores = vecs_list_stores() for store_name in vecs_stores: logging.info("Verifing vCenter Server VECS store: %s ", store_name) list_of_certs = vecs_list_certificates(store_name) bad_subjects = validate_certificates(list_of_certs) for subject, thumbprint, subject_keyid, signature_algorithm in bad_subjects: error = ("The certificate with subject '%s' in VECS store %s has " "weak signature algorithm %s. " "The certificate thumbprint is %s." % (get_subject_string(subject), store_name, signature_algorithm, thumbprint)) if subject_keyid: error += (" The certificate Subject Key Identifier is %s." % subject_keyid) if store_name.upper() == "TRUSTED_ROOTS": error += (" Caution: Verify that any certificates signed by the " "problematic certificate are not in use by " "vCenter Server.") error_spec_list.append(error) return error_spec_list def get_credentials(): global server global port_number global user global password server = None if os.name != 'nt': server = six.moves.input('Enter hostname [Default: localhost]:') if not server or server.strip() == '': server = "localhost" if server != "localhost" or os.name == "nt": port = six.moves.input('Enter port number [Default: 443]:') if not port or port.strip() == '': port_number = 443 else: port_number = int(port) user = six.moves.input('Enter username [Default: Administrator@vsphere.local]:') if not user or user.strip() == '': user = "Administrator@vsphere.local" password = getpass.getpass(prompt='Enter Password:') if not password or password.strip() == '': raise Exception('Password cannot be empty') def set_credentials(server_in, port_number_in, user_in, password_in): global server global port_number global user global password server = server_in port_number = port_number_in user = user_in password = password_in def vecs_refresh(): ''' Function to sync certificates from vmdir to vecs ''' if os.name == 'nt': bin_path = os.path.join(os.environ['VMWARE_CIS_HOME'], 'vmafdd', 'vecs-cli.exe') else: bin_path = '/usr/lib/vmware-vmafd/bin/vecs-cli' if server != "localhost": command = [bin_path, 'force-refresh', '--server', server, '--upn', user] else: command = [bin_path, 'force-refresh'] (rc, stdout, stderr) = run_command(command, password) if rc != 0: msg = "Failed to refresh vecs store." logging.error(stderr) raise Exception("%s\n Error: %s" % (msg, stderr)) def get_si(cert_store): '''Get the vCenter Server ServiceInstance object. When running this script on vCenter Server we can retrieve the ServiceInstance object without requesting credentials from the user. ''' if server == "localhost" and os.name != "nt": from hardware_discovery.services.vc_service import getVCServiceFromCertificate vcService = getVCServiceFromCertificate(cert_store) vcService.login() vcStub = vcService.stub.soapStub hostname = vcStub.host.split(':')[0] vcLatestVimStub = SoapStubAdapter( host=hostname, port=vcService.port, version=VmomiSupport.newestVersions.GetName('vim')) vcLatestVimStub.cookie = vcStub.cookie si = vim.ServiceInstance("ServiceInstance", vcLatestVimStub) if si is None: raise Exception("Failed to connect to vCenter Server with the latest " "vmodl version {}".format(vcService.host)) else: if PY3: sslContext = ssl._create_unverified_context() else: sslContext = None si = SmartConnect(host=server, port=port_number, user=user, pwd=password, sslContext=sslContext) if not si: raise Exception("Could not connect to '%s'." % server) logging.debug("Connected to %s (%s)" % (server, si.content.about.fullName)) atexit.register(Disconnect, si) return si def get_all_hosts(si): '''Find all ESX hosts and cache important properties. ''' # Create a view for all hosts in the inventory. hostViewRef = si.content.viewManager.CreateContainerView( si.content.rootFolder, [vim.HostSystem], True) # Create a traversal spec of type ContainverView. traversalSpec = vim.TraversalSpec() traversalSpec.name = 'host-entity-traversal' traversalSpec.type = vim.ContainerView traversalSpec.path = 'view' traversalSpec.skip = False # Create a object spec with selectSet of the hosts in the view. objectSpec = vmodl.query.PropertyCollector.ObjectSpec() objectSpec.obj = hostViewRef objectSpec.skip = True objectSpec.selectSet = [traversalSpec] propertySpec = vmodl.query.PropertyCollector.PropertySpec() propertySpec.type = vim.HostSystem propertySpec.pathSet = HOST_PROPS filterSpec = vmodl.query.PropertyCollector.FilterSpec() filterSpec.objectSet = [objectSpec] filterSpec.propSet = [propertySpec] collector = si.content.propertyCollector retrieveRes = collector.RetrievePropertiesEx(specSet=[filterSpec], options=vmodl.query.PropertyCollector.RetrieveOptions()) logging.debug("Result of RetrievePropertiesEx: %s", retrieveRes) objects = [] if retrieveRes: objects = retrieveRes.objects token = retrieveRes.token while token: retrieveRes = collector.ContinueRetrievePropertiesEx(token) logging.debug( "Result of ContinueRetrievePropertiesEx: %s with token: %s", retrieveRes, token) objects.extend(retrieveRes.objects) token = retrieveRes.token logging.debug("Total number of hosts in inventory: %d", len(objects)) hostsAndPropsDict = {} for object in objects: # Create a dict for the host. host = object.obj hostsAndPropsDict[host] = {} if len(object.missingSet) > 0: hostsAndPropsDict[host][FAULT_KEY] = object.missingSet[0].fault continue for prop in object.propSet: hostsAndPropsDict[host][prop.name] = prop.val logging.debug( "Generated a map of all hosts and their properties in the inventory %s", hostsAndPropsDict) return hostsAndPropsDict def validate_certificates(certificates): '''Check a set certificates for unsupported signature algorithms. Returns a list of bad certificates indicated by a tuple of subject name, thumbprint (SHA1), subject key ID (if present) and signature algorithm. ''' bad_subjects = [] for cert in certificates: x509_cert = OpenSSL.crypto.load_certificate( OpenSSL.crypto.FILETYPE_PEM, cert) signature_algorithm = x509_cert.get_signature_algorithm().decode() if signature_algorithm not in ALLOWED_SIGNATURE_ALGORITHMS: subject_keyid = None for idx in range(x509_cert.get_extension_count()): ext = x509_cert.get_extension(idx) if ext.get_short_name() == b'subjectKeyIdentifier': subject_keyid = str(ext) bad_subjects.append((x509_cert.get_subject(), x509_cert.digest("sha1").decode(), subject_keyid, signature_algorithm)) return bad_subjects def get_subject_string(subject): '''Convert a subject to string. ''' return "".join("/{:s}={:s}".format(name.decode(), value.decode()) \ for name, value in subject.get_components()) class TimeOutError(Exception): '''A timeout exception. ''' pass def timeout_handler(): '''Throw a timeout exception. ''' raise TimeOutError("Operation timed out") def validate_esx_ca_certificates(host): '''Validate ESX CA certificates. The ESX CA trust store is accessible through CertificateManager. ''' errors = [] timer = Timer(FETCH_CERTIFICATES_TIMEOUT_SECONDS, timeout_handler) try: timer.start() certificates = host.configManager.certificateManager. \ ListCACertificates() finally: timer.cancel() bad_subjects = validate_certificates(certificates) for subject, thumbprint, subject_keyid, signature_algorithm in bad_subjects: error = ("Host %s has a configured certificate authority (CA) with " "subject name '%s' that has weak signature algorithm %s. " "The certificate thumbprint is %s. " % (host.name, get_subject_string(subject), signature_algorithm, thumbprint)) if subject_keyid: error += ("The certificate Subject Key Identifier is %s. " % subject_keyid) error += ("Cleanup vCenter Server TRUSTED_ROOTS before explicitly " "removing certificates from the host.") errors.append(error) return errors def validate_esx_machine_certificate(hostname): '''Validate the ESX machine SSL certificate. The ESX machine SSL certificate chain can be retrieved only by starting the hankshake. The most portable way of doing this is by using the openssl command line tool. ''' errors = [] if os.name == 'nt': bin_path = os.environ['VMWARE_OPENSSL_BIN'] # Running the openssl app on Windows blocks for a long period # period of time. Only wait 5 seconds because the certificates # should be printed by then. timeout = 5 else: bin_path = 'openssl' timeout = FETCH_CERTIFICATES_TIMEOUT_SECONDS cmd = [bin_path, 's_client', '-showcerts', '-connect', '%s:443' % hostname ] # Ignore the return code because we may have killed the subprocess # due to reaching the timeout. rc, stdout, stderr = run_command(cmd, timeout=timeout) if rc != 0: logging.debug("OpenSSL command failed with return code %d" % rc) certs = [] cert_pattern = r'-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----' for cert in re.findall(cert_pattern, stdout, re.DOTALL): certs.append(cert) if len(certs) == 0: errors.append( "Failed to connect to host %s. The host's TLS certificate " "cannot be validated." % hostname) return errors bad_subjects = validate_certificates(certs) for subject, thumbprint, subject_keyid, signature_algorithm in bad_subjects: error = ("Host %s has a configured TLS certificate with subject name " "'%s' that has weak signature algorithm %s. " "The certificate thumbprint is %s." % (hostname, get_subject_string(subject), signature_algorithm, thumbprint)) if subject_keyid: error += (" The certificate Subject Key Identifier is %s." % subject_keyid) errors.append(error) return errors def validate_esx_hosts(hosts): '''Validate the specified ESX hosts. ''' all_errors = [] num_hosts = len(hosts) num_hosts_scanned = 0 for host, props in hosts.items(): logging.info("Verifying ESXi host %d of %d: %s ", num_hosts_scanned + 1, num_hosts, host.name) try: if props.get(FAULT_KEY): all_errors.append( "Could not retrieve properties for host %s: %s" % (host, props.get(FAULT_KEY))) num_hosts_scanned += 1 continue name = props[PROP_NAME] connectionState = props.get(PROP_RUNTIME_CONNECTIONSTATE) version = props.get(PROP_CONFIG_PRODUCT_VERSION) if props.get( PROP_CONFIG_PRODUCT_VERSION) else "unknown" build = props.get(PROP_CONFIG_PRODUCT_BUILD) if props.get( PROP_CONFIG_PRODUCT_BUILD) else "unknown" logging.debug("Host %s, Build: %s, Version: %s", name, build, version) if connectionState != vim.HostSystem.ConnectionState.connected: all_errors.append( "Host %s is not connected and cannot be validated." % name) num_hosts_scanned += 1 continue logging.debug("Verifing host %s certificate store", name) errors = validate_esx_ca_certificates(host) all_errors.extend(errors) logging.debug("Verifing host %s machine certificate", name) errors = validate_esx_machine_certificate(name) all_errors.extend(errors) except vmodl.fault.ManagedObjectNotFound: logging.warning("Host %s not found in vCenter Server inventory.", host.name) except vmodl.MethodFault as e: msg = e.msg if e.faultMessage and len(e.faultMessage) != 0: for faultMessage in e.faultMessage: msg = msg + " " + faultMessage.message all_errors.append( "Caught exception while validating host %s: %s" % (name, msg)) except Exception as e: all_errors.append( "Caught exception while validating host %s: %s" % (name, str(e))) num_hosts_scanned += 1 return all_errors def is_psc_node(): '''Check if the system is a Platform Services Controller (PSC). The PCS doesn't have access to the VC host inventory. ''' try: DEPLOYMENT_TYPE_FILE_PATH = os.path.join(os.environ["VMWARE_CFG_DIR"], "deployment.node.type") except KeyError: # Assume not PSC. e.g. Might be SDDC Manager. pass else: try: with open(DEPLOYMENT_TYPE_FILE_PATH) as f: node_type = f.read().rstrip() return node_type == "infrastructure" except FileNotFoundError: # Assume not PSC. e.g. Might be SDDC Manager. pass return False def validate_all_esx_hosts(): '''Validate all ESX hosts. ''' if is_psc_node(): logging.info("No host inventory to check on PSC node.") return [] cert_store = None try: if server == "localhost" and os.name != "nt": from hardware_discovery.services.utils.authentication import CertificateStore cert_store = CertificateStore() si = get_si(cert_store) hosts = get_all_hosts(si) logging.debug("Host list: %s", hosts) return validate_esx_hosts(hosts) except vmodl.MethodFault as e: logging.exception("Caught VMODL fault: %s", str(e)) sys.exit(1) except Exception as e: logging.exception("Caught exception: %s", str(e)) sys.exit(1) finally: if cert_store: cert_store.cleanup() def validate(): '''Perform all validations and return a list of error strings. ''' return validate_vecs() + validate_all_esx_hosts() def main(): logLevel = logging.INFO # Uncomment for debug logging #logLevel = logging.DEBUG logging.basicConfig(format='%(asctime)s.%(msecs)dZ %(levelname)s %(message)s', level=logLevel, datefmt='%Y-%m-%d %H:%M:%S') try: get_credentials() errors = validate() if errors: logging.error("") logging.error("#################### Errors Found ####################") logging.error("") logging.error("Support for certificates with weak signature algorithms " "has been removed in vSphere 8.0. " "Weak signature algorithm certificates must be replaced before upgrade. " "Refer to the vSphere release notes and VMware KB 89424 for more details. " "Correct the following %d issues before proceeding with upgrade." % len(errors)) logging.error("") for index, error_text in enumerate(errors): logging.error("%d. %s" % (index + 1, error_text)) logging.error("") logging.error("######################################################") sys.exit(1) else: logging.info("Validation was successful.") except Exception as error_msg: logging.error(error_msg) sys.exit(1) if __name__ == "__main__": main()