from dataclasses import replace from datetime import datetime, timezone import json from pathlib import Path import sys import time import unittest from urllib.parse import parse_qs, urlsplit import jwt from cryptography.hazmat.primitives.asymmetric import rsa sys.path.insert(0, str(Path(__file__).resolve().parents[1] / 'scripts')) from alert_identity import Login, SCOPES from alert_policy import Policy, request_digest from alert_ack import Actor class Issuer: def __init__(self): self.key = rsa.generate_private_key(public_exponent=65537, key_size=2048) self.jwk = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(self.key.public_key())) self.jwk.update(kid='fixture', use='sig', alg='RS256') self.nonce = '' self.change = lambda c: c def request(self, method, url, body=None, headers=None): if url.endswith('/.well-known/openid-configuration'): return 200, dict(issuer='https://issuer.example', authorization_endpoint='https://issuer.example/authorize', token_endpoint='https://issuer.example/token', jwks_uri='https://issuer.example/jwks', code_challenge_methods_supported=['S256']) if url.endswith('/jwks'): return 200, {'keys': [self.jwk]} now = int(time.time()) claims = dict(iss='https://issuer.example', sub='fixture-human', aud='railiance-telemetry-admin', iat=now, exp=now + 600, nonce=self.nonce, tenant='tenant:platform', tenant_source='registration', principal_type='human', groups=[], roles=['railiance-admin'], assurance=dict(level='aal2', mfa=True, source='key-cape', methods=['pwd', 'otp'], at=now)) claims = self.change(claims) identity = jwt.encode(claims, self.key, algorithm='RS256', headers={'kid': 'fixture'}) access = dict(claims, aud='railiance-telemetry', scope=' '.join(SCOPES)) return 200, dict(id_token=identity, access_token=jwt.encode(access, self.key, algorithm='RS256', headers={'kid': 'fixture'})) class IdentityTests(unittest.TestCase): def setUp(self): self.issuer = Issuer() self.login = Login('https://issuer.example', 'https://telemetry.example', transport=self.issuer) def start(self): url, browser = self.login.start('/ack/alerts?fingerprint=0123456789abcdef&starts_at=2026-09-28T00:00:00Z') params = parse_qs(urlsplit(url).query) self.issuer.nonce = params['nonce'][0] self.assertEqual(params['code_challenge_method'], ['S256']) return params['state'][0], browser def actor(self): state, browser = self.start() sid, target = self.login.finish(state, browser, 'fixture-code') return self.login.session(sid) def test_verified_human_session_and_logout(self): state, browser = self.start() sid, target = self.login.finish(state, browser, 'fixture-code') actor = self.login.session(sid) self.assertEqual(actor.subject, 'fixture-human') self.assertEqual(actor.roles, ('railiance-admin',)) self.assertLessEqual(actor.expires_at, time.time() + 600) self.login.logout(sid) self.assertIsNone(self.login.session(sid)) with self.assertRaises(ValueError): self.login.finish(state, browser, 'fixture-code') def test_wrong_browser_and_state_replay_refused(self): state, browser = self.start() with self.assertRaises(ValueError): self.login.finish(state, 'other', 'fixture-code') with self.assertRaises(ValueError): self.login.finish(state, browser, 'fixture-code') def test_nonce_wrong_issuer_wrong_tenant_and_weak_mfa(self): for changes in ({'nonce': 'wrong'}, {'iss': 'https://other.example'}, {'tenant': 'tenant:other'}, {'principal_type': 'service'}, {'assurance': {'mfa': False}}, {'exp': 1}): self.issuer.change = lambda claims, changes=changes: dict(claims, **changes) state, browser = self.start() with self.assertRaises(ValueError): self.login.finish(state, browser, 'fixture-code') def test_return_path_cannot_escape_surface(self): for value in ('https://attacker.example/ack/alerts', '//attacker.example/ack/alerts', '/other'): with self.assertRaises(ValueError): self.login.start(value) class DecisionTransport: def __init__(self): self.change = lambda response: response def request(self, method, url, body=None, headers=None): request = json.loads(body) now = time.time() stamp = lambda delta: datetime.fromtimestamp(now + delta, timezone.utc).isoformat() response = dict(id='fixture-decision', request_id=request['id'], contract_version='flex-auth.decision-record.v1', effect='allow', obligations=[], matched_policy_version='v1', subject=request['subject'], resource=request['resource'], binding=dict(tenant=request['tenant'], action=request['action'], subject=request['subject'], resource=request['resource'], context=request.get('context', {}), submitted_request_digest=request_digest(request), request_digest='sha256:' + 'b' * 64), provenance=dict(policy_package='telemetry.ack', policy_version='v1', policy_package_digest='sha256:' + 'a' * 64, registry_snapshot_digest='sha256:' + 'c' * 64, evaluator='flex-auth/fixture', decision_time=stamp(0)), lifetime=dict(kind='ttl', not_before=stamp(-1), expires_at=stamp(30))) return 200, self.change(response) class PolicyTests(unittest.TestCase): def test_policy_pins_and_action_bindings_fail_closed(self): actor = Actor('https://issuer.example', 'fixture-human', 'tenant:platform', ('railiance-admin',), time.time() + 600, 'c' * 32, tenant_source='directory') transport = DecisionTransport() policy = Policy('https://policy.example', lambda: 'fixture-workload-token', 'telemetry.ack', 'v1', 'sha256:' + 'a' * 64, transport=transport) resource = 'alert:11111111-1111-1111-1111-111111111111' self.assertEqual(policy(actor, 'acknowledge', resource)['effect'], 'allow') mutations = [lambda r: dict(r, effect='deny'), lambda r: dict(r, obligations=['unknown']), lambda r: dict(r, request_id='other'), lambda r: dict(r, binding=dict(r['binding'], action='read')), lambda r: dict(r, binding=dict(r['binding'], submitted_request_digest='sha256:' + '0' * 64)), lambda r: dict(r, provenance=dict(r['provenance'], policy_package_digest='sha256:' + '0' * 64)), lambda r: dict(r, lifetime=dict(r['lifetime'], expires_at='2000-01-01T00:00:00Z'))] for change in mutations: transport.change = change self.assertEqual(policy(actor, 'acknowledge', resource)['effect'], 'deny') self.assertEqual(policy(actor, 'delete', resource)['effect'], 'deny') if __name__ == '__main__': unittest.main()