import argparse
import os
import random
import string
import sys

sys.path.insert(0, os.path.dirname(__file__))

from siren_client import SirenClient
from hnp_solve import build_samples, solve_hnp, verify_privkey


def rand_msg(n=16):
    return "".join(random.choices(string.ascii_letters + string.digits, k=n))


def gather(host, port, use_tls, num_sigs):
    c = SirenClient(host, port, use_tls=use_tls)
    pk = c.pubkey()
    Qx = int(pk["Qx"], 16)
    Qy = int(pk["Qy"], 16)
    N = int(pk["n"], 16)
    song_id = pk["song_id"]
    pitch_bits = pk["pitch_bits"]
    priv_msg = pk["priv_msg"]
    print(f"[+] song_id={song_id} pitch_bits={pitch_bits} priv_msg={priv_msg!r}")

    sigs = []
    seen = set()
    while len(sigs) < num_sigs:
        m = rand_msg()
        if m == priv_msg or m in seen:
            continue
        seen.add(m)
        try:
            r, s = c.sign(m)
        except RuntimeError as e:
            print(f"[!] skip {m!r}: {e}")
            continue
        sigs.append((m, r, s))
        if len(sigs) % 20 == 0:
            print(f"[+] collected {len(sigs)}/{num_sigs}")

    return c, Qx, Qy, N, song_id, pitch_bits, priv_msg, sigs


def try_solve(sigs, N, song_id, pitch_bits, Qx, Qy, block_size):
    suffix_bits = N.bit_length() - pitch_bits
    samples = build_samples(sigs, song_id, N, pitch_bits, suffix_bits)
    candidates = solve_hnp(samples, N, suffix_bits, block_size=block_size)
    for d in candidates:
        if d and verify_privkey(d, Qx, Qy, N):
            return d
    return None


def forge_and_unlock(client, d, priv_msg, N):
    from ecdsa import SECP256k1
    import hashlib

    G = SECP256k1.generator
    z = int.from_bytes(hashlib.sha256(priv_msg.encode()).digest(), "big") % N
    k = random.randrange(1, N)
    R = k * G
    r = R.x() % N
    s = (pow(k, -1, N) * (z + r * d)) % N
    print(f"[+] forged signature r={hex(r)} s={hex(s)}")
    resp = client.unlock(r, s)
    print(f"[+] unlock response: {resp}")
    return resp


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--host", required=True)
    ap.add_argument("--port", type=int, required=True)
    ap.add_argument("--tls", action="store_true")
    ap.add_argument("--num-sigs", type=int, default=150)
    ap.add_argument("--block-size", type=int, default=20)
    ap.add_argument("--known-d", type=int, default=None,
                     help="ground truth D for local validation runs only")
    args = ap.parse_args()

    c, Qx, Qy, N, song_id, pitch_bits, priv_msg, sigs = gather(
        args.host, args.port, args.tls, args.num_sigs
    )

    print(f"[+] solving with {len(sigs)} signatures, BKZ block_size={args.block_size} ...")
    d = try_solve(sigs, N, song_id, pitch_bits, Qx, Qy, args.block_size)

    if d is None:
        print("[-] FAILED to recover private key with current sample count/block size.")
        if args.known_d is not None:
            print(f"    (ground truth D was {hex(args.known_d)})")
        sys.exit(1)

    print(f"[+] RECOVERED PRIVATE KEY d = {hex(d)}")
    if args.known_d is not None:
        ok = (d == args.known_d)
        print(f"[+] matches ground truth D: {ok}")
        if not ok:
            sys.exit(1)

    forge_and_unlock(c, d, priv_msg, N)


if __name__ == "__main__":
    main()
