"""
poly samurai - paper trading executor (crypto cross-exchange spread)
---------------------------------------------------------------------
Same paper-trading pattern as paper_trade_fed.py, applied to BTC/ETH
spot spread between Coinbase and Kraken.

Uses the REAL fee-adjusted threshold from our overnight measurement run
(21,256 checks, max spread ever seen was 0.07%, needed 0.15% to clear
round-trip fees) - so don't be surprised if this sits "flat" a lot too.
That's the honest answer, not a bug. Crypto trades 24/7 though, so it
gets far more chances per day than a once-a-month Fed event.
"""

import argparse
import asyncio
import json
import sqlite3
import time
from datetime import datetime, timezone

import httpx

DB_PATH = "paper_trades_crypto.sqlite"
STATS_PATH = "crypto_live_stats.json"

ENTRY_THRESHOLD_PCT = 0.15   # real fee-adjusted threshold from our measurement run
EXIT_THRESHOLD_PCT = 0.05    # close once spread has narrowed back under this
MAX_HOLD_SECONDS = 600       # crypto moves fast - 10 min safety timeout, not 1 hour
NOTIONAL = 100.0

SYMBOLS = {
    "BTC": {"coinbase": "BTC-USD", "kraken": "XBTUSD"},
    "ETH": {"coinbase": "ETH-USD", "kraken": "ETHUSD"},
}

open_positions = {}   # {asset: {"entry_spread":, "entry_ts":, "buy_ex":, "sell_ex":}}
recent_log = []
cum_pnl = 0.0
fills = 0
attempts = 0


def init_db():
    conn = sqlite3.connect(DB_PATH)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS paper_trades (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            asset TEXT NOT NULL,
            buy_ex TEXT NOT NULL,
            sell_ex TEXT NOT NULL,
            entry_ts TEXT NOT NULL,
            exit_ts TEXT,
            entry_spread_pct REAL,
            exit_spread_pct REAL,
            pnl REAL,
            reason TEXT
        )
    """)
    conn.commit()
    return conn


async def get_coinbase(client, product):
    try:
        r = await client.get(f"https://api.exchange.coinbase.com/products/{product}/ticker", timeout=8)
        r.raise_for_status()
        d = r.json()
        return float(d["bid"]), float(d["ask"])
    except Exception:
        return None, None


async def get_kraken(client, pair):
    try:
        r = await client.get("https://api.kraken.com/0/public/Ticker", params={"pair": pair}, timeout=8)
        r.raise_for_status()
        d = r.json()
        result = d.get("result", {})
        if not result:
            return None, None
        key = next(iter(result))
        book = result[key]
        return float(book["b"][0]), float(book["a"][0])
    except Exception:
        return None, None


def log_event(msg):
    global recent_log
    ts = datetime.now(timezone.utc).strftime("%H:%M:%S")
    recent_log.insert(0, f"{ts}  {msg}")
    recent_log = recent_log[:30]
    print(f"{ts}  {msg}", flush=True)


def write_stats(asset_snapshots):
    stats = {
        "updated": datetime.now(timezone.utc).isoformat(),
        "net_pnl": round(cum_pnl, 2),
        "fills": fills,
        "attempts": attempts,
        "fill_rate": round(100 * fills / attempts, 1) if attempts else 0.0,
        "open_positions": len(open_positions),
        "assets": asset_snapshots,
        "log": recent_log,
    }
    with open(STATS_PATH, "w") as f:
        json.dump(stats, f, indent=2)


async def poll_asset(client, conn, asset, syms):
    global cum_pnl, fills, attempts
    ts_now = datetime.now(timezone.utc).isoformat()

    (cb_bid, cb_ask), (kr_bid, kr_ask) = await asyncio.gather(
        get_coinbase(client, syms["coinbase"]),
        get_kraken(client, syms["kraken"]),
    )
    quotes = {"coinbase": (cb_bid, cb_ask), "kraken": (kr_bid, kr_ask)}
    quotes = {k: v for k, v in quotes.items() if v[0] is not None and v[1] is not None}

    if len(quotes) < 2:
        snapshot = {"asset": asset, "status": "no data"}
        print(f"{asset:<4} not enough exchanges responded", flush=True)
        return snapshot

    # best buy-low/sell-high across the two exchanges
    ex_names = list(quotes.keys())
    ex_a, ex_b = ex_names[0], ex_names[1]
    bid_a, ask_a = quotes[ex_a]
    bid_b, ask_b = quotes[ex_b]

    # option 1: buy on A, sell on B
    spread1 = (bid_b - ask_a) / ask_a * 100
    # option 2: buy on B, sell on A
    spread2 = (bid_a - ask_b) / ask_b * 100

    if spread1 >= spread2:
        best_spread, buy_ex, sell_ex, buy_price, sell_price = spread1, ex_a, ex_b, ask_a, bid_b
    else:
        best_spread, buy_ex, sell_ex, buy_price, sell_price = spread2, ex_b, ex_a, ask_b, bid_a

    status = "flat"

    if asset not in open_positions and best_spread > ENTRY_THRESHOLD_PCT:
        open_positions[asset] = {
            "entry_spread": best_spread, "entry_ts": time.time(),
            "buy_ex": buy_ex, "sell_ex": sell_ex,
        }
        attempts += 1
        status = "opened"
        log_event(f"OPEN {asset}: buy {buy_ex} @ {buy_price:.2f} sell {sell_ex} @ {sell_price:.2f} spread={best_spread:.3f}%")

    elif asset in open_positions:
        pos = open_positions[asset]
        held_for = time.time() - pos["entry_ts"]
        converged = best_spread < EXIT_THRESHOLD_PCT
        timed_out = held_for > MAX_HOLD_SECONDS

        if converged or timed_out:
            pnl = (best_spread / 100) * NOTIONAL  # simplistic: captured spread on notional
            cum_pnl += pnl
            fills += 1
            reason = "converged" if converged else "timeout"
            status = f"closed ({reason})"
            log_event(f"CLOSE {asset}: pnl={pnl:+.2f} reason={reason}")

            conn.execute(
                """INSERT INTO paper_trades
                   (asset, buy_ex, sell_ex, entry_ts, exit_ts, entry_spread_pct, exit_spread_pct, pnl, reason)
                   VALUES (?,?,?,?,?,?,?,?,?)""",
                (asset, pos["buy_ex"], pos["sell_ex"],
                 datetime.fromtimestamp(pos["entry_ts"], tz=timezone.utc).isoformat(),
                 ts_now, pos["entry_spread"], best_spread, pnl, reason),
            )
            conn.commit()
            del open_positions[asset]
        else:
            status = "holding"

    print(f"{asset:<4} {ex_a}={quotes[ex_a][0]:.2f}/{quotes[ex_a][1]:.2f}  "
          f"{ex_b}={quotes[ex_b][0]:.2f}/{quotes[ex_b][1]:.2f}  "
          f"best_spread={best_spread:+.3f}%  [{status}]", flush=True)

    return {
        "asset": asset, "buy_ex": buy_ex, "sell_ex": sell_ex,
        "spread_pct": round(best_spread, 4), "status": status,
    }


async def main(interval: int):
    conn = init_db()
    async with httpx.AsyncClient() as client:
        log_event(f"crypto paper trader started, entry>{ENTRY_THRESHOLD_PCT}%, exit<{EXIT_THRESHOLD_PCT}%")
        while True:
            start = time.time()
            snapshots = []
            try:
                for asset, syms in SYMBOLS.items():
                    snap = await poll_asset(client, conn, asset, syms)
                    snapshots.append(snap)
                write_stats(snapshots)
            except Exception as e:
                log_event(f"poll error: {e}")
            elapsed = time.time() - start
            await asyncio.sleep(max(0, interval - elapsed))


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--interval", type=int, default=5)
    args = ap.parse_args()
    asyncio.run(main(args.interval))
