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

POLY_GAMMA = "https://gamma-api.polymarket.com/markets"
KALSHI_BASE = "https://external-api.kalshi.com/trade-api/v2"
DB_PATH = "paper_trades.sqlite"
STATS_PATH = "live_stats.json"

ENTRY_THRESHOLD = 0.03
EXIT_THRESHOLD = 0.01
MAX_HOLD_SECONDS = 3600
NOTIONAL = 100.0

MARKET_PAIRS = [
    {
        "label": "fed-maintain-sept",
        "poly_slug": "will-there-be-no-change-in-fed-interest-rates-after-the-september-2026-meeting-615",
        "kalshi_ticker": "KXFEDDECISION-26SEP-H0",
    },
    {
        "label": "fed-hike25-sept",
        "poly_slug": "will-the-fed-increase-interest-rates-by-25-bps-after-the-september-2026-meeting-649",
        "kalshi_ticker": "KXFEDDECISION-26SEP-H25",
    },
    {
        "label": "fed-cut25-sept",
        "poly_slug": "will-the-fed-decrease-interest-rates-by-25-bps-after-the-september-2026-meeting-586",
        "kalshi_ticker": "KXFEDDECISION-26SEP-C25",
    },
]

open_positions = {}
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,
            label TEXT NOT NULL,
            side TEXT NOT NULL,
            entry_ts TEXT NOT NULL,
            exit_ts TEXT,
            entry_spread REAL,
            exit_spread REAL,
            pnl REAL,
            reason TEXT
        )
    """)
    conn.commit()
    return conn


async def get_poly_yes_price(client, slug):
    try:
        r = await client.get(POLY_GAMMA, params={"slug": slug}, timeout=10)
        r.raise_for_status()
        data = r.json()
        if not data:
            return None
        prices = json.loads(data[0]["outcomePrices"])
        return float(prices[0])
    except Exception:
        return None


async def get_kalshi_yes_price(client, ticker):
    try:
        r = await client.get(f"{KALSHI_BASE}/markets/{ticker}", timeout=10)
        r.raise_for_status()
        m = r.json().get("market", {})
        yes_ask = m.get("yes_ask_dollars")
        if yes_ask is None:
            return None
        return float(yes_ask)
    except Exception:
        return 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(pair_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),
        "pairs": pair_snapshots,
        "log": recent_log,
    }
    with open(STATS_PATH, "w") as f:
        json.dump(stats, f, indent=2)


async def poll_once(client, conn):
    global cum_pnl, fills, attempts
    ts_now = datetime.now(timezone.utc).isoformat()
    pair_snapshots = []

    for pair in MARKET_PAIRS:
        label = pair["label"]
        poly_yes, kalshi_yes = await asyncio.gather(
            get_poly_yes_price(client, pair["poly_slug"]),
            get_kalshi_yes_price(client, pair["kalshi_ticker"]),
        )
        if poly_yes is None or kalshi_yes is None:
            pair_snapshots.append({"label": label, "status": "no data"})
            print(f"{label:<20} no data (poly={poly_yes} kalshi={kalshi_yes})", flush=True)
            continue

        spread = poly_yes - kalshi_yes
        status = "flat"

        if label not in open_positions and abs(spread) > ENTRY_THRESHOLD:
            side = "long_poly_short_kalshi" if spread < 0 else "short_poly_long_kalshi"
            open_positions[label] = {"entry_spread": spread, "entry_ts": time.time(), "side": side}
            attempts += 1
            status = "opened"
            log_event(f"OPEN {label}: spread={spread:+.3f} side={side}")

        elif label in open_positions:
            pos = open_positions[label]
            held_for = time.time() - pos["entry_ts"]
            converged = abs(spread) < EXIT_THRESHOLD
            timed_out = held_for > MAX_HOLD_SECONDS

            if converged or timed_out:
                move = pos["entry_spread"] - spread if pos["side"] == "long_poly_short_kalshi" else spread - pos["entry_spread"]
                pnl = move * NOTIONAL
                cum_pnl += pnl
                fills += 1
                reason = "converged" if converged else "timeout"
                status = f"closed ({reason})"
                log_event(f"CLOSE {label}: pnl={pnl:+.2f} reason={reason}")

                conn.execute(
                    """INSERT INTO paper_trades
                       (label, side, entry_ts, exit_ts, entry_spread, exit_spread, pnl, reason)
                       VALUES (?,?,?,?,?,?,?,?)""",
                    (label, pos["side"],
                     datetime.fromtimestamp(pos["entry_ts"], tz=timezone.utc).isoformat(),
                     ts_now, pos["entry_spread"], spread, pnl, reason),
                )
                conn.commit()
                del open_positions[label]
            else:
                status = "holding"

        pair_snapshots.append({
            "label": label, "poly_yes": round(poly_yes, 3), "kalshi_yes": round(kalshi_yes, 3),
            "spread": round(spread, 3), "status": status,
        })
        print(f"{label:<20} poly={poly_yes:.3f} kalshi={kalshi_yes:.3f} spread={spread:+.3f}  [{status}]", flush=True)

    write_stats(pair_snapshots)


async def main(interval):
    conn = init_db()
    async with httpx.AsyncClient() as client:
        log_event(f"paper trader started, {len(MARKET_PAIRS)} pairs, entry>{ENTRY_THRESHOLD}, exit<{EXIT_THRESHOLD}")
        while True:
            start = time.time()
            try:
                await poll_once(client, conn)
            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=15)
    args = ap.parse_args()
    asyncio.run(main(args.interval))
