#!/usr/bin/env python3
"""Read-only Solana SPL-USDC payment verification.

The verifier uses ordinary Solana JSON-RPC calls and never signs, submits, or
funds a transaction.  It is deliberately dependency-free so the same module
works as a small CLI or as an importable library.
"""

from __future__ import annotations

import argparse
import json
import time
import urllib.error
import urllib.request
from dataclasses import asdict, dataclass
from decimal import Decimal, InvalidOperation
from typing import Any, Callable


# Circle's canonical Solana mainnet USDC mint. Pass --mint explicitly for
# devnet or a different SPL token.
MAINNET_USDC_MINT = "EPjFWdd5AufqSSqeM2qN1xzybapC8G4wEGGkZwyTDt1v"
DEFAULT_RPC_URL = "https://api.mainnet-beta.solana.com"
DEFAULT_DECIMALS = 6


class VerificationError(RuntimeError):
    """Raised for malformed RPC responses or invalid verifier inputs."""


@dataclass(frozen=True)
class VerificationResult:
    verified: bool
    reason: str
    signature: str
    recipient: str
    mint: str
    expected_amount_units: int
    observed_amount_units: int
    confirmations: int
    min_confirmations: int
    slot: int | None
    confirmation_status: str | None
    attempts: int

    def as_dict(self) -> dict[str, Any]:
        return asdict(self)


def amount_to_units(value: str | int | Decimal, decimals: int = DEFAULT_DECIMALS) -> int:
    """Convert an exact decimal USDC amount into its integer base units."""
    if decimals < 0 or decimals > 18:
        raise ValueError("decimals must be between 0 and 18")
    try:
        amount = Decimal(str(value))
    except (InvalidOperation, ValueError) as exc:
        raise ValueError("amount must be a decimal number") from exc
    if not amount.is_finite() or amount < 0:
        raise ValueError("amount must be finite and non-negative")
    scale = Decimal(10) ** decimals
    units = amount * scale
    if units != units.to_integral_value():
        raise ValueError(f"amount has more than {decimals} decimal places")
    return int(units)


class SolanaRpcClient:
    """Tiny JSON-RPC client with an injectable opener for tests."""

    def __init__(self, rpc_url: str, opener: Callable[..., Any] | None = None, timeout: float = 15.0):
        if not rpc_url.startswith(("https://", "http://")):
            raise ValueError("rpc_url must use http:// or https://")
        self.rpc_url = rpc_url
        self.opener = opener or urllib.request.urlopen
        self.timeout = timeout
        self._request_id = 0

    def call(self, method: str, params: list[Any]) -> Any:
        self._request_id += 1
        payload = json.dumps({"jsonrpc": "2.0", "id": self._request_id, "method": method, "params": params}).encode()
        request = urllib.request.Request(
            self.rpc_url,
            data=payload,
            headers={"content-type": "application/json", "accept": "application/json"},
            method="POST",
        )
        try:
            with self.opener(request, timeout=self.timeout) as response:
                raw = response.read()
        except (urllib.error.URLError, TimeoutError, OSError) as exc:
            raise VerificationError(f"rpc_error: {exc}") from exc
        try:
            data = json.loads(raw)
        except (TypeError, ValueError) as exc:
            raise VerificationError("rpc_error: response was not JSON") from exc
        if not isinstance(data, dict):
            raise VerificationError("rpc_error: response was not an object")
        if data.get("error") is not None:
            error = data["error"]
            message = error.get("message") if isinstance(error, dict) else str(error)
            raise VerificationError(f"rpc_error: {message}")
        return data.get("result")

    def get_transaction(self, signature: str) -> dict[str, Any] | None:
        result = self.call(
            "getTransaction",
            [signature, {"encoding": "jsonParsed", "commitment": "confirmed", "maxSupportedTransactionVersion": 0}],
        )
        return result if isinstance(result, dict) else None

    def get_signature_status(self, signature: str) -> dict[str, Any] | None:
        result = self.call("getSignatureStatuses", [[signature], {"searchTransactionHistory": True}])
        if not isinstance(result, dict):
            return None
        values = result.get("value")
        if not isinstance(values, list) or not values or not isinstance(values[0], dict):
            return None
        return values[0]


def _amount_from_balance(item: dict[str, Any]) -> int:
    ui_amount = item.get("uiTokenAmount")
    if not isinstance(ui_amount, dict):
        return 0
    raw = ui_amount.get("amount")
    try:
        return int(raw)
    except (TypeError, ValueError):
        return 0


def _balance_deltas(transaction: dict[str, Any], recipient: str, mint: str) -> int:
    meta = transaction.get("meta")
    if not isinstance(meta, dict):
        return 0
    before: dict[tuple[str, int], tuple[str | None, int]] = {}
    after: dict[tuple[str, int], tuple[str | None, int]] = {}
    for key, target in (("preTokenBalances", before), ("postTokenBalances", after)):
        rows = meta.get(key)
        if not isinstance(rows, list):
            continue
        for row in rows:
            if not isinstance(row, dict) or row.get("mint") != mint:
                continue
            try:
                index = int(row.get("accountIndex"))
            except (TypeError, ValueError):
                continue
            owner = row.get("owner") if isinstance(row.get("owner"), str) else None
            target[(mint, index)] = (owner, _amount_from_balance(row))

    total = 0
    for key in set(before) | set(after):
        owner = (after.get(key) or before.get(key))[0]
        if owner != recipient:
            continue
        total += (after.get(key) or (None, 0))[1] - (before.get(key) or (None, 0))[1]
    return max(total, 0)


def _transaction_slot(transaction: dict[str, Any]) -> int | None:
    try:
        return int(transaction["slot"])
    except (KeyError, TypeError, ValueError):
        return None


def _status_info(status: dict[str, Any] | None) -> tuple[str | None, int]:
    if not status:
        return None, 0
    confirmation_status = status.get("confirmationStatus")
    confirmations = status.get("confirmations")
    if confirmations is None and confirmation_status == "finalized":
        return confirmation_status, 2**31 - 1
    try:
        return confirmation_status if isinstance(confirmation_status, str) else None, max(0, int(confirmations or 0))
    except (TypeError, ValueError):
        return confirmation_status if isinstance(confirmation_status, str) else None, 0


def verify_payment(
    client: SolanaRpcClient,
    signature: str,
    recipient: str,
    expected_amount_units: int,
    mint: str = MAINNET_USDC_MINT,
    min_confirmations: int = 1,
    timeout_seconds: float = 30.0,
    poll_interval_seconds: float = 2.0,
    *,
    clock: Callable[[], float] = time.monotonic,
    sleep: Callable[[float], None] = time.sleep,
) -> VerificationResult:
    """Poll a signature until it is valid, confirmed, or times out."""
    if not signature or not recipient or not mint:
        raise ValueError("signature, recipient, and mint are required")
    if expected_amount_units < 0:
        raise ValueError("expected amount cannot be negative")
    if min_confirmations < 1:
        raise ValueError("min_confirmations must be at least 1")
    if timeout_seconds < 0 or poll_interval_seconds < 0:
        raise ValueError("timeouts cannot be negative")

    started = clock()
    attempts = 0
    last_slot: int | None = None
    last_status: str | None = None
    last_confirmations = 0
    last_observed = 0
    while True:
        attempts += 1
        try:
            status = client.get_signature_status(signature)
            transaction = client.get_transaction(signature)
        except VerificationError as exc:
            return VerificationResult(False, str(exc).split(":", 1)[0], signature, recipient, mint, expected_amount_units, last_observed, last_confirmations, min_confirmations, last_slot, last_status, attempts)

        last_status, status_confirmations = _status_info(status)
        if transaction is not None:
            last_slot = _transaction_slot(transaction)
            meta = transaction.get("meta")
            if isinstance(meta, dict) and meta.get("err") is not None:
                return VerificationResult(False, "transaction_failed", signature, recipient, mint, expected_amount_units, last_observed, status_confirmations, min_confirmations, last_slot, last_status, attempts)
            last_observed = _balance_deltas(transaction, recipient, mint)
            if last_observed != expected_amount_units:
                return VerificationResult(False, "amount_or_recipient_mismatch", signature, recipient, mint, expected_amount_units, last_observed, status_confirmations, min_confirmations, last_slot, last_status, attempts)
            last_confirmations = status_confirmations
            if last_status in {"confirmed", "finalized"} and last_confirmations >= min_confirmations:
                return VerificationResult(True, "verified", signature, recipient, mint, expected_amount_units, last_observed, last_confirmations, min_confirmations, last_slot, last_status, attempts)

        if clock() - started >= timeout_seconds:
            return VerificationResult(False, "timeout", signature, recipient, mint, expected_amount_units, last_observed, last_confirmations, min_confirmations, last_slot, last_status, attempts)
        sleep(min(poll_interval_seconds, max(0.0, timeout_seconds - (clock() - started))))


def _build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--rpc-url", default=DEFAULT_RPC_URL)
    parser.add_argument("--signature", required=True)
    parser.add_argument("--recipient", required=True)
    parser.add_argument("--amount-usdc", required=True, help="Exact decimal USDC amount, for example 1.25")
    parser.add_argument("--mint", default=MAINNET_USDC_MINT)
    parser.add_argument("--min-confirmations", type=int, default=1)
    parser.add_argument("--timeout", type=float, default=30.0, dest="timeout_seconds")
    parser.add_argument("--poll-interval", type=float, default=2.0, dest="poll_interval_seconds")
    parser.add_argument("--decimals", type=int, default=DEFAULT_DECIMALS)
    return parser


def main(argv: list[str] | None = None) -> int:
    args = _build_parser().parse_args(argv)
    try:
        expected = amount_to_units(args.amount_usdc, args.decimals)
        result = verify_payment(
            SolanaRpcClient(args.rpc_url),
            args.signature,
            args.recipient,
            expected,
            mint=args.mint,
            min_confirmations=args.min_confirmations,
            timeout_seconds=args.timeout_seconds,
            poll_interval_seconds=args.poll_interval_seconds,
        )
    except (ValueError, VerificationError) as exc:
        print(json.dumps({"verified": False, "reason": "invalid_input", "error": str(exc)}, sort_keys=True))
        return 2
    print(json.dumps(result.as_dict(), sort_keys=True))
    return 0 if result.verified else 1


if __name__ == "__main__":
    raise SystemExit(main())
