diff --git a/gp-okta.py b/gp-okta.py index fced07f..f80b117 100755 --- a/gp-okta.py +++ b/gp-okta.py @@ -28,7 +28,7 @@ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. """ -from __future__ import print_function, unicode_literals +from __future__ import print_function import io, os, sys, time, re, json, base64, getpass, subprocess, shlex, signal, tempfile, traceback from lxml import etree @@ -78,8 +78,35 @@ try: except ImportError: pass -to_b = lambda v: v if isinstance(v, binary_type) else v.encode('utf-8') -to_u = lambda v: v if isinstance(v, text_type) else v.decode('utf-8') +def _type_err(v, target): + # type: (Any, text_type) -> TypeError + return TypeError('cannot convert {0} to {1}'.format(type(v), target)) + +def to_b(v): + # type: (Union[binary_type, text_type]) -> binary_type + if isinstance(v, binary_type): + return v + if isinstance(v, text_type): + return v.encode('utf-8') + raise _type_err(v, 'bytes') + +def to_u(v): + # type: (Union[text_type, binary_type]) -> text_type + if isinstance(v, text_type): + return v + if isinstance(v, binary_type): + return v.decode('utf-8') + raise _type_err(v, 'unicode text') + +def to_n(v): + # type: (Union[text_type, binary_type]) -> str + if isinstance(v, str): + return v + if isinstance(v, text_type): + return v.encode('utf-8') + if isinstance(v, binary_type): + return v.decode('utf-8') + raise _type_err(v, 'native text') quiet = False @@ -94,7 +121,7 @@ def warn(s): print(u'[WARN] {0}'.format(s)) def dbg(d, h, *xs): - # type: (Any, str, Union[str, List, Dict]) -> None + # type: (Any, str, Union[str, List[str], Dict[str, Any]]) -> None if quiet: return if not d: @@ -107,15 +134,15 @@ def dbg(d, h, *xs): print(u'[DEBUG] {0}: {1}'.format(h, x)) def dbg_form(conf, name, data): - # type: (Dict[str, str], str, Dict[str, str]) -> None - if not conf.get('debug'): + # type: (Conf, str, Dict[str, str]) -> None + if not conf.debug: return for k in data: - dbg(True, name, u'{0}: {1}'.format(k, data[k])) + dbg(True, name, '{0}: {1}'.format(k, data[k])) if k in [u'SAMLRequest', u'SAMLResponse']: try: - saml_raw = to_u(base64.b64decode(data[k])) - dbg(True, name, u'{0}.decoded: {1}'.format(k, saml_raw)) + saml_raw = to_n(base64.b64decode(data[k]).decode('ascii')) + dbg(True, name, '{0}.decoded: {1}'.format(k, saml_raw)) except Exception: pass @@ -126,7 +153,7 @@ def err(s): sys.exit(1) def parse_xml(xml): - # type: (str) -> etree.ElementTree + # type: (str) -> etree._Element try: rxml = bytes(bytearray(xml, encoding='utf-8')) parser = etree.XMLParser(ns_clean=True, recover=True) @@ -135,7 +162,7 @@ def parse_xml(xml): err('failed to parse xml: {0}'.format(e)) def parse_html(html): - # type: (str) -> etree.ElementTree + # type: (str) -> etree._Element try: parser = etree.HTMLParser() return etree.fromstring(html, parser) @@ -143,14 +170,15 @@ def parse_html(html): err('failed to parse html: {0}'.format(e)) def parse_rjson(r): - # type: (requests.Response) -> Dict + # type: (requests.Response) -> Dict[str, Any] try: - return r.json() + j = r.json() # type: Dict[str, Any] + return j except Exception as e: err('failed to parse json: {0}'.format(e)) def parse_form(html, current_url = None): - # type: (etree.ElementTree, Optional[str]) -> Tuple[str, Dict[str, str]] + # type: (etree._Element, Optional[str]) -> Tuple[str, Dict[str, str]] xform = html.find('.//form') url = xform.attrib.get('action', '').strip() if not url.startswith('http') and current_url: @@ -163,80 +191,121 @@ def parse_form(html, current_url = None): data[k] = v return url, data -def load_conf(cf): - # tpye: (str) -> Dict[str, str] - log('load conf') - conf = {} - keys = ['vpn_url', 'username', 'password', 'okta_url'] - if isinstance(cf, binary_type): - cf = cf.decode('utf-8') - line_nr = 0 - for rline in cf.split('\n'): - line_nr += 1 - line = rline.strip() - mx = re.match(r'^\s*([^=\s]+)\s*=\s*(.*?)\s*(?:#\s+.*)?\s*$', line) - if mx: - k, v = mx.group(1).lower(), mx.group(2) - if k.startswith('#'): - continue - for q in '"\'': - if re.match(r'^{0}.*{0}$'.format(q), v): - v = v[1:-1] - conf[k] = v - conf['{0}.line'.format(k)] = line_nr - for k, v in os.environ.items(): - k = k.lower() - if k.startswith('gp_'): - k = k[3:] - if len(k) == 0: - continue - conf[k] = v.strip() - if len(conf.get('username', '').strip()) == 0: - conf['username'] = input('username: ').strip() - if len(conf.get('password', '').strip()) == 0: - conf['password'] = getpass.getpass('password: ').strip() - if len(conf.get('openconnect_certs', '').strip()) == 0: - conf['openconnect_certs'] = tempfile.NamedTemporaryFile() - log('will collect openconnect_certs in temporary file: {0} and verify against them'.format(conf['openconnect_certs'].name)) - else: - conf['openconnect_certs'] = open(conf['openconnect_certs'], 'wb') - if len(conf.get('vpn_url_cert', '').strip()) != 0: - if not os.path.exists(conf.get('vpn_url_cert')): - err('configured vpn_url_cert file does not exist') - conf['openconnect_certs'].write(open(conf.get('vpn_url_cert'), 'rb').read()) # copy it - if len(conf.get('okta_url_cert', '').strip()) != 0: - if not os.path.exists(conf.get('okta_url_cert')): - err('configured okta_url_cert file does not exist') - conf['openconnect_certs'].write(open(conf.get('okta_url_cert'), 'rb').read()) # copy it - if len(conf.get('client_cert', '').strip()) != 0: - if not os.path.exists(conf.get('client_cert')): - err('configured client_cert file does not exist') - for k in keys: - if k not in conf: - err('missing configuration key: {0}'.format(k)) - else: - if len(conf[k].strip()) == 0: - err('empty configuration key: {0}'.format(k)) - conf['debug'] = conf.get('debug', '').lower() in ['1', 'true'] - conf['openconnect_certs'].flush() - return conf +class Conf(object): + def __init__(self): + # type: () -> None + self._store = {} # type: Dict[str, str] + self._lines = {} # type: Dict[str, int] + self.debug = False + self.vpn_url = '' # for reassignment + self.certs = '' # for type + + def __getattr__(self, name): + # type: (str) -> str + if name in self._store: + return self._store[name] + return '' + + def get_value(self, name): + # type: (str) -> str + return to_n(getattr(self, name)) + + def get_bool(self, name): + # type: (str) -> bool + v = self.get_value(name) + return v.lower() in ['1', 'true'] + + def add_cert(self, cert): + # type: (str) -> None + if not cert: + return + if not self.certs: + if 'openconnect_certs' in self._store: + self.certs_fh = io.open(self._store['openconnect_certs'], 'wb') + else: + self.certs_fh = tempfile.NamedTemporaryFile() + self.certs = self.certs_fh.name + self.certs_fh.write(cert) + + def get_cert(self, name, default_verify=True): + # type: (str, bool) -> Union[str, bool] + name = '{0}_cert'.format(name) + if name in self._store: + return self._store[name] + return default_verify + + def get_line(self, name): + # type: (str) -> int + if name in self._lines: + return self._lines[name] + return 0 + + @classmethod + def from_data(cls, content): + # type: (str) -> Conf + conf = cls() + log('load conf') + keys = ['vpn_url', 'username', 'password', 'okta_url'] + line_nr = 0 + for rline in to_n(content).split('\n'): + line_nr += 1 + line = rline.strip() + mx = re.match(r'^\s*([^=\s]+)\s*=\s*(.*?)\s*(?:#\s+.*)?\s*$', line) + if mx: + k, v = mx.group(1).lower(), mx.group(2) + if k.startswith('#'): + continue + for q in '"\'': + if re.match(r'^{0}.*{0}$'.format(q), v): + v = v[1:-1] + conf._store[k] = v + conf._lines[k] = line_nr + for k, v in os.environ.items(): + k = k.lower() + if k.startswith('gp_'): + k = k[3:] + if len(k) == 0: + continue + conf._store[k] = v.strip() + if len(conf._store.get('username', '').strip()) == 0: + conf._store['username'] = input('username: ').strip() + if len(conf._store.get('password', '').strip()) == 0: + conf._store['password'] = getpass.getpass('password: ').strip() + for cert_name in ['vpn_url_cert', 'okta_url_cert', 'client_cert']: + cert_file = conf._store.get(k, '').strip() + if cert_file: + cert_file = os.path.expandvars(os.path.expanduser(cert_file)) + if not os.path.exists(cert_file): + err('configured "{0}" file "{1}" does not exist'.format(cert_name, cert_file)) + with io.open(cert_file, 'rb') as fp: + conf.add_cert(fp.read()) + for k in keys: + if k not in conf._store: + err('missing configuration key: {0}'.format(k)) + else: + if len(conf._store[k].strip()) == 0: + err('empty configuration key: {0}'.format(k)) + if k == 'vpn_url': + setattr(conf, k, conf._store[k].strip()) + conf.debug = conf._store.get('debug', '').lower() in ['1', 'true'] + return conf def mfa_priority(conf, ftype, fprovider): - # type: (Dict[str, str], str, str) -> int + # type: (Conf, str, str) -> int if ftype == 'token:software:totp' or (ftype, fprovider) == ('token', 'symantec'): ftype = 'totp' if ftype not in ['totp', 'sms', 'push', 'webauthn']: return 0 - mfa_order = conf.get('mfa_order', '') + mfa_order = conf.mfa_order.split() if ftype in mfa_order: priority = (10 - mfa_order.index(ftype)) * 100 else: priority = 0 - value = conf.get('{0}.{1}'.format(ftype, fprovider)) + value = conf.get_value('{0}.{1}'.format(ftype, fprovider)) # type: Optional[str] if ftype in ('sms', 'webauthn'): if not (value or '').lower() in ['1', 'true']: value = None - line_nr = int(conf.get('{0}.{1}.line'.format(ftype, fprovider), 0)) + line_nr = conf.get_line('{0}.{1}'.format(ftype, fprovider)) if value is None: priority += 0 elif len(value) == 0: @@ -246,30 +315,30 @@ def mfa_priority(conf, ftype, fprovider): return priority def get_state_token(conf, c): - # tpye: (Dict[str, str], str) -> str + # type: (Conf, str) -> Optional[str] rx_state_token = re.search(r'var\s*stateToken\s*=\s*\'([^\']+)\'', c) if not rx_state_token: - dbg(conf.get('debug'), 'not found', 'stateToken') + dbg(conf.debug, 'not found', 'stateToken') return None - state_token = to_b(rx_state_token.group(1)).decode('unicode_escape').strip() + state_token = to_n(to_b(rx_state_token.group(1)).decode('unicode_escape').strip()) return state_token def get_redirect_url(conf, c, current_url = None): - # type: (Dict[str, str], str, Optional[str]) -> Optional[str] + # type: (Conf, str, Optional[str]) -> Optional[str] rx_base_url = re.search(r'var\s*baseUrl\s*=\s*\'([^\']+)\'', c) rx_from_uri = re.search(r'var\s*fromUri\s*=\s*\'([^\']+)\'', c) if not rx_from_uri: - dbg(conf.get('debug'), 'not found', 'formUri') + dbg(conf.debug, 'not found', 'formUri') return None - from_uri = to_b(rx_from_uri.group(1)).decode('unicode_escape').strip() + from_uri = to_n(to_b(rx_from_uri.group(1)).decode('unicode_escape').strip()) if from_uri.startswith('http'): return from_uri if not rx_base_url: - dbg(conf.get('debug'), 'not found', 'baseUri') + dbg(conf.debug, 'not found', 'baseUri') if current_url: return urljoin(current_url, from_uri) return from_uri - base_url = to_b(rx_base_url.group(1)).decode('unicode_escape').strip() + base_url = to_n(to_b(rx_base_url.group(1)).decode('unicode_escape').strip()) return base_url + from_uri def parse_url(url): @@ -278,8 +347,8 @@ def parse_url(url): return (purl[0], purl[1].split(':')[0]) def _send_req_pre(conf, name, url, data, expected_url=None): - # type: (Dict[str, str], str, str, Dict[str, Any], Optional[str]) -> None - dbg(conf.get('debug'), '{0}.request'.format(name), url) + # type: (Conf, str, str, Dict[str, Any], Optional[str]) -> None + dbg(conf.debug, '{0}.request'.format(name), url) dbg_form(conf, 'send.req.data', data) if expected_url: purl, pexp = parse_url(url), parse_url(expected_url) @@ -287,15 +356,15 @@ def _send_req_pre(conf, name, url, data, expected_url=None): err('{0}: unexpected url found {1} != {2}'.format(name, purl, pexp)) def _send_req_post(conf, r, name, can_fail=False): - # type: (Dict[str, str], requests.Response, str, bool) -> None + # type: (Conf, requests.Response, str, bool) -> None hdump = '\n'.join([k + ': ' + v for k, v in sorted(r.headers.items())]) rr = 'status: {0}\n\n{1}\n\n{2}'.format(r.status_code, hdump, r.text) if not can_fail and r.status_code != 200: err('{0}.request failed.\n{1}'.format(name, rr)) - dbg(conf.get('debug'), '{0}.response'.format(name), rr) + dbg(conf.debug, '{0}.response'.format(name), rr) def send_req(conf, s, name, url, data, get=False, expected_url=None, verify=True, can_fail=False): - # type: (Dict[str, str], requests.Session, str, str, Dict[str, Any], bool, str, Union[bool, str], bool) -> Tuple[int, requests.structures.CaseInsensitiveDict[str], str] + # type: (Conf, requests.Session, str, str, Dict[str, Any], bool, Optional[str], Union[bool, str], bool) -> Tuple[int, requests.structures.CaseInsensitiveDict[str], str] _send_req_pre(conf, name, url, data, expected_url) if get: r = s.get(url, verify=verify) @@ -305,7 +374,7 @@ def send_req(conf, s, name, url, data, get=False, expected_url=None, verify=True return r.status_code, r.headers, r.text def send_json_req(conf, s, name, url, data, get=False, expected_url=None, verify=True, can_fail=False): - # type: (Dict[str, str], requests.Session, str, str, Dict[str, Any], bool, str, Union[bool, str], bool) -> Tuple[int, requests.structures.CaseInsensitiveDict[str], Dict[str, Any]] + # type: (Conf, requests.Session, str, str, Dict[str, Any], bool, Optional[str], Union[bool, str], bool) -> Tuple[int, requests.structures.CaseInsensitiveDict[str], Dict[str, Any]] _send_req_pre(conf, name, url, data, expected_url) headers = {'Accept': 'application/json', 'Content-Type': 'application/json'} if get: @@ -316,19 +385,21 @@ def send_json_req(conf, s, name, url, data, get=False, expected_url=None, verify return r.status_code, r.headers, parse_rjson(r) def paloalto_prelogin(conf, s, gateway_url=None): - # type: (Dict[str, str], requests.Session, Optional[str]) -> str - verify = False + # type: (Conf, requests.Session, Optional[str]) -> etree._Element + verify = True # type: Union[str, bool] if gateway_url: # 2nd round or direct gateway: use gateway log('prelogin request [gateway_url]') url = '{0}/ssl-vpn/prelogin.esp'.format(gateway_url) - verify = conf.get('vpn_url_cert') + verify = conf.get_cert('vpn_url', True) else: # 1st round: use portal log('prelogin request [vpn_url]') - url = '{0}/global-protect/prelogin.esp'.format(conf.get('vpn_url')) - if conf.get('openconnect_certs') and os.path.getsize(conf.get('openconnect_certs').name) > 0: - verify = conf.get('openconnect_certs').name + url = '{0}/global-protect/prelogin.esp'.format(conf.vpn_url) + if conf.certs: + verify = conf.certs + else: + verify = True _, _h, c = send_req(conf, s, 'prelogin', url, {}, get=True, verify=verify) x = parse_xml(c) saml_req = x.find('.//saml-request') @@ -344,32 +415,32 @@ def paloalto_prelogin(conf, s, gateway_url=None): if len(saml_req.text.strip()) == 0: err('empty saml request') try: - saml_raw = to_u(base64.b64decode(saml_req.text)) + saml_raw = to_n(base64.b64decode(saml_req.text).decode('ascii')) except Exception as e: err('failed to decode saml request: {0}'.format(e)) - dbg(conf.get('debug'), 'prelogin.decoded', saml_raw) + dbg(conf.debug, 'prelogin.decoded', saml_raw) saml_xml = parse_html(saml_raw) return saml_xml def okta_saml(conf, s, saml_xml): - # type: (Dict[str, str], requests.Session, str) -> str + # type: (Conf, requests.Session, str) -> str log('okta saml request [okta_url]') url, data = parse_form(saml_xml) dbg_form(conf, 'okta.saml request', data) _, _h, c = send_req(conf, s, 'saml', url, data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) redirect_url = get_redirect_url(conf, c, url) if redirect_url is None: err('did not find redirect url') return redirect_url def okta_auth(conf, s, stateToken = None): - # type: (Dict[str, str], requests.Session, Optional[str]) -> Any + # type: (Conf, requests.Session, Optional[str]) -> Any log('okta auth request [okta_url]') - url = '{0}/api/v1/authn'.format(conf.get('okta_url')) + url = '{0}/api/v1/authn'.format(conf.okta_url) data = { - 'username': conf.get('username'), - 'password': conf.get('password'), + 'username': conf.username, + 'password': conf.password, 'options': { 'warnBeforePasswordExpired':True, 'multiOptionalFactorEnroll':True @@ -377,7 +448,7 @@ def okta_auth(conf, s, stateToken = None): } if stateToken is None else { 'stateToken': stateToken } - _, _h, j = send_json_req(conf, s, 'auth', url, data, verify=conf.get('okta_url_cert', True)) + _, _h, j = send_json_req(conf, s, 'auth', url, data, verify=conf.get_cert('okta_url', True)) while True: ok, r = okta_transaction_state(conf, s, j) @@ -386,10 +457,10 @@ def okta_auth(conf, s, stateToken = None): j = r def okta_transaction_state(conf, s, j): - # type: (Dict[str, str], requests.Session, Dict[str, Any]) -> Tuple[bool, Dict[str, Any]] + # type: (Conf, requests.Session, Dict[str, Any]) -> Tuple[bool, Dict[str, Any]] # https://developer.okta.com/docs/api/resources/authn#transaction-state status = j.get('status', '').strip().lower() - dbg(conf.get('debug'), 'status', status) + dbg(conf.debug, 'status', status) # status: unauthenticated # status: password_warn if status == 'password_warn': @@ -402,7 +473,7 @@ def okta_transaction_state(conf, s, j): err('empty state token') data = {'stateToken': state_token} _, _h, j = send_json_req(conf, s, 'skip', url, data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) return False, j # status: password_expired # status: recovery @@ -424,10 +495,9 @@ def okta_transaction_state(conf, s, j): return True, session_token print(j) err('unknown status: {0}'.format(status)) - return False, None def okta_mfa(conf, s, j): - # type: (Dict[str, str], requests.Session, Dict[str, Any]) -> Dict[str, Any] + # type: (Conf, requests.Session, Dict[str, Any]) -> Dict[str, Any] state_token = j.get('stateToken', '').strip() if len(state_token) == 0: err('empty state token') @@ -448,7 +518,7 @@ def okta_mfa(conf, s, j): 'provider': provider, 'priority': mfa_priority(conf, factor_type, provider), 'url': factor_url}) - dbg(conf.get('debug'), 'factors', *factors) + dbg(conf.debug, 'factors', *factors) if len(factors) == 0: err('no factors found') for f in sorted(factors, key=lambda x: x.get('priority', 0), reverse=True): @@ -468,9 +538,9 @@ def okta_mfa(conf, s, j): err('no factors processed') def okta_mfa_totp(conf, s, factor, state_token): - # type: (Dict[str, str], requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] + # type: (Conf, requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] provider = factor.get('provider', '') - secret = conf.get('totp.{0}'.format(provider), '') or '' + secret = conf.get_value('totp.{0}'.format(provider)) code = None if len(secret) == 0: code = input('{0} TOTP: '.format(provider)).strip() @@ -489,11 +559,11 @@ def okta_mfa_totp(conf, s, factor, state_token): } log('mfa {0} totp request: {1} [okta_url]'.format(provider, code)) _, _h, j = send_json_req(conf, s, 'totp mfa', factor.get('url', ''), data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) return j def okta_mfa_sms(conf, s, factor, state_token): - # type: (Dict[str, str], requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] + # type: (Conf, requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] provider = factor.get('provider', '') data = { 'factorId': factor.get('id'), @@ -501,18 +571,18 @@ def okta_mfa_sms(conf, s, factor, state_token): } log('mfa {0} sms request [okta_url]'.format(provider)) _, _h, j = send_json_req(conf, s, 'sms mfa (1)', factor.get('url', ''), data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) code = input('{0} SMS verification code: '.format(provider)).strip() if len(code) == 0: return None data['passCode'] = code log('mfa {0} sms request [okta_url]'.format(provider)) _, _h, j = send_json_req(conf, s, 'sms mfa (2)', factor.get('url', ''), data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) return j def okta_mfa_push(conf, s, factor, state_token): - # type: (Dict[str, str], requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] + # type: (Conf, requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] provider = factor.get('provider', '') data = { 'factorId': factor.get('id'), @@ -524,32 +594,31 @@ def okta_mfa_push(conf, s, factor, state_token): while status == 'MFA_CHALLENGE': _, _h, j = send_json_req(conf, s, 'push mfa ({0})'.format(counter), factor.get('url', ''), data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) status = j.get('status', '').strip() - dbg(conf.get('debug'), 'status', status) + dbg(conf.debug, 'status', status) if status == 'MFA_CHALLENGE': time.sleep(3.33) counter += 1 return j def okta_mfa_webauthn(conf, s, factor, state_token): - # type: (Dict[str, str], requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] + # type: (Conf, requests.Session, Dict[str, str], str) -> Optional[Dict[str, Any]] if not have_fido: err('Need fido2 package(s) for webauthn. Consider doing `pip install fido2` (or similar)') devices = list(CtapHidDevice.list_devices()) if not devices: err('webauthn configured, but no U2F devices found') - return None provider = factor.get('provider', '') log('mfa {0} challenge request [okta_url]'.format(provider)) data = { 'stateToken': state_token } _, _h, j = send_json_req(conf, s, 'webauthn mfa challenge', factor.get('url', ''), data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) rfactor = j['_embedded']['factor'] profile = rfactor['profile'] - purl = parse_url(conf.get('okta_url', '')) + purl = parse_url(conf.okta_url) origin = '{0}://{1}'.format(purl[0], purl[1]) challenge = rfactor['_embedded']['challenge']['challenge'] credentialId = websafe_decode(profile['credentialId']) @@ -559,7 +628,7 @@ def okta_mfa_webauthn(conf, s, factor, state_token): print('!!! Touch the flashing U2F device to authenticate... !!!') try: result = client.get_assertion(purl[1], challenge, allow_list) - dbg(conf.get('debug'), 'assertion.result', result) + dbg(conf.debug, 'assertion.result', result) break except Exception: traceback.print_exc(file=sys.stderr) @@ -569,38 +638,39 @@ def okta_mfa_webauthn(conf, s, factor, state_token): assertion, client_data = result[0][0], result[1] # only one cred in allowList, so only one response. data = { 'stateToken': state_token, - 'clientData': to_u(base64.b64encode(client_data)), - 'signatureData': to_u(base64.b64encode(assertion.signature)), - 'authenticatorData': to_u(base64.b64encode(assertion.auth_data)) + 'clientData': to_n((base64.b64encode(client_data)).decode('ascii')), + 'signatureData': to_n((base64.b64encode(assertion.signature)).decode('ascii')), + 'authenticatorData': to_n((base64.b64encode(assertion.auth_data)).decode('ascii')) } log('mfa {0} signature request [okta_url]'.format(provider)) _, _h, j = send_json_req(conf, s, 'uf2 mfa signature', j['_links']['next']['href'], data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) return j def okta_redirect(conf, s, session_token, redirect_url): - # type: (Dict[str, str], requests.Session, str, str) -> Tuple[str, str] + # type: (Conf, requests.Session, str, str) -> Tuple[str, str] rc = 0 form_url = None # type: Optional[str] form_data = {} # type: Dict[str, str] + rurl = redirect_url # type: Optional[str] while True: if rc > 10: err('redirect rabbit hole is too deep...') rc += 1 - if redirect_url: + if rurl: data = { 'checkAccountSetupComplete': 'true', 'report': 'true', 'token': session_token, - 'redirectUrl': redirect_url + 'redirectUrl': rurl } - url = '{0}/login/sessionCookieRedirect'.format(conf.get('okta_url')) + url = '{0}/login/sessionCookieRedirect'.format(conf.okta_url) log('okta redirect request {0} [okta_url]'.format(rc)) _, h, c = send_req(conf, s, 'redirect', url, data, - verify=conf.get('okta_url_cert', True)) + verify=conf.get_cert('okta_url', True)) state_token = get_state_token(conf, c) - redirect_url = get_redirect_url(conf, c, url) - if redirect_url: + rurl = get_redirect_url(conf, c, url) + if rurl: form_url, form_data = None, {} else: xhtml = parse_html(c) @@ -611,26 +681,26 @@ def okta_redirect(conf, s, session_token, redirect_url): okta_auth(conf, s, state_token) elif form_url: log('okta redirect form request [vpn_url]') - purl, pexp = parse_url(form_url), parse_url(conf.get('vpn_url', '')) + purl, pexp = parse_url(form_url), parse_url(conf.vpn_url) if purl != pexp: # NOTE: redirect to nearest (geo) portal/gateway without any prior knowledge warn('{0}: unexpected url found {1} != {2}'.format('redirect form', purl, pexp)) _, h, c = send_req(conf, s, 'redirect form', form_url, form_data) else: _, h, c = send_req(conf, s, 'redirect form', form_url, form_data, - expected_url=conf.get('vpn_url'), verify=conf.get('vpn_url_cert', True)) + expected_url=conf.vpn_url, verify=conf.get_cert('vpn_url', True)) saml_username = h.get('saml-username', '').strip() prelogin_cookie = h.get('prelogin-cookie', '').strip() if saml_username and prelogin_cookie: saml_auth_status = h.get('saml-auth-status', '').strip() saml_slo = h.get('saml-slo', '').strip() - dbg(conf.get('debug'), 'saml prop', [saml_auth_status, saml_slo]) + dbg(conf.debug, 'saml prop', [saml_auth_status, saml_slo]) return saml_username, prelogin_cookie def paloalto_getconfig(conf, s, username = None, prelogin_cookie = None, can_fail = False): - # type: (Dict[str, str], requests.Session, Optional[str], Optional[str], bool) -> Tuple[int, str, Dict[str, str]] + # type: (Conf, requests.Session, Optional[str], Optional[str], bool) -> Tuple[int, str, Dict[str, str]] log('getconfig request [vpn_url]') - url = '{0}/global-protect/getconfig.esp'.format(conf.get('vpn_url')) + url = '{0}/global-protect/getconfig.esp'.format(conf.vpn_url) data = { #'jnlpReady': 'jnlpReady', #'ok': 'Login', @@ -643,15 +713,15 @@ def paloalto_getconfig(conf, s, username = None, prelogin_cookie = None, can_fai 'computer': 'DESKTOP', #'preferred-ip': '', 'inputStr': '', - 'user': username or conf['username'], - 'passwd': '' if prelogin_cookie else conf['password'], + 'user': username or conf.username, + 'passwd': '' if prelogin_cookie else conf.password, 'clientgpversion': '4.1.0.98', # 'host-id': '00:11:22:33:44:55' 'prelogin-cookie': prelogin_cookie or '', 'ipv6-support': 'yes' } sc, _h, c = send_req(conf, s, 'getconfig', url, data, - verify=conf.get('vpn_url_cert', True), can_fail=can_fail) + verify=conf.get_cert('vpn_url', True), can_fail=can_fail) if sc != 200: return sc, '', {} x = parse_xml(c) @@ -672,25 +742,24 @@ def paloalto_getconfig(conf, s, username = None, prelogin_cookie = None, can_fai if xtmp is not None: for entry in xtmp: cert = entry.find('.//cert').text - conf['openconnect_certs'].write(to_b(cert)) - conf['openconnect_certs'].flush() + conf.add_cert(cert) return 200, portal_userauthcookie, gateways # Combined first half of okta_saml with second half of okta_redirect def okta_saml_2(conf, s, gateway_url, saml_xml): - # type: (Dict[str, str], requests.Session, str, str) -> Tuple[str, str] + # type: (Conf, requests.Session, str, str) -> Tuple[str, str] log('okta saml request (2) [okta_url]') url, data = parse_form(saml_xml) dbg_form(conf, 'okta.saml request(2)', data) _, h, c = send_req(conf, s, 'okta saml request (2)', url, data, - expected_url=conf.get('okta_url'), verify=conf.get('okta_url_cert', True)) + expected_url=conf.okta_url, verify=conf.get_cert('okta_url', True)) xhtml = parse_html(c) url, data = parse_form(xhtml) dbg_form(conf, 'okta.saml request(2)', data) log('okta redirect form request (2) [gateway]') - verify = False - if conf.get('openconnect_certs') and os.path.getsize(conf.get('openconnect_certs').name) > 0: - verify = conf.get('openconnect_certs').name + verify = False # type: Union[str, bool] + if conf.certs: + verify = conf.certs _, h, c = send_req(conf, s, 'okta redirect form (2)', url, data, expected_url=gateway_url, verify=verify) saml_username = h.get('saml-username', '').strip() @@ -708,13 +777,12 @@ def output_gateways(gateways): print("\t{0} {1}".format(k, gateways[k])) def choose_gateway_url(conf, gateways): - # type: (Dict[str, str], Dict[str, str]) -> str - gateway_url = conf.get('gateway_url', '').strip() - if gateway_url: - return gateway_url + # type: (Conf, Dict[str, str]) -> str + if conf.gateway_url: + return conf.gateway_url if len(gateways) == 0: err('no available gateways') - gateway_name = conf.get('gateway', '').strip() + gateway_name = conf.gateway gateway_host = None for k in gateways.keys(): if gateways[k] == gateway_name: @@ -769,13 +837,13 @@ def main(): config_contents = decrypted_contents.data - conf = load_conf(config_contents) + conf = Conf.from_data(config_contents) s = requests.Session() s.headers['User-Agent'] = 'PAN GlobalProtect' - if conf.get('client_cert'): - s.cert = conf.get('client_cert') + if conf.client_cert: + s.cert = conf.client_cert if args.list_gateways: log('listing gateways') @@ -785,15 +853,15 @@ def main(): return 0 log('gateway list requires authentication') - another_dance = conf.get('another_dance', '').lower() in ['1', 'true'] - gateway_url = conf.get('gateway_url', '').strip() - gateway_name = conf.get('gateway', '').strip() + another_dance = conf.another_dance.lower() in ['1', 'true'] + gateway_url = conf.gateway_url + gateway_name = conf.gateway if gateway_url and not another_dance: vpn_url = gateway_url - if vpn_url != conf.get('vpn_url'): + if vpn_url != conf.vpn_url: log('Discarding \'vpn_url\', as concrete \'gateway_url\' is given and another_dance = 0') - conf['vpn_url'] = vpn_url + conf.vpn_url = vpn_url userauthcookie = None @@ -835,25 +903,25 @@ def main(): cookie = prelogin_cookie else: cookie_type = 'portal:portal-userauthcookie' - cookie = userauthcookie + cookie = userauthcookie or '' username = saml_username - cmd = conf.get('openconnect_cmd', 'openconnect') + cmd = conf.openconnect_cmd or 'openconnect' cmd += ' --protocol=gp -u \'{0}\'' cmd += ' --usergroup {1}' - if conf.get('client_cert'): - cmd += ' --certificate=\'{0}\''.format(conf.get('client_cert')) - if conf.get('openconnect_certs') and os.path.getsize(conf.get('openconnect_certs').name) > 0: - cmd += ' --cafile=\'{0}\''.format(conf.get('openconnect_certs').name) - cmd += ' --passwd-on-stdin ' + conf.get('openconnect_args', '') + ' \'{2}\'' + if conf.client_cert: + cmd += ' --certificate=\'{0}\''.format(conf.client_cert) + if conf.certs: + cmd += ' --cafile=\'{0}\''.format(conf.certs) + cmd += ' --passwd-on-stdin ' + conf.openconnect_args + ' \'{2}\'' cmd = cmd.format(username, cookie_type, - gateway_url if conf.get('another_dance', '').lower() in ['1', 'true'] else conf.get('vpn_url')) + gateway_url if conf.get_bool('another_dance') else conf.vpn_url) - pfmt = conf.get('openconnect_fmt') - if pfmt is None: + pfmt = conf.openconnect_fmt + if not pfmt: pfmt = '' - openconnect_bin = conf.get('openconnect_cmd', 'openconnect').split()[-1:][0] + openconnect_bin = (conf.openconnect_cmd or 'openconnect').split()[-1:][0] openconnect_bin = os.path.expandvars(os.path.expanduser(openconnect_bin)) with open(os.devnull, 'wb') as fnull: p = subprocess.Popen(['command', '-v', openconnect_bin], stdin=subprocess.PIPE, stdout=fnull, stderr=subprocess.STDOUT) @@ -870,18 +938,18 @@ def main(): pfmt = pfmt.replace('', cookie + '\\n') pfmt = pfmt.replace('', gateway_name + '\\n' if len(gateway_name) > 0 else '') for k in ['username', 'password', 'gateway_url']: - pfmt = pfmt.replace('<{0}>'.format(k), conf.get(k) + '\\n' if len(conf.get(k, '').strip()) > 0 else '') + v = conf.get_value(k).strip() + pfmt = pfmt.replace('<{0}>'.format(k), v + '\\n' if len(v) > 0 else '') pfmt = pfmt.replace('', saml_username + '\\n') if rmnl and pfmt.endswith('\\n'): pfmt = pfmt[:-2] pcmd = 'printf \'{0}\''.format(pfmt) print() - if conf.get('execute', '').lower() in ['1', 'true']: - cmd = shlex.split(cmd) - cmd = [os.path.expandvars(os.path.expanduser(x)) for x in cmd] + if conf.get_bool('execute'): + ecmd = [os.path.expandvars(os.path.expanduser(x)) for x in shlex.split(cmd)] pp = subprocess.Popen(shlex.split(pcmd), stdout=subprocess.PIPE) - cp = subprocess.Popen(cmd, stdin=pp.stdout, stdout=sys.stdout) + cp = subprocess.Popen(ecmd, stdin=pp.stdout, stdout=sys.stdout) if pp.stdout is not None: pp.stdout.close() # Do not abort on SIGINT. openconnect will perform proper exit & cleanup diff --git a/test/mypy.sh b/test/mypy.sh new file mode 100755 index 0000000..07c50e7 --- /dev/null +++ b/test/mypy.sh @@ -0,0 +1,11 @@ +#!/bin/sh +_cdir=$(cd -- "$(dirname "$0")" && pwd) +command -v mypy > /dev/null 2>&1 +if [ $? -ne 0 ]; then + echo "err: mypy (Optional Static Typing for Python) not found." + exit 1 +fi +for _pyver in 2.7 3.7; do + echo "--- Python ${_pyver}" + mypy --no-error-summary --strict --warn-unreachable --python-version "${_pyver}" "${_cdir}/../gp-okta.py" +done