#!/usr/bin/env python3
"""End-to-end SIGN_TX driver for the PSBT-style state machine (INS 0x06).

Streams an unsigned Handshake tx to the Ledger app as structured sub-ops
(BEGIN / ADD_INPUT / ADD_OUTPUT / REVIEW / SIGN_INPUT), then assembles the
signed tx with one witness per input.

Usage:
  ./scripts/sign-tx \
    --path "44'/5353'/0'/0/0" \
    --input 865d9e8011815fbfeb168bff83767f039cd6ac71bd75893d74fb89a3d932322d:0:10000000 \
    --output hs1qdydhjsla4ahvy5tn6w44ax7y62mpsd5ny4jlej:9999000

Amounts are in dollarydoos. 1 HNS = 1,000,000 dollarydoos.
Exit code 0 = signed; nonzero = error or device rejection.
"""
import argparse
import sys

import bech32  # type: ignore[import-untyped]
from ledgerwallet.transport import enumerate_devices

CLA = 0xE0
INS_GET_PUBKEY = 0x05
INS_SIGN_TX = 0x06

P1_BEGIN = 0x00
P1_ADD_INPUT = 0x01
P1_ADD_OUTPUT = 0x02
P1_REVIEW = 0x03
P1_SIGN_INPUT = 0x04

P2_LAST = 0x00
P2_MORE = 0x80

SCRIPT_KIND_P2WPKH = 0x00
MAX_DATA = 240

HRP = {"main": "hs", "testnet": "ts", "regtest": "rs", "simnet": "ss"}


def open_device():
    devs = enumerate_devices()
    if not devs:
        print("no Ledger device found", file=sys.stderr)
        sys.exit(1)
    t = devs[0]
    t.open()
    return t


def parse_path(s: str) -> bytes:
    """`count(1) || components[count](BE u32, hardened bit set on quoted parts)`."""
    parts = [p for p in s.lstrip("m/").split("/") if p]
    out = bytearray([len(parts)])
    for p in parts:
        n = int(p.rstrip("'"))
        if p.endswith("'"):
            n |= 0x80000000
        out.extend(n.to_bytes(4, "big"))
    return bytes(out)


def encode_varint(n: int) -> bytes:
    if n < 0xFD:
        return bytes([n])
    if n <= 0xFFFF:
        return b"\xfd" + n.to_bytes(2, "little")
    if n <= 0xFFFFFFFF:
        return b"\xfe" + n.to_bytes(4, "little")
    return b"\xff" + n.to_bytes(8, "little")


def apdu(ins: int, p1: int, p2: int, data: bytes) -> bytes:
    if len(data) > 255:
        raise RuntimeError(f"APDU data too long: {len(data)}")
    return bytes([CLA, ins, p1, p2, len(data)]) + data


def exchange(t, apdu_bytes: bytes, timeout_ms: int = 120000) -> bytes:
    """Same convention as scripts/send-apdu: transport returns data || sw."""
    print(f"-> {apdu_bytes.hex().upper()}", file=sys.stderr)
    response = t.exchange(apdu_bytes, timeout=timeout_ms)
    if len(response) < 2:
        raise RuntimeError(f"short response: {response.hex()}")
    sw = int.from_bytes(response[-2:], "big")
    data = bytes(response[:-2])
    print(f"<- sw={sw:04x} data={data.hex().upper()} ({len(data)} bytes)", file=sys.stderr)
    if sw != 0x9000:
        raise RuntimeError(f"APDU failed: sw={sw:04x}")
    return data


def get_compressed_pubkey(t, path_bytes: bytes) -> bytes:
    """GET_PUBKEY (no display): pk_len(1)||pk||cc_len(1)||cc||addr_len(1)||addr."""
    data = exchange(t, apdu(INS_GET_PUBKEY, 0, 0, path_bytes))
    return data[1 : 1 + data[0]]


def parse_input(s: str):
    """`<txid_hex>:<vout>:<value_dollarydoos>`."""
    txid_str, vout_str, value_str = s.split(":")
    txid = bytes.fromhex(txid_str)
    if len(txid) != 32:
        raise ValueError(f"input txid must be 32 bytes, got {len(txid)}")
    return txid, int(vout_str), int(value_str)


def parse_output(s: str, hrp: str) -> bytes:
    """`<bech32_addr>:<value>` -> value(8 LE)||ver(1)||len(1)||hash||cov(0x00)||varint(0)."""
    addr, amount_str = s.rsplit(":", 1)
    got_hrp, data = bech32.bech32_decode(addr)
    if got_hrp != hrp:
        raise ValueError(f"address HRP {got_hrp!r} != expected {hrp!r}")
    if not data:
        raise ValueError(f"invalid bech32 address: {addr}")
    witver = data[0]
    program = bech32.convertbits(data[1:], 5, 8, False)
    if program is None:
        raise ValueError(f"invalid bech32 program: {addr}")
    return (
        int(amount_str).to_bytes(8, "little")
        + bytes([witver, len(program)])
        + bytes(program)
        + b"\x00"           # covenant kind NONE
        + encode_varint(0)  # covenant item count
    )


def main() -> int:
    ap = argparse.ArgumentParser(
        description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
    )
    ap.add_argument("--path", required=True, help="BIP-32 path, e.g. 44'/5353'/0'/0/0")
    ap.add_argument("--input", action="append", required=True,
                    help="`<txid>:<vout>:<value>` (repeat for multiple inputs)")
    ap.add_argument("--output", action="append", default=[],
                    help="`<bech32_addr>:<value>` (repeat for multiple outputs)")
    ap.add_argument("--tx-version", type=int, default=0)
    ap.add_argument("--locktime", type=int, default=0)
    ap.add_argument("--sequence", type=lambda s: int(s, 0), default=0xFFFFFFFF)
    ap.add_argument("--sighash-type", type=lambda s: int(s, 0), default=0x01)
    ap.add_argument("--network", choices=list(HRP), default="main")
    args = ap.parse_args()

    hrp = HRP[args.network]
    path_bytes = parse_path(args.path)
    inputs = [parse_input(i) for i in args.input]
    outputs = [parse_output(o, hrp) for o in args.output]

    t = open_device()
    try:
        pubkey = get_compressed_pubkey(t, path_bytes)
        print(f"# compressed pubkey: {pubkey.hex().upper()}", file=sys.stderr)

        # BEGIN
        exchange(t, apdu(INS_SIGN_TX, P1_BEGIN, 0,
                         args.tx_version.to_bytes(4, "little")
                         + encode_varint(len(inputs))
                         + encode_varint(len(outputs))
                         + args.locktime.to_bytes(4, "little")))

        # ADD_INPUT (one APDU each)
        for txid, vout, value in inputs:
            exchange(t, apdu(INS_SIGN_TX, P1_ADD_INPUT, 0,
                             txid
                             + vout.to_bytes(4, "little")
                             + args.sequence.to_bytes(4, "little")
                             + value.to_bytes(8, "little")
                             + bytes([args.sighash_type])
                             + path_bytes
                             + bytes([SCRIPT_KIND_P2WPKH])))

        # ADD_OUTPUT (chunked; P2 = MORE until the last chunk of each output)
        for payload in outputs:
            chunks = [payload[i:i + MAX_DATA] for i in range(0, len(payload), MAX_DATA)]
            for i, c in enumerate(chunks):
                p2 = P2_LAST if i == len(chunks) - 1 else P2_MORE
                exchange(t, apdu(INS_SIGN_TX, P1_ADD_OUTPUT, p2, c))

        # REVIEW (blocks on the user)
        exchange(t, apdu(INS_SIGN_TX, P1_REVIEW, 0, b""))

        # SIGN_INPUT per input -> sig(64)||sighash(1)||pk_len(1)||pk
        witnesses = []
        for idx in range(len(inputs)):
            r = exchange(t, apdu(INS_SIGN_TX, P1_SIGN_INPUT, 0, idx.to_bytes(4, "big")))
            if len(r) < 66 or len(r) != 66 + r[65]:
                raise RuntimeError(f"bad SIGN_INPUT response: {r.hex()}")
            sig_with_sighash = r[:65]
            pk = r[66:]
            print(f"# input {idx} sig: {sig_with_sighash.hex().upper()}", file=sys.stderr)
            witnesses.append(
                encode_varint(2)
                + encode_varint(len(sig_with_sighash)) + sig_with_sighash
                + encode_varint(len(pk)) + pk
            )

        # version || inputs || outputs || locktime || witnesses (shd Transaction.write)
        tx = bytearray(args.tx_version.to_bytes(4, "little"))
        tx += encode_varint(len(inputs))
        for txid, vout, _ in inputs:
            tx += txid + vout.to_bytes(4, "little") + args.sequence.to_bytes(4, "little")
        tx += encode_varint(len(outputs))
        for payload in outputs:
            tx += payload
        tx += args.locktime.to_bytes(4, "little")
        for w in witnesses:
            tx += w

        print(tx.hex())
        return 0
    finally:
        t.close()


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