#!/usr/bin/env python3
"""CS09 deterministic teaching checks. All keys are public test fixtures, NEVER use live.
Requires Python 3.10+ and an already-installed cryptography package exposing
X25519, Ed25519, AESGCM, and Argon2id (author tested 50.0.0). No network,
credential store, filesystem persistence, or production TLS implementation.
The allocator's persist event and receiver snapshots are abstract atomic durable
transitions; this program does not test a disk, fsync, certificate path, erasure,
CSPRNG, cryptographic security proof, or machine-code constant-time behavior.
"""
import argparse
import copy
import hashlib
import hmac
import itertools
import json
from dataclasses import dataclass
try:
    from cryptography.exceptions import InvalidSignature, InvalidTag
    from cryptography.hazmat.primitives.asymmetric import ed25519, x25519
    from cryptography.hazmat.primitives.ciphers.aead import AESGCM
    from cryptography.hazmat.primitives.kdf.argon2 import Argon2id
    from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat
except ImportError as exc:
    raise SystemExit('Missing cryptography: no toy AEAD fallback; use an environment with that library.') from exc


def check(ok, message):
    if not ok:
        raise AssertionError(message)


def H(data):
    return hashlib.sha256(data).digest()


def mac(key, data):
    return hmac.new(key, data, hashlib.sha256).digest()


def enc(*fields):
    """u16 field count, then each field u32 big-endian length followed by bytes."""
    if len(fields) > 65535:
        raise ValueError('count')
    parts = [len(fields).to_bytes(2, 'big')]
    for field in fields:
        if not isinstance(field, bytes) or len(field) > 65535:
            raise ValueError('field')
        parts.extend([len(field).to_bytes(4, 'big'), field])
    return b''.join(parts)


def dec(data, count):
    if len(data) < 2 or int.from_bytes(data[:2], 'big') != count:
        raise ValueError('count')
    fields, pos = [], 2
    for _ in range(count):
        if pos + 4 > len(data):
            raise ValueError('truncated length')
        size = int.from_bytes(data[pos:pos + 4], 'big')
        pos += 4
        if size > 65535 or pos + size > len(data):
            raise ValueError('length')
        fields.append(data[pos:pos + size])
        pos += size
    if pos != len(data):
        raise ValueError('trailing bytes')
    return tuple(fields)


def extract(salt, ikm):
    return mac(bytes(32) if salt is None else salt, ikm)


def expand(prk, info, length):
    if not isinstance(length, int) or not 0 <= length <= 255 * 32:
        raise ValueError('HKDF length')
    t, blocks = b'', []
    for counter in range(1, (length + 31) // 32 + 1):
        t = mac(prk, t + info + bytes([counter]))
        blocks.append(t)
    return b''.join(blocks)[:length]


def public(key):
    return key.public_key().public_bytes(Encoding.Raw, PublicFormat.Raw)


PROTOCOL = b'CS09-demo'
VERSION = b'1'
SUITE = b'X25519-Ed25519-HKDF-SHA256-AES128GCM'
CLIENT, SERVER = b'client.test', b'server.test'


def derive(prk, transcript_hash, direction, purpose, length):
    info = enc(PROTOCOL, VERSION, SUITE, CLIENT, SERVER, direction, purpose,
               length.to_bytes(2, 'big'), transcript_hash)
    return expand(prk, info, length)


def fixture():
    # Deterministic, disclosed private bytes solely for reproducibility.
    cdh = x25519.X25519PrivateKey.from_private_bytes(bytes(range(1, 33)))
    sdh = x25519.X25519PrivateKey.from_private_bytes(bytes(range(33, 65)))
    csig = ed25519.Ed25519PrivateKey.from_private_bytes(bytes(range(65, 97)))
    ssig = ed25519.Ed25519PrivateKey.from_private_bytes(bytes(range(97, 129)))
    nc, ns = bytes(range(16)), bytes(range(16, 32))
    t0 = enc(PROTOCOL, VERSION, SUITE, CLIENT, SERVER, nc, ns,
             public(cdh), public(sdh), public(csig), public(ssig))
    zc, zs = cdh.exchange(sdh.public_key()), sdh.exchange(cdh.public_key())
    check(zc == zs, 'DH agreement')
    salt = H(enc(PROTOCOL, VERSION, nc, ns))
    prk = extract(salt, zc)
    cm = enc(PROTOCOL, VERSION, b'client-auth', H(t0))
    sm = enc(PROTOCOL, VERSION, b'server-auth', H(t0))
    cs, ss = csig.sign(cm), ssig.sign(sm)
    # Pinned public keys are the external identity assumption in this model.
    csig.public_key().verify(cs, cm)
    ssig.public_key().verify(ss, sm)
    t1 = enc(t0, cs, ss)
    kc = derive(prk, H(t1), b'c2s', b'finished', 32)
    ks = derive(prk, H(t1), b's2c', b'finished', 32)
    fc = mac(kc, enc(b'client-finished', H(t1)))
    fs = mac(ks, enc(b'server-finished', H(t1)))
    check(hmac.compare_digest(fc, mac(kc, enc(b'client-finished', H(t1)))), 'client finished')
    check(hmac.compare_digest(fs, mac(ks, enc(b'server-finished', H(t1)))), 'server finished')
    tf = enc(t1, fc, fs)
    th = H(tf)
    keys = {d: derive(prk, th, d.encode(), b'application', 16) for d in ['c2s', 's2c']}
    return locals()


class Confirmation:
    """Local acceptance state; receiving the peer value is an explicit event."""
    def __init__(self, role, prk, t1):
        if role not in ('client', 'server'):
            raise ValueError('role')
        self.role, self.prk, self.t1 = role, prk, t1
        self.state = 'await_auth'

    def authenticated(self, identity_and_signatures_valid):
        if self.state != 'await_auth':
            raise ValueError('phase')
        self.state = 'await_peer_finished' if identity_and_signatures_valid else 'failed'

    def receive(self, peer_value):
        if self.state != 'await_peer_finished':
            return False
        direction, label = (b's2c', b'server-finished') if self.role == 'client' else (b'c2s', b'client-finished')
        key = derive(self.prk, H(self.t1), direction, b'finished', 32)
        if not hmac.compare_digest(peer_value, mac(key, enc(label, H(self.t1)))):
            self.state = 'failed'
            return False
        self.state = 'established'
        return True

    def application_allowed(self):
        return self.state == 'established'


@dataclass
class Allocator:
    """One writer per directional key; durable high-water mark is exclusive."""
    durable: int = 0
    next: int = 0
    end: int = 0
    block: int = 4
    limit: int = 2**64

    def __post_init__(self):
        if self.block < 1 or not 0 <= self.durable <= self.limit <= 2**64:
            raise ValueError('allocator bounds')

    def reserve(self):
        if self.next != self.end:
            raise ValueError('unconsumed reservation')
        old = self.durable
        new = min(old + self.block, self.limit)
        if old == new:
            raise OverflowError('rekey required')
        # This assignment models atomic durable commit before first use.
        self.durable = new
        self.next, self.end = old, new

    def take(self):
        if self.next == self.end:
            self.reserve()
        seq = self.next
        self.next += 1
        return seq

    def restart(self):
        self.next = self.end = self.durable


def nonce(seq):
    if not isinstance(seq, int) or not 0 <= seq < 2**64:
        raise ValueError('sequence')
    return bytes(4) + seq.to_bytes(8, 'big')


def header(th, direction, seq, version=VERSION, kind=b'data'):
    return enc(PROTOCOL, version, th, direction.encode(), kind, seq.to_bytes(8, 'big'))


def seal(key, th, direction, seq, message):
    if len(message) > 65535:
        raise ValueError('record length')
    a = header(th, direction, seq)
    return a, AESGCM(key).encrypt(nonce(seq), message, a)


@dataclass
class Window:
    width: int = 4
    high: int = -1
    bits: int = 0

    def __post_init__(self):
        if self.width < 1:
            raise ValueError('width')

    def eligible(self, seq):
        if not isinstance(seq, int) or not 0 <= seq < 2**64:
            return False
        return seq > self.high or (self.high - seq < self.width and not (self.bits >> (self.high-seq)) & 1)

    def commit(self, seq):
        if not self.eligible(seq):
            raise ValueError('replay or too old')
        if seq > self.high:
            gap = seq - self.high
            self.bits = 0 if gap >= self.width else (self.bits << gap) & ((1 << self.width)-1)
            self.high = seq
            self.bits |= 1
        else:
            self.bits |= 1 << (self.high-seq)


class Receiver:
    def __init__(self, key, th, direction, window=None):
        self.key, self.th, self.direction = key, th, direction
        self.window = window if window is not None else Window()

    def receive(self, packet):
        a, ciphertext = packet
        if not 16 <= len(ciphertext) <= 65535 + 16:
            return 'length', None
        try:
            p, v, th, direction, kind, sb = dec(a, 6)
            if (p, v, th, direction, kind) != (PROTOCOL, VERSION, self.th, self.direction.encode(), b'data') or len(sb) != 8:
                return 'header', None
            seq = int.from_bytes(sb, 'big')
            if not self.window.eligible(seq):
                return 'replay', None
            plain = AESGCM(self.key).decrypt(nonce(seq), ciphertext, a)
        except (ValueError, InvalidTag):
            return 'authentication', None
        # Sequential model. In concurrent code, test-and-commit must be atomic
        # and recheck eligibility after authentication. Snapshot before delivery.
        self.window.commit(seq)
        return 'accepted', plain


def expect_exception(fn, cls):
    try:
        fn()
    except cls:
        return
    raise AssertionError('expected ' + cls.__name__)


def run():
    cases = []
    # Published known-answer tests, not vectors produced only by our own code.
    hm = mac(bytes([0x0b])*20, b'Hi There').hex()
    check(hm == 'b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7', 'RFC4231 case1')
    ikm, salt, info = bytes([0x0b])*22, bytes(range(13)), bytes(range(0xf0,0xfa))
    prk = extract(salt, ikm)
    okm = expand(prk, info, 42)
    check(prk.hex() == '077709362c2e32df0ddc3f0dc47bba6390b6c73bb50f9c3122ec844ad7c2b3e5', 'RFC5869 PRK')
    check(okm.hex() == '3cb25f25faacd57a90434f64d0362f2a2d2d0a90cf1a5a4c5db02d56ecc4c5bf34007208d5b887185865', 'RFC5869 OKM')
    expect_exception(lambda: expand(prk,b'',8161), ValueError)
    cases += ['RFC4231-HMAC-SHA256-case1','RFC5869-HKDF-SHA256-A1','HKDF-length-bound']
    check(b'ab'+b'c' == b'a'+b'bc', 'raw collision')
    check(enc(b'ab',b'c') != enc(b'a',b'bc'), 'encoded distinction')
    for fs in [(),(b'',),(b'ab',b'c'),(b'a',b'bc')]:
        check(dec(enc(*fs),len(fs)) == fs, 'encoding roundtrip')
    for data in [enc(b'a')[:-1],enc(b'a')+b'x',b'\x00\x01\xff\xff\xff\xff']:
        expect_exception(lambda d=data:dec(d,1), ValueError)
    cases += ['injective-example-and-roundtrips','malformed-length-trailing-rejected']
    f=fixture(); keys=f['keys'];th=f['th'];tc=H(f['t1'])
    cp=Confirmation('client',f['prk'],f['t1']);sp=Confirmation('server',f['prk'],f['t1'])
    check(not cp.receive(f['fs']) and cp.state=='await_auth','Finished before authentication')
    cp.authenticated(True);sp.authenticated(True)
    check(not cp.application_allowed() and not sp.application_allowed(),'lost Finished cannot establish')
    check(cp.receive(f['fs']) and cp.application_allowed() and not sp.application_allowed(),'one-sided delivery')
    check(sp.receive(f['fc']) and sp.application_allowed(),'both delivered')
    badp=Confirmation('client',f['prk'],f['t1']);badp.authenticated(True)
    check(not badp.receive(f['fc']) and badp.state=='failed','wrong Finished no establishment')
    cases += ['confirmation-auth-before-finished','confirmation-loss-and-one-sided-delivery','confirmation-wrong-message-fails-closed']
    check(keys['c2s'] != keys['s2c'] and f['kc'] != f['ks'], 'direction separation')
    check(derive(f['prk'],th,b'c2s',b'application',16) != derive(f['prk'],th,b'c2s',b'exporter',16), 'purpose separation')
    # Remove ONLY direction from KDF info: both calls are identical by definition.
    def no_direction(direction):
        # BUG: direction is accepted by the interface but omitted from info.
        return expand(f['prk'],enc(PROTOCOL,VERSION,SUITE,CLIENT,SERVER,b'application',(16).to_bytes(2,'big'),th),16)
    check(no_direction(b'c2s') == no_direction(b's2c'), 'omitted label equal keys')
    expect_exception(lambda:f['ssig'].public_key().verify(f['ss'],enc(PROTOCOL,VERSION,b'client-auth',H(f['t0']))), InvalidSignature)
    expect_exception(lambda:f['ssig'].public_key().verify(f['ss'],enc(PROTOCOL,b'2',b'server-auth',H(f['t0']))), InvalidSignature)
    check(not hmac.compare_digest(f['fs'], mac(f['kc'],enc(b'client-finished',tc))), 'reflected finished rejected')
    for index in [1,3,4,7,8,9,10]:
        fields=list(dec(f['t0'],11)); fields[index] = fields[index][:-1]+bytes([fields[index][-1]^1])
        changed_t0=enc(*fields)
        expect_exception(lambda t=changed_t0:f['ssig'].public_key().verify(f['ss'],enc(PROTOCOL,VERSION,b'server-auth',H(t))),InvalidSignature)
        changed_t1=enc(changed_t0,f['cs'],f['ss'])
        changed_k=derive(f['prk'],H(changed_t1),b's2c',b'finished',32)
        check(not hmac.compare_digest(f['fs'],mac(changed_k,enc(b'server-finished',H(changed_t1)))),'modified transcript finished')
    cases += ['seven-transcript-field-mutations-reject-signature-and-finished']
    cases += ['role-purpose-keys-distinct','omitted-direction-collapses-key-domain','wrong-signature-role-version-rejected','reflected-finished-rejected']
    al=Allocator(); used=[]
    used += [al.take(),al.take()]; before={'durable':al.durable,'next':al.next,'end':al.end};al.restart();used += [al.take(),al.take()]
    check(used == [0,1,4,5], 'crash skip')
    # Enumerate every crash/allocate pattern of length 9, 512 traces.
    for actions in itertools.product([0,1],repeat=9):
        a=Allocator(); seen=set()
        for action in actions:
            if action:a.restart()
            else:
                n=a.take();check(n not in seen,'allocator uniqueness');seen.add(n)
    limited=Allocator(block=4,limit=5)
    check([limited.take() for _ in range(5)] == list(range(5)), 'end clipping')
    expect_exception(limited.take,OverflowError)
    cases += ['reserve-before-use-restart-0-1-4-5','512-crash-allocation-traces','exhaustion-fails-closed']
    maxpacket=seal(keys['c2s'],th,'c2s',1000,b'x'*65535)
    check(Receiver(keys['c2s'],th,'c2s').receive(maxpacket)[0]=='accepted','maximum record')
    expect_exception(lambda:seal(keys['c2s'],th,'c2s',1001,b'x'*65536),ValueError)
    oversize=(header(th,'c2s',1001),AESGCM(keys['c2s']).encrypt(nonce(1001),b'x'*65536,header(th,'c2s',1001)))
    check(Receiver(keys['c2s'],th,'c2s').receive(oversize)[0]=='length','oversized authentic ciphertext')
    check(Receiver(keys['c2s'],th,'c2s').receive((header(th,'c2s',1002),b'x'*15))[0]=='length','short tag')
    cases += ['record-65535-accepted-65536-rejected-short-tag-rejected']
    packets={s:seal(keys['c2s'],th,'c2s',s,('m'+str(s)).encode()) for s in [0,1,2,4,5,99]}
    rx=Receiver(keys['c2s'],th,'c2s'); trace=[]
    for s in [0,2,1,2,5,0,4]:
        status, plain=rx.receive(packets[s]);trace.append({'seq':s,'status':status,'high':rx.window.high,'bits':format(rx.window.bits,'04b')})
    check([x['status'] for x in trace] == ['accepted','accepted','accepted','replay','accepted','replay','accepted'], 'window trace')
    state=copy.deepcopy(rx.window)
    a,c=packets[99];bad=(a,c[:-1]+bytes([c[-1]^1]))
    check(rx.receive(bad)[0] == 'authentication' and rx.window==state, 'bad high seq no slide')
    check(Receiver(keys['s2c'],th,'s2c').receive(packets[0])[0]=='header', 'opposite direction reject')
    a,c=packets[0]
    check(Receiver(keys['s2c'],th,'s2c').receive((header(th,'s2c',0),c))[0]=='authentication','direction rewrite fails')
    check(rx.receive((header(th,'c2s',99,b'2'),packets[99][1]))[0]=='header', 'version reject')
    restored=Receiver(keys['c2s'],th,'c2s',copy.deepcopy(rx.window))
    check(restored.receive(packets[5])[0]=='replay','receiver snapshot replay')
    rolled=Receiver(keys['c2s'],th,'c2s')
    check(rolled.receive(packets[0])[0]=='accepted','receiver rollback replay witness')
    cases += ['window-out-of-order-duplicate-too-old','forged-high-sequence-does-not-slide','record-role-version-rejected','receiver-snapshot-retains-replay-state','receiver-rollback-replay-witness']
    # Fresh receivers vs ideal accepted-set oracle for all short delivery traces.
    checks=0
    for seqs in itertools.product(range(6),repeat=5):
        w=Window();accepted=set();hi=-1
        for s in seqs:
            expected=s not in accepted and (hi < 0 or s > hi-4)
            check(w.eligible(s)==expected,'window oracle')
            if expected:w.commit(s);accepted.add(s);hi=max(hi,s)
        checks+=1
    cases += ['7776-window-traces-against-set-oracle']
    # Counter rollback reuses the actual AES-GCM keystream: ciphertext-body XOR.
    m1,m2=b'balance=100',b'balance=900';n=nonce(0);aad=header(th,'c2s',0)
    c1=AESGCM(keys['c2s']).encrypt(n,m1,aad);c2=AESGCM(keys['c2s']).encrypt(n,m2,aad)
    xor=lambda a,b:bytes(x^y for x,y in zip(a,b))
    check(xor(c1[:-16],c2[:-16])==xor(m1,m2),'nonce-reuse plaintext XOR')
    cases += ['actual-AESGCM-nonce-reuse-XOR-witness']
    argon=Argon2id(salt=bytes([2])*16,length=32,iterations=3,lanes=4,memory_cost=32,secret=bytes([3])*8,ad=bytes([4])*12).derive(bytes([1])*32)
    check(argon.hex()=='0d640df58d78766c08c037a34a8b53c9d01ef0452d75b65eb52520e96b01e659','RFC9106 Argon2id')
    cases += ['RFC9106-Argon2id-5.3-with-secret-and-associated-data']
    vectors={'argon2id_rfc9106':argon.hex(),'protocol':'CS09-demo, synthetic fixed public fixtures',
      'hmac_rfc4231':hm,'hkdf_rfc5869_prk':prk.hex(),'hkdf_rfc5869_okm':okm.hex(),
      'client_dh_public':public(f['cdh']).hex(),'server_dh_public':public(f['sdh']).hex(),
      'shared_secret_TEST_ONLY':f['zc'].hex(),'transcript0':f['t0'].hex(),'transcript0_sha256':H(f['t0']).hex(),
      'client_signature':f['cs'].hex(),'server_signature':f['ss'].hex(),'transcript1_sha256':H(f['t1']).hex(),
      'client_finished':f['fc'].hex(),'server_finished':f['fs'].hex(),'final_transcript_sha256':th.hex(),
      'application_keys_TEST_ONLY':{k:v.hex() for k,v in keys.items()},
      'record0_header':packets[0][0].hex(),'record0_nonce':nonce(0).hex(),'record0_ciphertext_and_tag':packets[0][1].hex(),
      'allocator_before_crash':before,'allocator_used':used,'window_trace':trace,
      'nonce_reuse_plaintext_xor':xor(m1,m2).hex(),'nonce_reuse_ciphertext_body_xor':xor(c1[:-16],c2[:-16]).hex()}
    return {'status':'passed','checks':cases,'check_count':len(cases),'allocator_traces':512,'window_traces':checks,
      'scope':'Real standard-library HMAC/HKDF and installed cryptography X25519/Ed25519/AESGCM/Argon2id; finite state model, synthetic public keys only; no deployable protocol or cryptographic proof.', 'vectors':vectors}


if __name__=='__main__':
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--vectors',action='store_true',help='print full reproducible fixture bytes')
    args=parser.parse_args()
    result=run()
    if not args.vectors:result.pop('vectors')
    print(json.dumps(result,ensure_ascii=False,indent=2))
