#!/usr/bin/env python3
"""Verify one Travis.TX merkle proof against the posted root.

Usage:
  python3 verify_travis_proof.py proof_0100030105.json
  python3 verify_travis_proof.py proof_0100030105.json --rpc
"""

from __future__ import annotations

import argparse
import hashlib
import json
import struct
import sys
import urllib.request
from pathlib import Path

FIPS = 48453
EXPECTED_ROOT = "3f875b9faf6f2410647790ca8d76ba391a3deeb691f1eba73d29b559982ac577"
PDA = "3sHiaLWMMbmkPzVYAFxYvrZaADqwtbLby3oNptzVh6oZ"
RPC = "https://api.devnet.solana.com"


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


def leaf_digest(apn: str, acres_q32: int, landuse: int) -> bytes:
    material = (
        FIPS.to_bytes(4, "big")
        + b"\x00"
        + apn.encode()
        + b"\x00"
        + acres_q32.to_bytes(8, "big")
        + b"\x00"
        + landuse.to_bytes(2, "big")
    )
    return sha256(material)


def walk(leaf: bytes, proof: list[dict]) -> bytes:
    acc = leaf
    for step in proof:
        sib = bytes.fromhex(step["sibling"])
        if step["side"] == "L":
            acc = sha256(sib + acc)
        elif step["side"] == "R":
            acc = sha256(acc + sib)
        else:
            raise SystemExit(f"bad side {step['side']}")
    return acc


def fetch_onchain_root(pda: str) -> str:
    body = json.dumps(
        {
            "jsonrpc": "2.0",
            "id": 1,
            "method": "getAccountInfo",
            "params": [pda, {"encoding": "base64"}],
        }
    ).encode()
    req = urllib.request.Request(
        RPC,
        data=body,
        headers={"Content-Type": "application/json"},
    )
    with urllib.request.urlopen(req, timeout=30) as resp:
        val = json.loads(resp.read().decode())["result"]["value"]
    if not val:
        raise SystemExit("PDA not found on devnet")
    import base64

    raw = base64.b64decode(val["data"][0])
    # 8 disc + fips4 + epoch4 + n8 + root32
    root = raw[8 + 4 + 4 + 8 : 8 + 4 + 4 + 8 + 32]
    fips, epoch = struct.unpack_from("<II", raw, 8)
    n = struct.unpack_from("<Q", raw, 16)[0]
    print(f"on-chain fips={fips} epoch={epoch} n={n}")
    return root.hex()


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("proof_json")
    ap.add_argument("--rpc", action="store_true", help="also read the Travis PDA on devnet")
    args = ap.parse_args()
    doc = json.loads(Path(args.proof_json).read_text())

    leaf = leaf_digest(doc["parcel_id"], int(doc["acres_q32"]), int(doc["landuse"]))
    if leaf.hex() != doc["leaf"]:
        raise SystemExit("leaf rebuild != proof.leaf")
    print("1 leaf rebuilt     ", leaf.hex())

    root = walk(leaf, doc["proof"])
    print("2 walked proof     ", root.hex())
    if root.hex() != EXPECTED_ROOT:
        raise SystemExit("proof does not land on posted Travis root")
    print("3 matches packet   ", EXPECTED_ROOT)

    if args.rpc:
        chain = fetch_onchain_root(doc.get("pda") or PDA)
        print("4 on-chain PDA     ", chain)
        if chain != root.hex():
            raise SystemExit("on-chain root mismatch")
        print("PASS  parcel", doc["parcel_id"], "is in Travis.TX epoch 1")
    else:
        print("PASS  local proof (add --rpc to check the PDA)")
        print("      parcel", doc["parcel_id"])


if __name__ == "__main__":
    main()
