Assistant: codex Assistant-Model: gpt-6-astra Assistant-Session: 01a0e6f1-443f-7783-9920-a16b2ffc467f
136 lines
7.3 KiB
Python
136 lines
7.3 KiB
Python
"""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)
|