#!/usr/bin/env python3
"""Offline public-vector reference, RFC 9380 and RFC 9497 P256-SHA256 modes 0/1.
Not constant time. No secure key generation, networking, batching or secret inputs.
Single-element VOPRF fails closed on zero composite coefficient or identity
reconstruction. PrivateInput follows the stricter RFC 9497 section 5.1 bound.
Random scalars are supplied explicitly and nonzero (Verified erratum 8392).
"""
import hashlib
import json

P = 2**256 - 2**224 + 2**192 + 2**96 - 1
Q = int('ffffffff00000000ffffffffffffffffbce6faada7179e84f3b9cac2fc632551', 16)
A = P - 3
B = int('5ac635d8aa3a93e7b3ebbd55769886bc651d06b0cc53b0f63bce3c3e27d2604b', 16)
Z = P - 10
G = (int('6b17d1f2e12c4247f8bce6e563a440f277037d812deb33a0f4a13945d898c296', 16),
     int('4fe342e2fe1a7f9b8ee7eb4a7c0f9e162bce33576b315ececbb6406837bf51f5', 16))
IDENTITY = None

def require(ok, message):
    if not ok:
        raise ValueError(message)

def integer(x, low, high):
    require(type(x) is int and low <= x < high, 'noncanonical integer')
    return x

def octets(x):
    require(type(x) is bytes, 'expected bytes')
    return x

def i2(n, size):
    integer(n, 0, 256**size)
    return n.to_bytes(size, 'big')

def frame(x):
    return i2(len(x), 2) + x

def sha(x):
    octets(x); require(len(x) < 2**61, 'SHA256 input length')
    return hashlib.sha256(x).digest()

def input_bytes(x):
    octets(x)
    require(len(x) < 65535, 'private input too long')
    return x

def scalar(x, nonzero=False):
    return integer(x, int(nonzero), Q)

def point(t, nonzero=False):
    if t is None:
        require(not nonzero, 'identity point')
        return t
    require(type(t) is tuple and len(t) == 2, 'point pair required')
    x, y = t
    integer(x, 0, P); integer(y, 0, P)
    require(y*y % P == (x*x*x + A*x + B) % P, 'not on curve')
    return t

def add(t, u):
    if t is None: return u
    if u is None: return t
    x, y = t; a, b = u
    if x == a:
        if (y+b) % P == 0: return None
        slope = (3*x*x+A) * pow(2*y, -1, P) % P
    else:
        slope = (b-y) * pow(a-x, -1, P) % P
    c = (slope*slope-x-a) % P
    return c, (slope*(x-c)-y) % P

def mul(k, t):
    scalar(k); point(t)
    out = None
    while k:
        if k & 1: out = add(out, t)
        t = add(t, t)
        k >>= 1
    return out

def encode(t):
    point(t, True)
    x, y = t
    return bytes((2+(y & 1),)) + i2(x, 32)

def decode(buf):
    octets(buf)
    require(len(buf) == 33 and buf[0] in (2, 3), 'compressed point encoding')
    x = int.from_bytes(buf[1:], 'big')
    integer(x, 0, P)
    v = (x*x*x+A*x+B) % P
    y = pow(v, (P+1)//4, P)
    require(y*y % P == v, 'non-curve x')
    if (y & 1) != (buf[0] & 1): y = (-y) % P
    t = point((x, y), True)
    require(encode(t) == buf, 'noncanonical point')
    return t

def decode_scalar(buf):
    octets(buf); require(len(buf) == 32, 'scalar length')
    return scalar(int.from_bytes(buf, 'big'))

def xmd(msg, dst, length):
    octets(msg); octets(dst); integer(length, 1, 8161)
    require(0 < len(dst) <= 255, 'short nonempty DST required')
    dstp = dst + i2(len(dst), 1)
    b0 = sha(bytes(64) + msg + i2(length, 2) + b'\x00' + dstp)
    previous = sha(b0 + b'\x01' + dstp)
    blocks = [previous]
    for i in range(2, (length+31)//32+1):
        previous = sha(bytes(a ^ b for a, b in zip(b0, previous)) + i2(i, 1) + dstp)
        blocks.append(previous)
    return b''.join(blocks)[:length]

def hash_field(msg, dst, count=2, modulus=P):
    integer(count, 1, 171)
    require(modulus in (P, Q), 'fixed P256 moduli only')
    buf = xmd(msg, dst, 48*count)
    return [int.from_bytes(buf[48*i:48*(i+1)], 'big') % modulus for i in range(count)]

def swu(u):
    integer(u, 0, P)
    t = Z*u*u % P
    denom = (t*t+t) % P
    if denom == 0:
        x1 = B * pow(Z*A % P, -1, P) % P
    else:
        x1 = (-B * pow(A, -1, P) * (1+pow(denom, -1, P))) % P
    gx1 = (x1*x1*x1+A*x1+B) % P
    y = pow(gx1, (P+1)//4, P)
    first = y*y % P == gx1
    x = x1 if first else t*x1 % P
    if not first:
        y = pow((x*x*x+A*x+B) % P, (P+1)//4, P)
    if (y & 1) != (u & 1): y = (-y) % P
    result = point((x, y))
    return result, {'exception': denom == 0, 'branch': 1 if first else 2}

def hash_curve(msg, dst):
    us = hash_field(msg, dst)
    q0, a0 = swu(us[0]); q1, a1 = swu(us[1])
    return add(q0, q1), {'u': us, 'points': [q0, q1], 'maps': [a0, a1]}

def context(mode):
    integer(mode, 0, 2)
    return b'OPRFV1-' + bytes((mode,)) + b'-P256-SHA256'

def hscalar(msg, mode, prefix=b'HashToScalar-'):
    return hash_field(msg, prefix+context(mode), 1, Q)[0]

def hgroup(msg, mode):
    t, log = hash_curve(input_bytes(msg), b'HashToGroup-'+context(mode))
    return point(t, True), log

def derive(seed, info, mode):
    octets(seed); require(len(seed) == 32, 'seed length')
    input_bytes(info)
    for counter in range(256):
        key = hscalar(seed+frame(info)+i2(counter, 1), mode, b'DeriveKeyPair')
        if key:
            return key, mul(key, G), counter
    raise ValueError('DeriveKeyPairError')

def blind(msg, r, mode):
    scalar(r, True)
    t, _ = hgroup(msg, mode)
    return encode(mul(r, t))

def evaluate_blind(key, request):
    scalar(key, True)
    return encode(mul(key, decode(request)))

def finish(msg, r, response):
    input_bytes(msg); scalar(r, True)
    n = mul(pow(r, -1, Q), decode(response))
    return sha(frame(msg)+frame(encode(n))+b'Finalize')

def evaluate(key, msg, mode):
    scalar(key, True)
    t, _ = hgroup(msg, mode)
    return sha(frame(msg)+frame(encode(mul(key, t)))+b'Finalize')

def composites(public, request, response, hash_scalar=hscalar):
    pk, c0, d0 = decode(public), decode(request), decode(response)
    seed = sha(frame(public)+frame(b'Seed-'+context(1)))
    transcript = frame(seed)+b'\x00\x00'+frame(request)+frame(response)+b'Composite'
    coefficient = scalar(hash_scalar(transcript, 1))
    require(coefficient != 0, 'degenerate composite')
    m, z = mul(coefficient, c0), mul(coefficient, d0)
    return pk, m, z, coefficient

def challenge(public, m, z, t2, t3):
    return b''.join(frame(x) for x in (public, encode(m), encode(z), encode(t2), encode(t3)))+b'Challenge'

def prove(key, public, request, response, nonce):
    scalar(key, True); scalar(nonce, True)
    pk, m, z, coefficient = composites(public, request, response)
    require(pk == mul(key, G) and z == mul(key, m), 'key/response mismatch')
    t2, t3 = mul(nonce, G), mul(nonce, m)
    transcript = challenge(public, m, z, t2, t3)
    c = hscalar(transcript, 1)
    s = (nonce-c*key) % Q
    return i2(c, 32)+i2(s, 32), {
        'seed': sha(frame(public)+frame(b'Seed-'+context(1))).hex(),
        'coefficient': coefficient, 'm': encode(m).hex(), 'n': encode(z).hex(),
        't2': encode(t2).hex(), 't3': encode(t3).hex(),
        'challenge_input': transcript.hex(), 'challenge_bytes': len(transcript)}

def verify(public, request, response, proof):
    try:
        octets(proof); require(len(proof) == 64, 'proof length')
        c, s = decode_scalar(proof[:32]), decode_scalar(proof[32:])
        pk, m, z, _ = composites(public, request, response)
        t2, t3 = add(mul(s, G), mul(c, pk)), add(mul(s, m), mul(c, z))
        return hscalar(challenge(public, m, z, t2, t3), 1) == c
    except ValueError:
        return False

def finalize_verified(msg, r, public, request, response, proof):
    require(blind(msg, r, 1) == request, 'local request mismatch')
    require(verify(public, request, response, proof), 'VerifyError')
    return finish(msg, r, response)


def main():
    dst = b'QUUX-V01-CS02-with-P256_XMD:SHA-256_SSWU_RO_'
    hashes = []
    expected = [('2c15230b26dbc6fc9a37051158c95b79656e17a1a920b11394ca91c44247d3e4','8a7a74985cc5c776cdfe4b1f19884970453912e9d31528c060be9ab5c43e8415'),
                ('0bb8b87485551aa43ed54f009230450b492fead5f1cc91658775dac4a3388a0f','5c41b3d0731a27a7b14bc0bf0ccded2d8751f83493404c84a88e71ffd424212e')]
    for msg, coords in zip((b'', b'abc'), expected):
        t, log = hash_curve(msg, dst)
        require(t == tuple(int(x, 16) for x in coords), 'RFC9380 point vector')
        hashes.append({'input': msg.hex(), 'dst': dst.hex(), 'xmd_96': xmd(msg, dst, 96).hex(),
                       'point': encode(t).hex(), 'u': [hex(x) for x in log['u']],
                       'map_points': [encode(x).hex() for x in log['points']]})
    seed, info = bytes.fromhex('a3'*32), b'test key'
    r = int('3338fa65ec36e0290022b48eb562889d89dbfa691d1cde91517fa222ed7ad364', 16)
    nonce = int('f9db001266677f62c095021db018cd8cbb55941d4073698ce45c405d1348b7b1', 16)
    vectors = []
    for mode in (0, 1):
        key, pk, ctr = derive(seed, info, mode)
        request = blind(b'\x00', r, mode)
        response = evaluate_blind(key, request)
        output = finish(b'\x00', r, response)
        require(output == evaluate(key, b'\x00', mode), 'unblind/direct equality')
        unblinded = encode(mul(pow(r, -1, Q), decode(response)))
        final_bytes = frame(b'\x00')+frame(unblinded)+b'Finalize'
        item = {'mode': mode, 'context': context(mode).hex(), 'key': hex(key),
                'key_bytes': i2(key,32).hex(), 'blind_bytes': i2(r,32).hex(),
                'public': encode(pk).hex(), 'request': request.hex(), 'response': response.hex(),
                'unblinded': unblinded.hex(), 'finalize_input': final_bytes.hex(),
                'output': output.hex(), 'derive_counter': ctr}
        if mode == 1:
            proof, log = prove(key, encode(pk), request, response, nonce)
            require(finalize_verified(b'\x00', r, encode(pk), request, response, proof) == output, 'VOPRF final')
            cc, ss = decode_scalar(proof[:32]), decode_scalar(proof[32:])
            pk1, mm, nn, _ = composites(encode(pk), request, response)
            restored = [encode(add(mul(ss,G),mul(cc,pk1))).hex(),
                        encode(add(mul(ss,mm),mul(cc,nn))).hex()]
            require(restored == [log['t2'],log['t3']], 'verification reconstruction log')
            item.update(proof=proof.hex(), proof_log=log, verification_points=restored)
        vectors.append(item)
    require(vectors[0]['output'] == 'a0b34de5fa4c5b6da07e72af73cc507cceeb48981b97b7285fc375345fe495dd', 'RFC9497 OPRF vector')
    require(vectors[1]['output'] == '0412e8f78b02c415ab3a288e228978376f99927767ff37c5718d420010a645a1', 'RFC9497 VOPRF vector')
    require(vectors[1]['proof'] == 'e7c2b3c5c954c035949f1f74e6bce2ed539a3be267d1481e9ddb178533df4c2664f69d065c604a4fd953e100b856ad83804eb3845189babfa5a702090d6fc5fa', 'RFC9497 proof vector')
    # Byte-exact checkpoints, not only final equality.
    require(vectors[0]['request'] == '03723a1e5c09b8b9c18d1dcbca29e8007e95f14f4732d9346d490ffc195110368d', 'OPRF request')
    require(vectors[0]['response'] == '030de02ffec47a1fd53efcdd1c6faf5bdc270912b8749e783c7ca75bb412958832', 'OPRF response')
    require(vectors[1]['public'] == '03e17e70604bcabe198882c0a1f27a92441e774224ed9c702e51dd17038b102462', 'VOPRF public')
    require(vectors[1]['request'] == '02dd05901038bb31a6fae01828fd8d0e49e35a486b5c5d4b4994013648c01277da', 'VOPRF request')
    require(vectors[1]['response'] == '0209f33cab60cf8fe69239b0afbcfcd261af4c1c5632624f2e9ba29b90ae83e4a2', 'VOPRF response')
    for mode, wanted in ((0, 'c748ca6dd327f0ce85f4ae3a8cd6d4d5390bbb804c9e12dcf94f853fece3dcce'), (1, '771e10dcd6bcd3664e23b8f2a710cfaaa8357747c4a8cbba03133967b5c24f18')):
        key, pk, _ = derive(seed, info, mode)
        msg = b'Z'*17; req = blind(msg, r, mode); resp = evaluate_blind(key, req)
        require(finish(msg, r, resp).hex() == wanted, 'RFC9497 second vector')
        if mode:
            pr, _ = prove(key, encode(pk), req, resp, nonce)
            require(pr.hex() == '2787d729c57e3d9512d3aa9e8708ad226bc48e0f1750b0767aaff73482c44b8d2873d74ec88aebd3504961acea16790a05c542d9fbff4fe269a77510db00abab', 'second proof')
    key, pk, _ = derive(seed, info, 1); public = encode(pk); msg = b'\x00'
    first = blind(msg, r, 1); second = blind(msg, r+1, 1)
    v1, v2 = evaluate_blind(key, first), evaluate_blind(key, second)
    require(first != second and v1 != v2 and finish(msg, r, v1) == finish(msg, r+1, v2), 'fresh blinding')
    wrongkey = (key+1) % Q; wrongpublic = encode(mul(wrongkey, G))
    wrongresponse = evaluate_blind(wrongkey, first)
    wrongproof, _ = prove(wrongkey, wrongpublic, first, wrongresponse, nonce)
    require(not verify(public, first, wrongresponse, wrongproof), 'fixed key mismatch')
    require(verify(wrongpublic, first, wrongresponse, wrongproof), 'changed identity changes claim')
    proof1, _ = prove(key, public, first, v1, nonce)
    proof2, _ = prove(key, public, second, v2, nonce)
    c1, s1 = decode_scalar(proof1[:32]), decode_scalar(proof1[32:])
    c2, s2 = decode_scalar(proof2[:32]), decode_scalar(proof2[32:])
    require(c1 != c2, 'nonce leakage challenge collision')
    recovered = (s1-s2)*pow((c2-c1) % Q, -1, Q) % Q
    require(recovered == key, 'nonce reuse extraction')
    # Wrong H1(x)=h(x)G still passes the DLEQ relation, but one query suffices.
    h0, h1 = hscalar(b'first', 1), hscalar(b'unqueried', 1)
    require(h0 != 0 and h1 != 0, 'public attack samples')
    badrequest = encode(mul(r, mul(h0, G))); badresponse = evaluate_blind(key, badrequest)
    n0 = mul(pow(r, -1, Q), decode(badresponse))
    recovered_public = mul(pow(h0, -1, Q), n0)
    forged_n = mul(h1, recovered_public)
    require(forged_n == mul(key, mul(h1, G)), 'one-query known-log attack')
    badproof, _ = prove(key, public, badrequest, badresponse, nonce)
    require(verify(public, badrequest, badresponse, badproof), 'DLEQ cannot fix wrong map')
    forged_output = sha(frame(b'unqueried')+frame(encode(forged_n))+b'Finalize')
    root = pow(pow(10, -1, P), (P+1)//4, P)
    exceptions = [0, root, (-root) % P]
    require(root*root % P == pow(10, -1, P), 'exception root')
    for u in exceptions:
        _, log = swu(u); require(log['exception'], 'all exceptional denominators')
    for u in range(512):
        t, _ = swu(u); point(t); require(t[1] % 2 == u % 2, 'SSWU sign')
        require(decode(encode(t)) == t, 'canonical round trip')
    altered, altered_log = hash_curve(b'', dst+b'-other')
    require(altered != hash_curve(b'', dst)[0] and altered != hash_curve(b'', dst)[1]['points'][0], 'changed DST')
    noncurve = next(x for x in range(100) if pow((x*x*x+A*x+B) % P, (P-1)//2, P) == P-1)
    rejected = []
    def reject(label, fn):
        try: fn()
        except ValueError: rejected.append(label); return
        raise ValueError('invalid input accepted: '+label)
    for label, fn in [
        ('zero blind', lambda: blind(msg, 0, 1)), ('bool scalar', lambda: blind(msg, True, 1)),
        ('zero key', lambda: evaluate_blind(0, first)), ('zero nonce', lambda: prove(key, public, first, v1, 0)),
        ('long input', lambda: blind(bytes(65535), r, 1)), ('empty DST', lambda: xmd(b'', b'', 96)),
        ('oversize DST', lambda: xmd(b'', bytes(256), 96)), ('oversize expansion', lambda: xmd(b'', dst, 8161)),
        ('short point', lambda: decode(first[:-1])), ('identity encoding', lambda: decode(b'\x00')),
        ('x equals p', lambda: decode(b'\x02'+i2(P,32))), ('not on curve', lambda: decode(b'\x02'+i2(noncurve,32))),
        ('scalar equals q', lambda: decode_scalar(i2(Q,32))),
        ('zero composite', lambda: composites(public, first, v1, lambda *_: 0)),
        ('changed local request', lambda: finalize_verified(msg,r,public,second,v2,proof2)),
        ('wrong fixed key', lambda: finalize_verified(msg,r,public,first,wrongresponse,wrongproof))]: reject(label,fn)
    require(not verify(public,first,v1,bytes(64)), 'identity reconstructed commitments')
    require(not verify(public,first,v1,proof1[:-1]), 'short proof')
    require(not verify(public,first,v1,i2(Q,32)+proof1[32:]), 'noncanonical proof scalar')
    require(not verify(public,blind(msg,r,0),v1,proof1), 'wrong mode request')
    # This finite additive group checks the uniform-bijection algebra only.
    views = [sorted((rr*h % 11, (7*(rr*h % 11)+coin) % 11, coin)
                    for rr in range(1,11) for coin in range(3)) for h in (2,7)]
    require(views[0] == views[1], 'joint server-view enumeration')
    require(len(blind(bytes(65534),r,1)) == 33, 'largest accepted input')
    print(json.dumps({'hash_vectors': hashes, 'protocol_vectors': vectors,
          'changed_dst': {'dst': (dst+b'-other').hex(),
                          'xmd_96': xmd(b'',dst+b'-other',96).hex(),
                          'u': [hex(x) for x in altered_log['u']],
                          'map_points': [encode(x).hex() for x in altered_log['points']],
                          'point': encode(altered).hex()},
          'migration': {'fresh_request':second.hex(),'fresh_response':v2.hex(),
                        'same_output':finish(msg,r+1,v2).hex(),'wrong_key_rejected':True,
                        'nonce_recovered_key':i2(recovered,32).hex(),
                        'nonce_reuse_proofs':[proof1.hex(),proof2.hex()],
                        'changed_key_public':wrongpublic.hex(),
                        'changed_key_response':wrongresponse.hex(),
                        'changed_key_proof':wrongproof.hex(),
                        'one_query_recovered_public':encode(recovered_public).hex(),
                        'unqueried_wrong_map_output':forged_output.hex(),
                        'wrong_map_dleq_still_valid':True,
                        'wrong_map_request':badrequest.hex(),
                        'wrong_map_response':badresponse.hex(),
                        'wrong_map_proof':badproof.hex(),
                        'wrong_map_first_scalar':i2(h0,32).hex(),
                        'wrong_map_unqueried_scalar':i2(h1,32).hex(),
                        'wrong_map_unblinded_first':encode(n0).hex(),
                        'wrong_map_unqueried_point':encode(forged_n).hex()},
          'self_checks': {'swu_regular_inputs':512, 'exceptional_inputs':[hex(x) for x in exceptions],
                          'joint_view_records_per_input':30, 'explicit_rejections':rejected,
                          'verification_rejections':4}}, indent=2))

if __name__ == '__main__':
    main()
