"""KeyCape OIDC code/PKCE sessions for the telemetry acknowledgment surface.""" import base64 import hashlib import secrets import threading import time from urllib.parse import urlencode, urlsplit import jwt from alert_ack import Actor, TENANT from telemetry_http import Transport, origin SCOPES = ('openid', 'profile', 'email', 'telemetry:read', 'telemetry:acknowledge') class Login: def __init__(self, issuer, public_origin, *, transport=None, clock=time.time): self.issuer = origin(issuer) self.public_origin = origin(public_origin) self.callback = public_origin + '/ack/auth/callback' self.client = 'railiance-telemetry-admin' self.transport = transport or Transport() self.clock = clock self.lock = threading.RLock() self.pending, self.sessions = {}, {} def prune(self): now = self.clock() self.pending = {k: v for k, v in self.pending.items() if v['expires'] > now} self.sessions = {k: v for k, v in self.sessions.items() if v.expires_at > now} def metadata(self): status, data = self.transport.request('GET', self.issuer + '/.well-known/openid-configuration') if status != 200 or data.get('issuer') != self.issuer: raise ValueError('issuer unavailable') for key in ('authorization_endpoint', 'token_endpoint', 'jwks_uri'): p = urlsplit(data[key]) if p.scheme + '://' + p.netloc != self.issuer or p.query or p.fragment or not p.path: raise ValueError('issuer endpoint refused') if 'S256' not in data.get('code_challenge_methods_supported', []): raise ValueError('PKCE unavailable') return data def start(self, return_to): parsed = urlsplit(return_to) if parsed.scheme or parsed.netloc or parsed.fragment or parsed.path != '/ack/alerts' or len(return_to) > 1024: raise ValueError('invalid return path') metadata = self.metadata() state, browser, nonce, verifier = (secrets.token_urlsafe(32) for _ in range(4)) with self.lock: self.prune() if len(self.pending) >= 1024: raise ValueError('login capacity') self.pending[state] = dict(browser=browser, nonce=nonce, verifier=verifier, expires=self.clock() + 300, return_to=return_to, metadata=metadata) challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b'=').decode() return metadata['authorization_endpoint'] + '?' + urlencode(dict(response_type='code', client_id=self.client, redirect_uri=self.callback, scope=' '.join(SCOPES), state=state, nonce=nonce, code_challenge=challenge, code_challenge_method='S256')), browser def decode(self, token, audience, keys): if not isinstance(token, str) or not 1 <= len(token) <= 32768: raise ValueError('invalid token') header = jwt.get_unverified_header(token) if header.get('alg') != 'RS256' or not header.get('kid'): raise ValueError('invalid signing key') matches = [key for key in keys if key.get('kid') == header['kid'] and key.get('kty') == 'RSA' and key.get('use', 'sig') == 'sig' and key.get('alg', 'RS256') == 'RS256'] if len(matches) != 1: raise ValueError('invalid signing key') claims = jwt.decode(token, jwt.PyJWK.from_dict(matches[0]).key, algorithms=['RS256'], issuer=self.issuer, audience=audience, options={ 'require': ['iss', 'aud', 'sub', 'iat', 'exp'], 'strict_aud': True}) if (not isinstance(claims['sub'], str) or not claims['sub'] or type(claims['exp']) is not int or type(claims['iat']) is not int or claims['exp'] <= self.clock() or claims['exp'] <= claims['iat']): raise ValueError('invalid claims') return claims def finish(self, state, browser, code): with self.lock: self.prune() pending = self.pending.pop(state, None) if not pending or not browser or not secrets.compare_digest(browser, pending['browser']) or not 1 <= len(code) <= 4096: raise ValueError('invalid login state') try: status, tokens = self.transport.request('POST', pending['metadata']['token_endpoint'], urlencode(dict(grant_type='authorization_code', client_id=self.client, redirect_uri=self.callback, code=code, code_verifier=pending['verifier'])).encode(), {'Content-Type': 'application/x-www-form-urlencoded'}) if status != 200: raise ValueError('code exchange failed') status, jwks = self.transport.request('GET', pending['metadata']['jwks_uri']) if status != 200 or not isinstance(jwks.get('keys'), list): raise ValueError('issuer keys unavailable') identity = self.decode(tokens['id_token'], self.client, jwks['keys']) access = self.decode(tokens['access_token'], 'railiance-telemetry', jwks['keys']) if identity.get('nonce') != pending['nonce']: raise ValueError('nonce mismatch') for key in ('sub', 'tenant', 'tenant_source', 'principal_type', 'roles', 'groups', 'assurance'): if key not in identity or identity[key] != access.get(key): raise ValueError('paired identity mismatch') if (access['tenant'] != TENANT or access['principal_type'] != 'human' or access['tenant_source'] not in ('directory', 'registration') or set(access.get('scope', '').split()) != set(SCOPES)): raise ValueError('human platform scope required') for key in ('roles', 'groups'): if not isinstance(access[key], list) or any(not isinstance(v, str) or not v for v in access[key]): raise ValueError('invalid identity claims') assurance = access['assurance'] if (assurance.get('level') != 'aal2' or assurance.get('mfa') is not True or assurance.get('source') != 'key-cape' or assurance.get('methods') != ['pwd', 'otp'] or type(assurance.get('at')) is not int or not -30 <= self.clock() - assurance['at'] <= 900): raise ValueError('fresh MFA required') actor = Actor(self.issuer, access['sub'], TENANT, tuple(access['roles']), min(access['exp'], identity['exp'], self.clock() + 900), secrets.token_urlsafe(32), assurance=dict(assurance), groups=tuple(access['groups']), tenant_source=access['tenant_source']) with self.lock: self.prune() if len(self.sessions) >= 1024: raise ValueError('session capacity') sid = secrets.token_urlsafe(32) self.sessions[sid] = actor return sid, pending['return_to'] except (jwt.PyJWTError, KeyError, TypeError, AttributeError): raise ValueError('invalid issuer response') from None def session(self, sid): with self.lock: self.prune() return self.sessions.get(sid) def logout(self, sid): with self.lock: self.sessions.pop(sid, None)