#!/usr/bin/env python3
"""Build a Handshake unsigned base transaction (no witnesses) from inputs and outputs.

Mirrors the wire format from shd `Sources/Protocol/Transaction.swift`:
  version(4 LE) || varint(in_count) || inputs[40B each] || varint(out_count) || outputs || locktime(4 LE)

For each output:
  value(8 LE) || address(version(1)+hashLen(1)+hash) || covenant(type(1)+varint(0))

Defaults: version=0, sequence=0xFFFFFFFF, locktime=0, covenant=.none.

Usage:
  ./scripts/build-unsigned-tx \\
    --input 865d9e8011815fbfeb168bff83767f039cd6ac71bd75893d74fb89a3d932322d:0 \\
    --output hs1qdydhjsla4ahvy5tn6w44ax7y62mpsd5nk3qnvw:9999000

(Amounts are in dollarydoos. 1 HNS = 1,000,000 dollarydoos.)
"""
import argparse
import sys

import bech32  # type: ignore[import-untyped]


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 parse_input(s: str) -> tuple[bytes, int]:
    """`<txid_hex>:<vout>`: txid is natural byte order (matches shd .hex)."""
    txid_str, vout_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)


def parse_output(s: str, network_hrp: str) -> tuple[int, int, bytes]:
    """`<bech32_addr>:<amount_in_doos>` → (amount, witness_version, hash_bytes)."""
    addr, amount_str = s.split(":")
    hrp, data = bech32.bech32_decode(addr)
    if hrp != network_hrp:
        raise ValueError(f"address HRP {hrp!r} != expected {network_hrp!r}")
    if data is None or 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), witver, bytes(program)


def main() -> int:
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--input", action="append", required=True,
                    help="`<txid>:<vout>` (repeat for multiple inputs)")
    ap.add_argument("--output", action="append", required=True,
                    help="`<bech32_addr>:<amount_doos>` (repeat for multiple outputs)")
    ap.add_argument("--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("--network", choices=["main", "testnet", "regtest", "simnet"],
                    default="main")
    args = ap.parse_args()

    hrp_map = {"main": "hs", "testnet": "ts", "regtest": "rs", "simnet": "ss"}
    hrp = hrp_map[args.network]

    out = bytearray()
    out += args.version.to_bytes(4, "little")

    inputs = [parse_input(i) for i in args.input]
    out += encode_varint(len(inputs))
    for txid, vout in inputs:
        out += txid
        out += vout.to_bytes(4, "little")
        out += args.sequence.to_bytes(4, "little")

    outputs = [parse_output(o, hrp) for o in args.output]
    out += encode_varint(len(outputs))
    for amount, witver, program in outputs:
        out += amount.to_bytes(8, "little")
        # Address: version(1) + hashLen(1) + hash
        out += bytes([witver, len(program)]) + program
        # Covenant.none: type(0) + varint(0 items)
        out += b"\x00\x00"

    out += args.locktime.to_bytes(4, "little")

    print(out.hex())
    return 0


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