# Gateway quickstart: create a user, fund their wallet, buy
# 2 USDC of SOL for them, and find the fill in your organization's
# 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. Running it again reuses the same user
# and buys another 2 USDC of SOL while the 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 cryptography.hazmat.primitives.serialization import (
    Encoding,
    PublicFormat,
)
from turnkey_api_key_stamper import (
    ApiKeyStamper,
    ApiKeyStamperConfig,
)

API = "https://api.truemarkets.co"
API_KEY_FILE = os.environ["TM_KEY_FILE"]
SIGNER_KEY_FILE = os.environ.get(
    "SIGNER_KEY_FILE", "signer-key.json"
)
# Your own id for this customer.
EXTERNAL_REF_ID = "quickstart_user_1"
ORDER_USDC = "2"

token = ""


class ApiError(Exception):
    def __init__(self, status, body, where):
        self.status = status
        reason = body.get("code") or body.get("type", "")
        detail = f"{status} {reason} {body.get('message', '')}"
        super().__init__(f"{where}: {detail}")


# Every call sends the organization token. Calls for one user
# add TM-On-Behalf-Of.
def call(method, path, body=None, headers=None):
    res = requests.request(
        method,
        f"{API}{path}",
        json=body,
        headers={
            **(
                {"Authorization": f"Bearer {token}"}
                if token
                else {}
            ),
            **(headers or {}),
        },
        timeout=60,
    )
    data = res.json() if res.content else {}
    if not res.ok:
        raise ApiError(res.status_code, data, f"{method} {path}")
    return data


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


def get(path, headers=None):
    return call("GET", path, None, headers)


def b64url_int(value):
    return int.from_bytes(
        base64.urlsafe_b64decode(value + "=" * (-len(value) % 4)),
        "big",
    )


# region signer
# The signer key signs every wallet transaction for the users you
# register it on. Keep this file: a user can only ever sign with
# the key it was created with.
if not os.path.exists(SIGNER_KEY_FILE):
    key = ec.generate_private_key(ec.SECP256R1())
    file = {
        "signer_public_key": key.public_key()
        .public_bytes(Encoding.X962, PublicFormat.CompressedPoint)
        .hex(),
        "signer_private_key": key.private_numbers()
        .private_value.to_bytes(32, "big")
        .hex(),
    }
    descriptor = os.open(
        SIGNER_KEY_FILE, os.O_WRONLY | os.O_CREAT, 0o600
    )
    with os.fdopen(descriptor, "w") as f:
        json.dump(file, f)
with open(SIGNER_KEY_FILE) as f:
    signer = json.load(f)
stamper = ApiKeyStamper(
    ApiKeyStamperConfig(
        api_public_key=signer["signer_public_key"],
        api_private_key=signer["signer_private_key"],
    )
)


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


# endregion signer


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
    with open(API_KEY_FILE) as f:
        api_key = json.load(f)
    jwk = api_key["private_key"]
    private_key = ec.EllipticCurvePrivateNumbers(
        b64url_int(jwk["d"]),
        ec.EllipticCurvePublicNumbers(
            b64url_int(jwk["x"]),
            b64url_int(jwk["y"]),
            ec.SECP256R1(),
        ),
    ).private_key()
    timestamp = int(time.time())
    der = private_key.sign(
        f"{api_key['key_id']}.{timestamp}".encode(),
        ec.ECDSA(hashes.SHA256()),
    )
    # The API expects r and s concatenated.
    r, s = decode_dss_signature(der)
    signature = (
        base64.urlsafe_b64encode(
            r.to_bytes(32, "big") + s.to_bytes(32, "big")
        )
        .rstrip(b"=")
        .decode()
    )

    minted = post(
        "/v1/auth/api-key/token",
        {
            "key_id": api_key["key_id"],
            "timestamp": timestamp,
            "signature": signature,
        },
    )
    token = minted["access_token"]

    # The token names your organization: no id to configure.
    claims = json.loads(
        base64.urlsafe_b64decode(token.split(".")[1] + "==")
    )
    organization_id = claims["tm"]["organization_id"]
    # endregion token
    print(f"✓ token     expires {minted['expires_in']}")

    # region user
    for attempt in range(1, 6):
        try:
            user = post(
                f"/v1/account/organizations/{organization_id}/users",
                {
                    "external_ref_id": EXTERNAL_REF_ID,
                    "signer_public_key": signer[
                        "signer_public_key"
                    ],
                },
            )
            if user["wallets"]:
                break
        except ApiError as err:
            # A 503 means the wallets aren't ready yet. The same
            # request finishes them.
            if err.status != 503 or attempt == 5:
                raise
        time.sleep(2)
    user_id = user["user_id"]
    for_user = {"TM-On-Behalf-Of": user_id}
    solana_address = next(
        w["address"]
        for w in user["wallets"]
        if w["chain_family"] == "solana"
    )
    # endregion user
    print(f"✓ user      {user_id}   solana {solana_address}")

    # region fund
    def usdc_available():
        balances = get("/v1/gateway/balances", for_user)["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(
            f"… send about 3 USDC on Solana to {solana_address}"
        )
        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",
        },
        for_user,
    )
    # An empty order_id means no order was created, and
    # quote.issues says why.
    if not order.get("order_id"):
        raise RuntimeError(
            f"no order: {order.get('quote', {}).get('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",
        },
        for_user,
    )
    # 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']}", for_user
        )
    print(
        f"✓ order     {filled['status']}  "
        f"{filled.get('executed_qty', '')} SOL for {ORDER_USDC} USDC"
    )

    # Your organization's transactions cover every user, so this
    # call has no header.
    feed = get(
        f"/v1/gateway/organizations/{organization_id}/transactions"
        "?type=order"
    )
    row = next(
        (t for t in feed["data"] if t["id"] == order["order_id"]),
        None,
    )
    # endregion fill
    print(
        f"✓ in your organization's transactions ({row['status']})"
        if row
        else "… not in the transactions list yet"
    )


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