# Trading API quickstart: mint a token, buy 2 USDC of SOL, and
# find the fill in your transactions.
#
#   pip install requests cryptography turnkey-api-key-stamper
#   TM_KEY_FILE=path/to/your-api-key.json python quickstart.py
#
# This trades real funds. Each run buys another 2 USDC of SOL
# while your wallet has funds.
import base64
import json
import os
import sys
import time

import requests
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import ec
from cryptography.hazmat.primitives.asymmetric.utils import (
    decode_dss_signature,
)
from turnkey_api_key_stamper import (
    ApiKeyStamper,
    ApiKeyStamperConfig,
)

API = "https://api.truemarkets.co"
KEY_FILE = os.environ["TM_KEY_FILE"]
ORDER_USDC = "2"

token = ""


def call(method, path, body=None):
    headers = (
        {"Authorization": f"Bearer {token}"} if token else {}
    )
    res = requests.request(
        method,
        f"{API}{path}",
        json=body,
        headers=headers,
        timeout=60,
    )
    data = res.json() if res.content else {}
    if not res.ok:
        reason = data.get("code") or data.get("type", "")
        message = data.get("message", "")
        detail = f"{res.status_code} {reason} {message}"
        raise RuntimeError(f"{method} {path}: {detail}")
    return data


def post(path, body):
    return call("POST", path, body)


def get(path):
    return call("GET", path)


def b64url_bytes(value):
    return base64.urlsafe_b64decode(
        value + "=" * (-len(value) % 4)
    )


# region key
# One key file does both jobs: it signs the token challenge, and,
# because the app registered it on your wallet, your payloads.
with open(KEY_FILE) as f:
    key_file = json.load(f)
key_id, jwk = key_file["key_id"], key_file["private_key"]
y = b64url_bytes(jwk["y"])
stamper = ApiKeyStamper(
    ApiKeyStamperConfig(
        # The compressed public key and the private key, as hex.
        api_public_key=(
            bytes([2 if y[31] % 2 == 0 else 3])
            + b64url_bytes(jwk["x"])
        ).hex(),
        api_private_key=b64url_bytes(jwk["d"]).hex(),
    )
)


def stamp(payload):
    return stamper.stamp(payload).stamp_header_value


# endregion key


def main():
    global token

    # region assets
    assets = get("/v1/gateway/assets")["data"]
    sol = next(
        a
        for a in assets
        if a["symbol"] == "SOL" and a["chain"] == "solana"
    )
    # endregion assets

    # region token
    private_key = ec.EllipticCurvePrivateNumbers(
        int.from_bytes(b64url_bytes(jwk["d"]), "big"),
        ec.EllipticCurvePublicNumbers(
            int.from_bytes(b64url_bytes(jwk["x"]), "big"),
            int.from_bytes(b64url_bytes(jwk["y"]), "big"),
            ec.SECP256R1(),
        ),
    ).private_key()
    timestamp = int(time.time())
    der = private_key.sign(
        f"{key_id}.{timestamp}".encode(),
        ec.ECDSA(hashes.SHA256()),
    )
    # The API expects r and s concatenated.
    r, s = decode_dss_signature(der)
    raw = r.to_bytes(32, "big") + s.to_bytes(32, "big")
    signature = (
        base64.urlsafe_b64encode(raw).rstrip(b"=").decode()
    )

    minted = post(
        "/v1/auth/api-key/token",
        {
            "key_id": key_id,
            "timestamp": timestamp,
            "signature": signature,
        },
    )
    token = minted["access_token"]
    # minted["refresh_token"] gets a new pair later, with no
    # signing.
    # endregion token
    print(f"✓ token     expires {minted['expires_in']}")

    # region fund
    def usdc_available():
        balances = get("/v1/gateway/balances")["data"]
        usdc = next(
            (
                b
                for b in balances
                if b["symbol"] == "USDC"
                and b["chain"] == "solana"
            ),
            None,
        )
        return float(usdc["available"]) if usdc else 0.0

    if usdc_available() < float(ORDER_USDC):
        print("… send about 3 USDC on Solana to your wallet")
        while usdc_available() < float(ORDER_USDC):
            time.sleep(10)
    # endregion fund
    print("✓ funded")

    # region order
    order = post(
        "/v1/gateway/orders",
        {
            "asset_id": sol["id"],
            "qty": ORDER_USDC,
            "qty_unit": "quote",
            "side": "buy",
            "type": "market",
        },
    )
    # An empty order_id means no order was created, and
    # quote.issues says why.
    if not order.get("order_id"):
        issues = order.get("quote", {}).get("issues")
        raise RuntimeError(f"no order: {issues}")

    # Sign each payload and execute straight away, before the
    # quote expires.
    signatures = [stamp(p["payload"]) for p in order["payloads"]]
    executed = post(
        f"/v1/gateway/orders/{order['order_id']}/execute",
        {"signatures": signatures, "auth_type": "api_key"},
    )
    # endregion order

    # region fill
    filled = executed
    while filled["status"] not in (
        "complete",
        "canceled",
        "failed",
    ):
        time.sleep(2)
        filled = get(f"/v1/gateway/orders/{order['order_id']}")
    print(
        f"✓ order     {filled['status']}  "
        f"{filled.get('executed_qty', '')} SOL "
        f"for {ORDER_USDC} USDC"
    )

    feed = get("/v1/gateway/transactions?type=order")
    row = next(
        (t for t in feed["data"] if t["id"] == order["order_id"]),
        None,
    )
    # endregion fill
    print(
        f"✓ in your transactions ({row['status']})"
        if row
        else "… not in your transactions yet"
    )


if __name__ == "__main__":
    try:
        main()
    except Exception as err:
        print(f"✗ {err}", file=sys.stderr)
        sys.exit(1)
