"""Dashing Crypto Prepaid Transaction Fee API: a complete client in one file.

    pip install coincurve pycryptodome base58 requests

Every request is signed by the account's key. A read may carry a read token instead, which a
signed POST /{address}/tokens returns: PrepaidReader reads with one and holds no key. GET /terms
needs neither. The signed message is five lines:

    Dashing Crypto prepaid v1
    METHOD
    /nile/prepaid/T.../transfers          the path as sent, network prefix included
    issued=1790467200&limit=25             the query without signature, encoded and sorted
    e3b0c442...                            SHA-256 of the body as sent, lower-case hex

signed as a TRON message (TIP-191): keccak-256 over "\\x19TRON Signed Message:\\n", the
message's length in bytes and the message; secp256k1; r || s || v with v 27 or 28.
"""

import hashlib
import json
import time
from urllib.parse import quote, urlsplit

import base58
import requests
from Crypto.Hash import keccak
from coincurve import PrivateKey

NETWORKS = {
    "mainnet": {
        "api": "https://api.crypto.dashing.ws/prepaid",
        "node": "https://api.trongrid.io",
        "usdt": "TR7NHqjeKQxGTCi8q8ZY4pL8otSzgjLj6t",
    },
    "nile": {
        "api": "https://api.crypto.dashing.ws/nile/prepaid",
        "node": "https://nile.trongrid.io",
        "usdt": "TXYZopYRdj2D9XRtbG411XZZ3kM5VkAeBf",
    },
}


def keccak256(data: bytes) -> bytes:
    return keccak.new(digest_bits=256, data=data).digest()


def address_of(key: PrivateKey) -> str:
    """The base58 TRON address of a key: 0x41 + the last 20 bytes of keccak(public key)."""
    public = key.public_key.format(compressed=False)[1:]
    return base58.b58encode_check(b"\x41" + keccak256(public)[-20:]).decode()


def encode(value: str) -> str:
    """RFC 3986: A-Z a-z 0-9 - . _ ~ stay; every other byte is %XX in upper-case hex."""
    return quote(value, safe="-._~")


def canonical_query(params: dict) -> str:
    """The query without signature: each name and value encoded, sorted, joined by &."""
    pairs = sorted((encode(str(k)), encode(str(v))) for k, v in params.items())
    return "&".join(f"{k}={v}" for k, v in pairs)


def sign_message(key: PrivateKey, message: str) -> str:
    """TIP-191, as TronWeb's signMessageV2 does it: 65 bytes r || s || v as hex, v 27 or 28."""
    body = message.encode("utf-8")
    digest = keccak256(b"\x19TRON Signed Message:\n" + str(len(body)).encode() + body)
    signature = key.sign_recoverable(digest, hasher=None)  # r || s || recovery id (0 or 1)
    return (signature[:64] + bytes([signature[64] + 27])).hex()


class PrepaidError(Exception):
    def __init__(self, status: int, body: dict):
        super().__init__(f"{status} {body.get('code')}: {body.get('message')}")
        self.status, self.body = status, body


class PrepaidClient:
    def __init__(self, api: str, node: str, usdt: str, private_key: str):
        self.api = api.rstrip("/")
        self.node = node
        self.usdt = usdt
        self.key = PrivateKey(bytes.fromhex(private_key))
        self.address = address_of(self.key)

    def request(self, method: str, path: str = "", query: dict = None, body: dict = None) -> dict:
        full_path = f"{urlsplit(self.api).path}/{self.address}{path}"
        params = dict(query or {}, issued=int(time.time()))
        canonical = canonical_query(params)
        body_bytes = b"" if body is None else json.dumps(body, separators=(",", ":")).encode()
        message = "\n".join([
            "Dashing Crypto prepaid v1",
            method,
            full_path,
            canonical,
            hashlib.sha256(body_bytes).hexdigest(),
        ])
        signature = sign_message(self.key, message)

        origin = "{0.scheme}://{0.netloc}".format(urlsplit(self.api))
        response = requests.request(
            method,
            f"{origin}{full_path}?{canonical}&signature={signature}",
            data=body_bytes if body is not None else None,  # the exact bytes that were hashed
            headers={"Content-Type": "application/json"} if body is not None else {},
            timeout=30,
        )
        return answer(response)

    def account(self):
        return self.request("GET")

    def mint_token(self) -> str:
        """A read token for this account, prt_...; shown once. Keep it for PrepaidReader."""
        return self.request("POST", "/tokens")["token"]

    def revoke_tokens(self):
        """Ends every read token this account holds, for a phone that was lost."""
        self.request("DELETE", "/tokens")

    def quote(self, sender: str, to: str, amount: str):
        return self.request("POST", "/quotes", body={"from": sender, "to": to, "amount": amount})

    def transfer(self, sender: str, to: str, amount: str, signed_transaction: str, max_price: str = None):
        body = {"from": sender, "to": to, "amount": amount, "signedTransaction": signed_transaction}
        if max_price:
            body["maxPrice"] = max_price
        return self.request("POST", "/transfers", body=body)

    def deposit_by_txid(self, tx_id: str):
        return self.request("POST", "/deposits", body={"txId": tx_id})

    def deposit_sponsored(self, signed_transaction: str, max_sponsored_fee: str = None):
        body = {"signedTransaction": signed_transaction}
        if max_sponsored_fee:
            body["maxSponsoredFee"] = max_sponsored_fee
        return self.request("POST", "/deposits", body=body)

    def transactions(self, limit: int = 25, cursor: str = None):
        query = {"limit": limit}
        if cursor:
            query["cursor"] = cursor
        return self.request("GET", "/transactions", query=query)

    def transaction(self, tx: str):
        return self.request("GET", f"/transactions/{tx}")


# ---- the USDT transfer a sender signs ------------------------------------------------------


def signed_usdt_transfer(node: str, usdt: str, private_key: str, to: str, amount: str) -> dict:
    """A USDT transfer built by the node, given a five-minute expiry, and signed by the sender.

    Returns {"txID", "hex"}: hex is the whole signed Transaction, as the API takes it.
    """
    key = PrivateKey(bytes.fromhex(private_key))
    sender = address_of(key)
    micros = round(float(amount) * 1_000_000)
    recipient = base58.b58decode_check(to)[1:]  # 20 bytes
    parameter = recipient.rjust(32, b"\0").hex() + micros.to_bytes(32, "big").hex()

    built = requests.post(f"{node}/wallet/triggersmartcontract", json={
        "owner_address": sender,
        "contract_address": usdt,
        "function_selector": "transfer(address,uint256)",
        "parameter": parameter,
        "fee_limit": 20_000_000,
        "call_value": 0,
        "visible": True,
    }, timeout=30).json()["transaction"]

    # A node stamps a one-minute expiry; the API wants two to ten minutes left. Field 8 of
    # raw_data is the expiry in milliseconds, and the transaction id is SHA-256 of raw_data.
    raw = with_expiration(bytes.fromhex(built["raw_data_hex"]), int(time.time() * 1000) + 5 * 60 * 1000)
    tx_id = hashlib.sha256(raw).digest()
    signature = PrivateKey(bytes.fromhex(private_key)).sign_recoverable(tx_id, hasher=None)

    # Transaction { raw_data = 1; signature = 2 }
    signed = b"\x0a" + varint(len(raw)) + raw + b"\x12" + varint(len(signature)) + signature
    return {"txID": tx_id.hex(), "hex": signed.hex()}


def varint(n: int) -> bytes:
    out = bytearray()
    while n > 0x7F:
        out.append((n & 0x7F) | 0x80)
        n >>= 7
    out.append(n)
    return bytes(out)


def read_varint(data: bytes, at: int):
    value = shift = 0
    while True:
        b = data[at]
        at += 1
        value |= (b & 0x7F) << shift
        if not b & 0x80:
            return value, at
        shift += 7


def with_expiration(raw: bytes, expiration_ms: int) -> bytes:
    """raw_data with field 8 (expiration) replaced and every other field copied as it is."""
    out, at = bytearray(), 0
    while at < len(raw):
        start = at
        tag, at = read_varint(raw, at)
        wire = tag & 7
        if wire == 0:
            _, at = read_varint(raw, at)
        elif wire == 2:
            length, at = read_varint(raw, at)
            at += length
        elif wire == 1:
            at += 8
        elif wire == 5:
            at += 4
        if tag >> 3 == 8 and wire == 0:
            out += varint(tag) + varint(expiration_ms)
        else:
            out += raw[start:at]
    return bytes(out)


def broadcast(node: str, signed_hex: str) -> dict:
    return requests.post(f"{node}/wallet/broadcasthex", json={"transaction": signed_hex}, timeout=30).json()


class PrepaidReader:
    """Reads one account with a read token, and holds no key.

    For a wallet that keeps the key behind a fingerprint. A token unused for 180 days, or revoked,
    raises PrepaidError with code TOKEN_EXPIRED; sign PrepaidClient.mint_token() for a new one.
    """

    def __init__(self, api: str, address: str, token: str):
        self.base = f"{api.rstrip('/')}/{address}"
        self.token = token

    def request(self, method: str, path: str = "", query: dict = None):
        response = requests.request(
            method,
            f"{self.base}{path}",
            params=query,
            headers={"Authorization": f"Bearer {self.token}"},
            timeout=30,
        )
        return answer(response)

    def account(self):
        return self.request("GET")

    def transactions(self, limit: int = 25, cursor: str = None):
        query = {"limit": limit}
        if cursor:
            query["cursor"] = cursor
        return self.request("GET", "/transactions", query=query)

    def transaction(self, tx: str):
        return self.request("GET", f"/transactions/{tx}")

    def revoke(self):
        """Ends this token, and only this one."""
        self.request("DELETE", "/tokens")


def terms(api: str) -> dict:
    """The deposit terms for an address with none of its own. No key and no token."""
    return answer(requests.get(f"{api.rstrip('/')}/terms", timeout=30))


def answer(response: requests.Response):
    """The body, or PrepaidError carrying the status and the body's code; None for 204."""
    if response.status_code == 204:
        return None
    payload = response.json()
    if not response.ok:
        raise PrepaidError(response.status_code, payload)
    return payload
