railiance-telemetry/scripts/alert_identity.py

137 lines
7.3 KiB
Python
Raw Normal View History

"""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)