#!/usr/bin/env python3
"""Independent Waxseal verifier. Standard library only.

Written from docs/COMMITMENT_FORMAT.md, not translated from the TypeScript SDK, so that a
disagreement between the two exposes an ambiguity in the specification. Use it to:

  selftest                       run every vector in conformance/vectors.json, plus a
                                 deliberately wrong implementation that must fail
  opening --bid N --salt HEX --bidder-id HEX [--commitment HEX]
  opening --file waxseal-opening.json
                                 recompute a bid commitment from an opening (a bid and salt
                                 a bidder chose to disclose) and compare it
  bidder-id --secret HEX --tender HEX
  metadata FILE [--expect HEX]   hash tender metadata (JSON file) as the contract does
  price --rule lowest|highest --min N --max N --step N --round R
  round --rule lowest|highest --min N --max N --step N --bid N
                                 the only clock round in which a bid can be claimed
  anchor --statement FILE --tsr FILE [--roots tsa-roots.pem]
                                 check an independent RFC 3161 time stamp: that it covers
                                 exactly this statement, when it was made, and (with
                                 OpenSSL installed) the authority's signature

Exit status 0 when everything checked matches, 1 otherwise.
"""
import argparse
import hashlib
import json
import os
import shutil
import subprocess
import sys

# --- Section 1: notation -------------------------------------------------------------

def H(data: bytes) -> bytes:
    return hashlib.sha256(data).digest()


def tag(s: str) -> bytes:
    raw = s.encode("utf-8")
    if len(raw) > 32:
        raise ValueError("tag longer than 32 bytes")
    return raw + b"\x00" * (32 - len(raw))


def le32(n: int) -> bytes:
    if n < 0 or n >= 1 << 256:
        raise ValueError("value out of range for 32 bytes")
    return n.to_bytes(32, "little")


def b32(hex_str: str) -> bytes:
    raw = bytes.fromhex(hex_str)
    if len(raw) != 32:
        raise ValueError("expected 32 bytes (64 hex characters)")
    return raw


# --- Section 2: identifiers ----------------------------------------------------------

def bidder_id(secret: bytes, tender: bytes) -> bytes:
    return H(tag("waxseal:bidder-id:v1") + secret + tender)


def owner_id(secret: bytes) -> bytes:
    return H(tag("waxseal:owner-id:v1") + secret)


def eligibility_key(secret: bytes, organiser_owner_id: bytes) -> bytes:
    return H(tag("waxseal:eligibility:v2") + secret + organiser_owner_id)


# --- Section 3: bid commitment -------------------------------------------------------

UINT64_MAX = (1 << 64) - 1


def commitment(bid: int, salt: bytes, bidder: bytes, encode_bid=le32) -> bytes:
    if bid < 0 or bid > UINT64_MAX:
        raise ValueError("bid must fit in 64 bits")
    if len(salt) != 32 or len(bidder) != 32:
        raise ValueError("salt and bidder id must be 32 bytes")
    return H(tag("waxseal:commitment:v1") + encode_bid(bid) + salt + bidder)


# --- Section 4: canonical metadata ---------------------------------------------------

_SHORT = {'"': '\\"', "\\": "\\\\", "\b": "\\b", "\f": "\\f", "\n": "\\n", "\r": "\\r", "\t": "\\t"}


def canonical_string(s: str) -> str:
    out = ['"']
    for ch in s:
        c = ord(ch)
        if ch in _SHORT:
            out.append(_SHORT[ch])
        elif c < 0x20 or 0xD800 <= c <= 0xDFFF:
            # Python joins valid surrogate pairs into one code point when parsing JSON, so
            # any surrogate seen here is unpaired.
            out.append("\\u%04x" % c)
        else:
            out.append(ch)
    out.append('"')
    return "".join(out)


def canonical(value) -> str:
    if value is None:
        return "null"
    if value is True:
        return "true"
    if value is False:
        return "false"
    if isinstance(value, int):
        return str(value)
    if isinstance(value, str):
        return canonical_string(value)
    if isinstance(value, list):
        return "[" + ",".join(canonical(v) for v in value) + "]"
    if isinstance(value, dict):
        keys = sorted(value.keys(), key=lambda k: k.encode("ascii"))
        return "{" + ",".join(canonical_string(k) + ":" + canonical(value[k]) for k in keys) + "}"
    raise TypeError("unsupported metadata value %r" % (value,))


def metadata_hash(metadata) -> bytes:
    return H(canonical(metadata).encode("utf-8", "surrogatepass"))


# --- Section 5: clock mode -----------------------------------------------------------

def price(rule: str, lo: int, hi: int, step: int, r: int) -> int:
    moved = min(r * step, hi - lo)
    return lo + moved if rule == "lowest" else hi - moved


def reaches(rule: str, bid: int, p: int) -> bool:
    return bid <= p if rule == "lowest" else bid >= p


def rounds_for(lo: int, hi: int, step: int) -> int:
    return (hi - lo + step - 1) // step


def qualifying_round(rule: str, lo: int, hi: int, step: int, bid: int) -> int:
    for r in range(rounds_for(lo, hi, step) + 1):
        if reaches(rule, bid, price(rule, lo, hi, step, r)):
            return r
    raise ValueError("bid outside the clock range")


# --- Self test -----------------------------------------------------------------------

def run_vectors(v, encode_bid=le32, canon=canonical):
    failures = []

    def expect(label, got, want):
        if got != want:
            failures.append("%s: expected %s, computed %s" % (label, want, got))

    for c in v["identifiers"]:
        s, t = b32(c["bidderSecret"]), b32(c["tender"])
        expect("bidderId", bidder_id(s, t).hex(), c["bidderId"])
        expect("ownerId", owner_id(s).hex(), c["ownerId"])
        expect("eligibilityKey", eligibility_key(s, b32(c["organiserOwnerId"])).hex(), c["eligibilityKey"])
    for c in v["clockCommitments"]:
        got = commitment(int(c["bid"]), b32(c["salt"]), b32(c["bidderId"]), encode_bid).hex()
        expect("clock commitment bid=%s" % c["bid"], got, c["commitment"])
    for c in v["committeeCommitments"]:
        expect("committee salt encoding", le32(int(c["saltField"])).hex(), c["salt"])
        got = commitment(int(c["bid"]), b32(c["salt"]), b32(c["bidderId"]), encode_bid).hex()
        expect("committee commitment bid=%s" % c["bid"], got, c["commitment"])
    for c in v["metadata"]:
        text = canon(c["metadata"])
        expect("metadata canonical %s" % c["name"], text, c["canonical"])
        expect("metadata hash %s" % c["name"], H(text.encode("utf-8", "surrogatepass")).hex(), c["hash"])
    for c in v["clock"]:
        rule, lo, hi, step = c["rule"], int(c["minBid"]), int(c["maxBid"]), int(c["step"])
        expect("rounds %s" % rule, str(rounds_for(lo, hi, step)), c["rounds"])
        for p in c["prices"]:
            expect("price %s r=%s" % (rule, p["round"]), str(price(rule, lo, hi, step, int(p["round"]))), p["price"])
        for q in c["qualifying"]:
            expect("qualifying %s bid=%s" % (rule, q["bid"]), str(qualifying_round(rule, lo, hi, step, int(q["bid"]))), q["round"])
    return failures


def cmd_selftest(args):
    with open(args.vectors, encoding="utf-8") as f:
        v = json.load(f)
    failures = run_vectors(v)
    total = (len(v["identifiers"]) * 3 + len(v["clockCommitments"]) + len(v["committeeCommitments"]) * 2
             + len(v["metadata"]) * 2 + sum(1 + len(c["prices"]) + len(c["qualifying"]) for c in v["clock"]))
    if failures:
        print("FAIL: %d of %d checks" % (len(failures), total))
        for line in failures[:20]:
            print("  " + line)
        return 1
    print("PASS: %d checks from %s" % (total, os.path.relpath(args.vectors)))

    # Negative controls: plausible bugs the vectors must catch.
    wrong_endian = run_vectors(v, encode_bid=lambda n: n.to_bytes(32, "big"))
    plain_json = run_vectors(v, canon=lambda m: json.dumps(m, sort_keys=True, separators=(",", ":")))
    for name, fails in (("big-endian bid encoding", wrong_endian), ("json.dumps with ASCII escaping", plain_json)):
        if not fails:
            print("FAIL: the vectors did not catch a wrong implementation (%s)" % name)
            return 1
        print("PASS: wrong implementation rejected (%s): %d mismatches" % (name, len(fails)))
    return 0


def cmd_opening(args):
    if args.file:
        with open(args.file, encoding="utf-8") as f:
            o = json.load(f)
        if o.get("format") != "waxseal-bid-opening/v1":
            print("not a waxseal-bid-opening/v1 file")
            return 1
        args.bid, args.salt, args.bidder_id, args.commitment = o["bid"], o["salt"], o["bidderId"], o["commitment"]
        print("tender %s (%s), bid %s" % (o["tender"], o["network"], o["bid"]))
    if not (args.bid and args.salt and args.bidder_id):
        print("give --file, or --bid, --salt and --bidder-id")
        return 2
    got = commitment(int(args.bid), b32(args.salt), b32(args.bidder_id)).hex()
    print(got)
    if args.commitment:
        ok = got == args.commitment.lower()
        print("MATCH: this opening is the committed bid" if ok else "MISMATCH: this opening does not produce the commitment")
        return 0 if ok else 1
    return 0


def cmd_bidder_id(args):
    print(bidder_id(b32(args.secret), b32(args.tender)).hex())
    return 0


def cmd_metadata(args):
    with open(args.file, encoding="utf-8") as f:
        m = json.load(f)
    got = metadata_hash(m).hex()
    print(got)
    if args.expect:
        ok = got == args.expect.lower()
        print("MATCH: metadata is the one fixed in the contract" if ok else "MISMATCH: metadata differs from the contract")
        return 0 if ok else 1
    return 0


def cmd_price(args):
    print(price(args.rule, args.min, args.max, args.step, args.round))
    return 0


def cmd_round(args):
    print(qualifying_round(args.rule, args.min, args.max, args.step, args.bid))
    return 0


# --- Section 8: independent time stamps (RFC 3161) ------------------------------------
# The standard library cannot check RSA signatures, so this reads the time-stamp token (DER)
# only far enough to compare what was time-stamped with the statement and to show the time.
# The authority's signature and certificate chain are checked by OpenSSL when installed.

SHA256_OID = bytes.fromhex("608648016503040201")  # 2.16.840.1.101.3.4.2.1


def der_read(buf: bytes, i: int = 0):
    tag, length, j = buf[i], buf[i + 1], i + 2
    if length & 0x80:
        n = length & 0x7F
        if n == 0 or n > 4:
            raise ValueError("not DER (length)")
        length, j = int.from_bytes(buf[j:j + n], "big"), j + n
    if j + length > len(buf):
        raise ValueError("not DER (truncated)")
    return tag, buf[j:j + length], j + length


def der_children(value: bytes):
    out, i = [], 0
    while i < len(value):
        tag, v, i = der_read(value, i)
        out.append((tag, v))
    return out


def tst_info(tsr: bytes):
    """messageImprint, serial and genTime from a DER TimeStampResp."""
    _, resp, _ = der_read(tsr)
    status, token = der_children(resp)[:2]
    if int.from_bytes(der_children(status[1])[0][1], "big") not in (0, 1):
        raise ValueError("the authority refused the request")
    content = der_children(token[1])[1]            # ContentInfo: contentType, [0] content
    signed_data = der_children(der_children(content[1])[0][1])
    e_wrap = der_children(signed_data[2][1])[1]    # encapContentInfo: eContentType, [0] eContent
    _, tst_value, _ = der_read(der_children(e_wrap[1])[0][1])  # OCTET STRING holding TSTInfo
    tst = der_children(tst_value)
    imprint = der_children(tst[2][1])
    return {
        "hash_oid": der_children(imprint[0][1])[0][1],
        "digest": imprint[1][1],
        "serial": int.from_bytes(tst[3][1], "big"),
        "gen_time": tst[4][1].decode("ascii"),
    }


def cmd_anchor(args):
    here = os.path.dirname(os.path.abspath(__file__))
    with open(args.statement, "rb") as f:
        statement = f.read()
    with open(args.tsr, "rb") as f:
        tsr = f.read()
    ok = True
    s = json.loads(statement.decode("utf-8"))
    if s.get("format") != "waxseal-anchor/v1":
        print("not a waxseal-anchor/v1 statement")
        return 1
    if canonical(s).encode("utf-8") != statement:
        print("statement is NOT in canonical form")
        ok = False
    print("tender %s (%s), state after tx %s in block %d: phase %s, %d sealed bid(s)%s" % (
        s["tender"], s["network"], s["tx"], s["block"]["height"], s["phase"], len(s["commitments"]),
        ", result " + canonical(s["result"]) if s.get("result") else ""))
    info = tst_info(tsr)
    digest = H(statement)
    if info["hash_oid"] != SHA256_OID or info["digest"] != digest:
        print("the time stamp does NOT cover this statement")
        ok = False
    else:
        print("the time stamp covers this statement (SHA-256 %s)" % digest.hex())
    t = info["gen_time"]
    print("time asserted by the authority: %s-%s-%sT%s:%s:%s%s UTC (serial %x)" % (t[0:4], t[4:6], t[6:8], t[8:10], t[10:12], t[12:14], t[14:-1], info["serial"]))
    roots = args.roots or os.path.join(here, "tsa-roots.pem")
    cmd = ["openssl", "ts", "-verify", "-data", args.statement, "-in", args.tsr, "-CAfile", roots]
    if shutil.which("openssl") and os.path.exists(roots):
        r = subprocess.run(cmd, capture_output=True, text=True)
        signed = "Verification: OK" in r.stdout
        print("authority signature and certificate chain (OpenSSL): %s" % ("OK" if signed else "FAILED"))
        ok = ok and signed
    else:
        print("authority signature not checked here; run:\n  " + " ".join(cmd))
    print("MATCH" if ok else "MISMATCH")
    return 0 if ok else 1


def main(argv=None):
    here = os.path.dirname(os.path.abspath(__file__))
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    sub = ap.add_subparsers(dest="cmd", required=True)
    p = sub.add_parser("selftest")
    p.add_argument("--vectors", default=os.path.join(here, "vectors.json"))
    p.set_defaults(fn=cmd_selftest)
    p = sub.add_parser("opening")
    p.add_argument("--file")
    p.add_argument("--bid")
    p.add_argument("--salt")
    p.add_argument("--bidder-id")
    p.add_argument("--commitment")
    p.set_defaults(fn=cmd_opening)
    p = sub.add_parser("bidder-id")
    p.add_argument("--secret", required=True)
    p.add_argument("--tender", required=True)
    p.set_defaults(fn=cmd_bidder_id)
    p = sub.add_parser("metadata")
    p.add_argument("file")
    p.add_argument("--expect")
    p.set_defaults(fn=cmd_metadata)
    p = sub.add_parser("anchor")
    p.add_argument("--statement", required=True)
    p.add_argument("--tsr", required=True)
    p.add_argument("--roots")
    p.set_defaults(fn=cmd_anchor)
    for name, fn, extra in (("price", cmd_price, "round"), ("round", cmd_round, "bid")):
        p = sub.add_parser(name)
        p.add_argument("--rule", choices=["lowest", "highest"], required=True)
        p.add_argument("--min", type=int, required=True)
        p.add_argument("--max", type=int, required=True)
        p.add_argument("--step", type=int, required=True)
        p.add_argument("--" + extra, type=int, required=True)
        p.set_defaults(fn=fn)
    args = ap.parse_args(argv)
    return args.fn(args)


if __name__ == "__main__":
    sys.exit(main())
