#!/usr/bin/env python3
"""Educational execution fixtures, NOT a cryptographic implementation.
All secret-looking values derive from PUBLIC deterministic labels. They are
intentionally reproducible and MUST NOT be used as keys. The stateful signer
simulates an atomic irreversible reservation CONTRACT in memory; it provides
no persistence, concurrency control, rollback protection, or production API.
Python 3 standard library. Checks remain enabled with python -O.
"""
from hashlib import sha256
from itertools import combinations, product
import json

N = 32

def check(condition, label):
    if not condition:
        raise RuntimeError(label)

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

def u32(i):
    if type(i) is not int or not 0 <= i < 2**32:
        raise ValueError('u32 outside range')
    return i.to_bytes(4, 'big')

def height(h):
    return type(h) is int and 0 <= h <= 31

def sequence(x):
    return type(x) in (list, tuple)

def word(x):
    return type(x) is bytes and len(x) == N

def leaf(i, value):
    if type(value) is not bytes:
        raise ValueError('leaf value must be bytes')
    return H(b'\x00' + u32(i) + u32(len(value)) + value)

def node(level, left, right, tag=b'\x01'):
    if not 1 <= level <= 31 or not word(left) or not word(right):
        raise ValueError('invalid internal node')
    return H(tag + bytes([level]) + left + right)

class Merkle:
    def __init__(self, values):
        n = len(values)
        if n == 0 or n & (n-1) or n >= 2**32:
            raise ValueError('positive power-of-two leaf count required')
        self.h = n.bit_length()-1
        self.values = list(values)
        self.levels = [[leaf(i, x) for i, x in enumerate(values)]]
        for level in range(1, self.h+1):
            a = self.levels[-1]
            self.levels.append([node(level, a[i], a[i+1]) for i in range(0, len(a), 2)])
        self.root = self.levels[-1][0]

    def proof(self, i):
        if type(i) is not int or not 0 <= i < len(self.values):
            raise ValueError('bad leaf index')
        return [self.levels[j][(i >> j) ^ 1] for j in range(self.h)]

    def multiproof(self, indices):
        indices = list(indices)
        if not indices or len(set(indices)) != len(indices) or any(type(i) is not int or not 0 <= i < len(self.values) for i in indices):
            raise ValueError('unique nonempty valid target indices required')
        active = set(indices)
        proof = []
        for level in range(self.h):
            for i in sorted(active):
                if i ^ 1 not in active:
                    proof.append((level, i ^ 1, self.levels[level][i ^ 1]))
            active = {i//2 for i in active}
        return proof

def verify_path(h, root, i, value, proof):
    if not height(h) or not word(root) or type(i) is not int or not 0 <= i < 2**h or type(value) is not bytes:
        return False
    if type(proof) not in (list, tuple) or len(proof) != h or not all(word(p) for p in proof):
        return False
    v = leaf(i, value)
    for level, sibling in enumerate(proof, 1):
        v = node(level, sibling, v) if i & 1 else node(level, v, sibling)
        i //= 2
    return v == root

def verify_multi(h, root, items, proof):
    """Proof is a sequence of (level,index,digest); reject duplicates and extras."""
    if not height(h) or not word(root) or not sequence(items) or not items or not sequence(proof):
        return False
    if any(not sequence(item) or len(item) != 2 for item in items) or any(not sequence(item) or len(item) != 3 for item in proof):
        return False
    active = {}
    for i, value in items:
        if type(i) is not int or i in active or not 0 <= i < 2**h or type(value) is not bytes:
            return False
        active[i] = leaf(i, value)
    supplied = {}
    for level, i, digest in proof:
        if type(level) is not int or type(i) is not int or not 0 <= level < h or not 0 <= i < 2**(h-level) or not word(digest):
            return False
        if (level, i) in supplied:
            return False
        supplied[level, i] = digest
    used = set()
    for level in range(h):
        parents = {}
        for p in sorted({i//2 for i in active}):
            pair = []
            for i in (2*p, 2*p+1):
                if i in active:
                    pair.append(active[i])
                elif (level, i) in supplied:
                    pair.append(supplied[level, i]); used.add((level, i))
                else:
                    return False
            parents[p] = node(level+1, *pair)
        active = parents
    return active == {0: root} and used == set(supplied)

class SparseMerkle:
    """Only nondefault node digests are stored. Key is an h-bit integer."""
    def __init__(self, h):
        if not height(h):
            raise ValueError('bad height')
        self.h = h
        self.defaults = [H(b'\x02')]
        for j in range(1, h+1):
            self.defaults.append(node(j, self.defaults[-1], self.defaults[-1], b'\x04'))
        self.nodes = {}

    def at(self, level, i):
        return self.nodes.get((level, i), self.defaults[level])

    @property
    def root(self):
        return self.at(self.h, 0)

    def set(self, key, value):
        if type(key) is not int or not 0 <= key < 2**self.h or (value is not None and type(value) is not bytes):
            raise ValueError('bad key/value; None means absent, b"" is present')
        digest = sparse_leaf(key, value)
        i = key
        for j in range(self.h+1):
            if digest == self.defaults[j]:
                self.nodes.pop((j, i), None)
            else:
                self.nodes[j, i] = digest
            if j < self.h:
                sibling = self.at(j, i ^ 1)
                digest = node(j+1, sibling, digest, b'\x04') if i & 1 else node(j+1, digest, sibling, b'\x04')
                i //= 2

    def proof(self, key):
        if type(key) is not int or not 0 <= key < 2**self.h:
            raise ValueError('bad key')
        return [self.at(j, (key >> j) ^ 1) for j in range(self.h)]

def sparse_leaf(key, value):
    return H(b'\x02') if value is None else H(b'\x03' + u32(key) + u32(len(value)) + value)

def verify_sparse(h, root, key, value, proof):
    if not height(h) or not word(root) or type(key) is not int or not 0 <= key < 2**h or (value is not None and type(value) is not bytes):
        return False
    if type(proof) not in (list, tuple) or len(proof) != h or not all(word(x) for x in proof):
        return False
    digest = sparse_leaf(key, value)
    for j, sibling in enumerate(proof, 1):
        digest = node(j, sibling, digest, b'\x04') if key & 1 else node(j, digest, sibling, b'\x04')
        key //= 2
    return digest == root

def bits_ok(message, length):
    return type(message) is str and len(message) == length and all(x in '01' for x in message)

def F(x):
    return H(b'\x10' + x)

def lamport_fixture(length=4, namespace=0):
    if type(length) is not int or length <= 0:
        raise ValueError('positive bit length required')
    sk = [[H(b'\xf0' + u32(namespace) + u32(i) + bytes([b])) for b in (0, 1)] for i in range(length)]
    return sk, [[F(x) for x in pair] for pair in sk]

def lamport_sign(sk, message):
    if not bits_ok(message, len(sk)):
        raise ValueError('wrong message bit length')
    return [sk[i][int(b)] for i, b in enumerate(message)]

def lamport_verify(pk, message, signature):
    if not sequence(pk) or not sequence(signature) or not pk or not bits_ok(message, len(pk)) or len(signature) != len(pk):
        return False
    if any(not sequence(pair) or len(pair) != 2 or not all(word(x) for x in pair) for pair in pk) or not all(word(x) for x in signature):
        return False
    return all(F(x) == pk[i][int(b)] for i, (b, x) in enumerate(zip(message, signature)))

def base_digits(value, base, length):
    digits = [0]*length
    for i in range(length-1, -1, -1):
        digits[i] = value % base; value //= base
    if value:
        raise ValueError('not enough digits')
    return digits

def wots_digits(message, base):
    if type(base) is not int or base < 2 or not sequence(message) or not message or any(type(x) is not int or not 0 <= x < base for x in message):
        raise ValueError('invalid base/digits')
    maximum = len(message)*(base-1)
    length = 1
    while base**length <= maximum:
        length += 1
    checksum = sum(base-1-x for x in message)
    return list(message) + base_digits(checksum, base, length)

def WF(x):
    # Classic fixed chaining function, not address-based LM-OTS or WOTS+.
    return H(b'\x20' + x)

def chain(x, steps):
    for _ in range(steps):
        x = WF(x)
    return x

def wots_fixture(message_length=2, base=4):
    length = len(wots_digits([0]*message_length, base))
    sk = [H(b'\xf1' + u32(i)) for i in range(length)]
    return sk, [chain(x, base-1) for x in sk]

def wots_sign(sk, message, base):
    digits = wots_digits(message, base)
    if len(digits) != len(sk):
        raise ValueError('wrong key/message parameters')
    return [chain(x, d) for x, d in zip(sk, digits)]

def wots_verify(pk, message, base, signature):
    if not sequence(pk) or not sequence(signature):
        return False
    try:
        digits = wots_digits(message, base)
    except ValueError:
        return False
    return len(pk) == len(signature) == len(digits) and all(word(x) for x in [*signature, *pk]) and all(chain(x, base-1-d) == y for x, d, y in zip(signature, digits, pk))

def encode_pk(pk):
    return u32(len(pk)) + b''.join(x for pair in pk for x in pair)

class StatefulFixture:
    def __init__(self, h=2, bits=4):
        self.h, self.bits, self.next = h, bits, 0
        pairs = [lamport_fixture(bits, namespace=100+i) for i in range(2**h)]
        self.keys = pairs
        self.tree = Merkle([encode_pk(pk) for sk, pk in pairs])

    def sign(self, message, fail_after_reservation=False):
        if not bits_ok(message, self.bits):
            raise ValueError('bad message before reservation')
        if self.next == len(self.keys):
            raise ValueError('exhausted')
        q = self.next
        self.next += 1  # Abstract atomic irreversible reservation, only simulated.
        if fail_after_reservation:
            return None
        sk, pk = self.keys[q]
        return q, pk, lamport_sign(sk, message), self.tree.proof(q)

def tree_verify(h, bits, root, message, signature):
    if not sequence(signature) or len(signature) != 4 or type(bits) is not int or bits <= 0:
        return False
    q, pk, sig, path = signature
    if not sequence(pk):
        return False
    return len(pk) == bits and bits_ok(message, bits) and lamport_verify(pk, message, sig) and verify_path(h, root, q, encode_pk(pk), path)

def hx(x):
    if type(x) is bytes:
        return x.hex()
    if type(x) in (tuple, list):
        return [hx(t) for t in x]
    if type(x) is dict:
        return {str(k): hx(v) for k, v in x.items()}
    return x

def main():
    tree = Merkle([bytes([65+i]) for i in range(8)])
    proof = tree.proof(2)
    check(verify_path(3, tree.root, 2, b'C', proof), 'main inclusion')
    path_tests = 0
    for h in range(6):
        t = Merkle([u32(i)+b':value' for i in range(2**h)])
        for i in range(2**h):
            p = t.proof(i)
            check(verify_path(h, t.root, i, t.values[i], p), 'path')
            check(not verify_path(h, t.root, i, t.values[i]+b'!', p), 'changed value')
            check(not verify_path(h, t.root, i, t.values[i], p+[b'0'*32]), 'extra node')
            path_tests += 1
    targets = [1, 2, 6]; mp = tree.multiproof(targets)
    check(len(mp) == 4, 'four siblings')
    subsets = 0
    for k in range(1, 9):
        for subset in combinations(range(8), k):
            p = tree.multiproof(subset); items = [(i, tree.values[i]) for i in subset]
            check(verify_multi(3, tree.root, items, p), 'multiproof')
            check(not verify_multi(3, tree.root, items+items[:1], p), 'duplicate target')
            if p:
                check(not verify_multi(3, tree.root, items, p[:-1]), 'missing multiproof')
                check(not verify_multi(3, tree.root, items, p+p[:1]), 'duplicate proof')
            check(not verify_multi(3, tree.root, items, p+[(0, subset[0], tree.levels[0][subset[0]])]), 'unused proof')
            subsets += 1
    sm = SparseMerkle(3); sm.set(1, b'red'); sm.set(6, b'blue')
    sparse_root = sm.root; absence = sm.proof(2); stored = len(sm.nodes)
    check(verify_sparse(3, sparse_root, 2, None, absence), 'absent key')
    check(not verify_sparse(3, sparse_root, 2, b'', absence), 'empty value not absent')
    sm.set(2, b'green'); updated_root = sm.root
    check(not verify_sparse(3, updated_root, 2, None, absence), 'stale absence')
    sm.set(2, None); check(sm.root == sparse_root, 'delete restores snapshot root')
    sparse_tests = 0
    # Exhaust all 256 occupancy maps at h=3, with empty byte strings among present values.
    for mask in range(256):
        t = SparseMerkle(3)
        for key in range(8):
            if mask >> key & 1: t.set(key, b'' if key & 1 else u32(key))
        for key in range(8):
            value = (b'' if key & 1 else u32(key)) if mask >> key & 1 else None
            check(verify_sparse(3, t.root, key, value, t.proof(key)), 'sparse map')
            opposite = b'fake' if value is None else None
            check(not verify_sparse(3, t.root, key, opposite, t.proof(key)), 'sparse opposite')
            sparse_tests += 1
    sk, pk = lamport_fixture(); sig = lamport_sign(sk, '1010')
    messages = [''.join(x) for x in product('01', repeat=4)]
    lamport_pairs = 0
    for message in messages:
        s = lamport_sign(sk, message) # Independent hypothetical one-use tests, NOT a safe signing workload.
        for other in messages:
            check(lamport_verify(pk, other, s) == (message == other), 'Lamport validation matrix')
            lamport_pairs += 1
    sig0, sig1 = lamport_sign(sk, '0000'), lamport_sign(sk, '1111')
    forge = [sig1[i] if b == '1' else sig0[i] for i, b in enumerate('0101')]
    check(lamport_verify(pk, '0101', forge), 'reuse forgery')
    wsk, wpk = wots_fixture(); wm=[1,1]; wd=wots_digits(wm,4); ws=wots_sign(wsk,wm,4)
    check(wots_verify(wpk,wm,4,ws), 'Winternitz')
    forward=[chain(ws[0],1),ws[1]]
    check(all(chain(x,3-d)==y for x,d,y in zip(forward,[2,1],wpk)), 'no-checksum forward forge')
    check(not wots_verify(wpk,[2,1],4,forward+ws[2:]), 'checksum blocks this forge')
    anti_domination = 0
    for a in product(range(4), repeat=2):
        for b in product(range(4), repeat=2):
            if a == b: continue
            da,db=wots_digits(a,4),wots_digits(b,4)
            check(any(y<x for x,y in zip(da,db)), 'checksum retreat')
            anti_domination += 1
    st=StatefulFixture(); s0=st.sign('1010'); lost=st.sign('1111',True); s2=st.sign('0101'); s3=st.sign('1100')
    check(lost is None and [s0[0],s2[0],s3[0]]==[0,2,3] and st.next==4, 'reserve/burn')
    for m,s in [('1010',s0),('0101',s2),('1100',s3)]:
        check(tree_verify(2,4,st.tree.root,m,s), 'tree signature')
    try:
        st.sign('0000')
    except ValueError:
        exhausted=True
    else:
        exhausted=False
    check(exhausted,'exhaustion')
    # Deliberately violate the state contract to produce an actual new-message signature.
    broken=StatefulFixture(); b0=broken.sign('0000'); broken.next=0; b1=broken.sign('1111')
    fq, fpk, _, fp=b0; fs=[b1[2][i] if b=='1' else b0[2][i] for i,b in enumerate('0101')]
    check(tree_verify(2,4,broken.tree.root,'0101',(fq,fpk,fs,fp)), 'rollback tree forgery')
    malformed_checks = [
        not verify_multi(3, tree.root, None, []),
        not verify_multi(3, tree.root, [(2, b'C')], [3]),
        not verify_multi(3, tree.root, [(2, b'C'), (2, b'C')], []),
        not lamport_verify(None, '1010', sig),
        not lamport_verify([7]*4, '1010', sig),
        not lamport_verify(pk, '1010', None),
        not wots_verify(wpk, 7, 4, ws),
        not wots_verify(wpk, [1,1], 4, None),
        not tree_verify(2,4,st.tree.root,'1010',[0]),
        not tree_verify(2,4,st.tree.root,'1010',(0,None,[],[])),
        not verify_path(3,tree.root,2,b'C',[H(b'\x01')]*3),
        not verify_sparse(3,sparse_root,True,None,absence),
    ]
    check(all(malformed_checks), 'malformed structured input rejection')
    report={
      'warning':'PUBLIC DETERMINISTIC FIXTURES: NOT SECURITY KEYS OR A PRODUCTION IMPLEMENTATION',
      'malformed_structured_cases':len(malformed_checks),
      'merkle':{'values':['A','B','C','D','E','F','G','H'],'h':3,'root':tree.root,'index':2,'value':'C','proof':proof,'verified_paths':path_tests,
                'wrong_index_rejected':not verify_path(3,tree.root,3,b'C',proof),'wrong_depth_rejected':not verify_path(2,tree.root,2,b'C',proof[:2]),'reversed_path_rejected':not verify_path(3,tree.root,2,b'C',proof[::-1])},
      'multiproof':{'targets':targets,'proof':mp,'hash_count':len(mp),'independent_path_hashes':9,'all_nonempty_subsets':subsets},
      'sparse':{'root':sparse_root,'proof_absent_010':absence,'stored_nondefault_nodes':stored,'updated_root':updated_root,'checked_map_key_pairs':sparse_tests},
      'lamport':{'message':'1010','public_key':pk,'signature':sig,'validation_pairs':lamport_pairs,'reuse_forged_message':'0101','reuse_forgery_accepts':True},
      'winternitz':{'base':4,'message':wm,'signed_digits':wd,'target_digits':wots_digits([2,1],4),'signature':ws,'public_key':wpk,'sign_hashes':sum(wd),'verify_hashes':len(wd)*3-sum(wd),'keygen_hashes':len(wd)*3,'ordered_distinct_digit_pairs':anti_domination},
      'tree_signature':{'root':st.tree.root,'returned_indices':[0,2,3],'burned_indices':[1],'next':st.next,'exhaustion_rejected':exhausted,'rollback_forged_message':'0101','rollback_forgery_accepts':True,'signature_index0':s0}
    }
    check(all(report['merkle'][k] for k in ('wrong_index_rejected','wrong_depth_rejected','reversed_path_rejected')), 'path negative checks')
    print(json.dumps(hx(report), ensure_ascii=False, indent=2))

if __name__ == '__main__':
    main()
