#!/usr/bin/env python3
"""
Collect live tick snapshots for Polymarket/Chainlink lag research.

Writes CSV rows compatible with:
  RenTech/strategy_stack/backtest_polymarket_chainlink_lag.py

Output columns:
  ts,binance_px,coinbase_px,chainlink_px,market_mid,orderbook_imbalance

Notes:
  - Binance and Coinbase are polled from public REST endpoints.
  - Chainlink BTC/USD is read on-chain via Ethereum JSON-RPC (latestRoundData).
  - Polymarket book is pulled from CLOB REST (`/book?token_id=...`).
"""

from __future__ import annotations

import argparse
import csv
import json
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

import requests

BINANCE_TICKER_URL = "https://api.binance.com/api/v3/ticker/price?symbol=BTCUSDT"
COINBASE_TICKER_URL = "https://api.exchange.coinbase.com/products/BTC-USD/ticker"
POLYMARKET_BOOK_URL = "https://clob.polymarket.com/book"

# Chainlink BTC / USD feed on Ethereum mainnet.
CHAINLINK_BTC_USD_FEED = "0xF4030086522a5bEEa4988F8cA5B36dbC97BeE88c"
CHAINLINK_DECIMALS_SELECTOR = "0x313ce567"
CHAINLINK_LATEST_ROUND_SELECTOR = "0xfeaf968c"


@dataclass
class TickRow:
    ts: str
    binance_px: float
    coinbase_px: float
    chainlink_px: float
    market_mid: float
    orderbook_imbalance: float


def _utc_now_iso() -> str:
    return datetime.now(timezone.utc).isoformat()


def _rpc_call(rpc_url: str, to_addr: str, data: str, timeout_sec: float) -> str:
    payload = {
        "jsonrpc": "2.0",
        "id": 1,
        "method": "eth_call",
        "params": [{"to": to_addr, "data": data}, "latest"],
    }
    r = requests.post(rpc_url, json=payload, timeout=timeout_sec)
    r.raise_for_status()
    body = r.json()
    if "error" in body:
        raise RuntimeError(f"rpc error: {body['error']}")
    out = body.get("result")
    if not isinstance(out, str) or not out.startswith("0x"):
        raise RuntimeError(f"unexpected rpc result: {body}")
    return out


def _hex_word_to_int(hex_data: str, word_idx: int) -> int:
    raw = hex_data[2:]
    start = word_idx * 64
    end = start + 64
    if len(raw) < end:
        raise ValueError("hex payload too short")
    return int(raw[start:end], 16)


def _hex_word_to_int_signed(hex_data: str, word_idx: int) -> int:
    v = _hex_word_to_int(hex_data, word_idx)
    bits = 256
    if v >= (1 << (bits - 1)):
        v -= 1 << bits
    return v


def fetch_chainlink_btc_usd(rpc_url: str, feed_address: str, timeout_sec: float) -> tuple[float, int]:
    dec_hex = _rpc_call(rpc_url, feed_address, CHAINLINK_DECIMALS_SELECTOR, timeout_sec)
    decimals = _hex_word_to_int(dec_hex, 0)
    latest_hex = _rpc_call(rpc_url, feed_address, CHAINLINK_LATEST_ROUND_SELECTOR, timeout_sec)
    # latestRoundData returns:
    # [0]=roundId, [1]=answer, [2]=startedAt, [3]=updatedAt, [4]=answeredInRound
    answer = _hex_word_to_int_signed(latest_hex, 1)
    updated_at = _hex_word_to_int(latest_hex, 3)
    px = float(answer) / float(10**decimals)
    return px, int(updated_at)


def fetch_binance_btc(timeout_sec: float) -> float:
    r = requests.get(BINANCE_TICKER_URL, timeout=timeout_sec)
    r.raise_for_status()
    body = r.json()
    return float(body["price"])


def fetch_coinbase_btc(timeout_sec: float) -> float:
    r = requests.get(COINBASE_TICKER_URL, timeout=timeout_sec, headers={"Accept": "application/json"})
    r.raise_for_status()
    body = r.json()
    return float(body["price"])


def _book_levels(body: dict[str, Any], side: str) -> list[dict[str, Any]]:
    v = body.get(side)
    if isinstance(v, list):
        return [x for x in v if isinstance(x, dict)]
    if isinstance(v, dict):
        levels = v.get("levels")
        if isinstance(levels, list):
            return [x for x in levels if isinstance(x, dict)]
    return []


def _level_num(level: dict[str, Any], keys: tuple[str, ...]) -> float:
    for k in keys:
        if k in level:
            try:
                return float(level[k])
            except (TypeError, ValueError):
                continue
    raise ValueError(f"no numeric key in level for keys={keys}: {level}")


def fetch_polymarket_mid_and_imbalance(token_id: str, depth_levels: int, timeout_sec: float) -> tuple[float, float]:
    r = requests.get(POLYMARKET_BOOK_URL, params={"token_id": token_id}, timeout=timeout_sec)
    r.raise_for_status()
    body = r.json()
    bids = _book_levels(body, "bids")
    asks = _book_levels(body, "asks")
    if not bids or not asks:
        raise RuntimeError(f"empty book: {body}")

    best_bid = max(_level_num(x, ("price", "p")) for x in bids)
    best_ask = min(_level_num(x, ("price", "p")) for x in asks)
    mid = 0.5 * (best_bid + best_ask)

    n = max(int(depth_levels), 1)
    bid_levels = sorted(bids, key=lambda x: _level_num(x, ("price", "p")), reverse=True)[:n]
    ask_levels = sorted(asks, key=lambda x: _level_num(x, ("price", "p")))[:n]
    bid_vol = sum(_level_num(x, ("size", "s", "quantity", "q")) for x in bid_levels)
    ask_vol = sum(_level_num(x, ("size", "s", "quantity", "q")) for x in ask_levels)
    if ask_vol <= 0:
        imbalance = 0.0
    else:
        imbalance = bid_vol / ask_vol
    return float(mid), float(imbalance)


def _append_row(path: Path, row: TickRow) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    file_exists = path.is_file()
    with path.open("a", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=["ts", "binance_px", "coinbase_px", "chainlink_px", "market_mid", "orderbook_imbalance"])
        if not file_exists:
            w.writeheader()
        w.writerow(
            {
                "ts": row.ts,
                "binance_px": f"{row.binance_px:.10f}",
                "coinbase_px": f"{row.coinbase_px:.10f}",
                "chainlink_px": f"{row.chainlink_px:.10f}",
                "market_mid": f"{row.market_mid:.10f}",
                "orderbook_imbalance": f"{row.orderbook_imbalance:.10f}",
            }
        )


def main() -> None:
    ap = argparse.ArgumentParser(description="Collect Polymarket/Chainlink lag tick CSV.")
    ap.add_argument("--token-id", required=True, help="Polymarket CLOB token id to query.")
    ap.add_argument("--out-csv", type=Path, default=Path("RenTech/data/logs/polymarket_chainlink_ticks.csv"))
    ap.add_argument("--rpc-url", default="https://ethereum.publicnode.com")
    ap.add_argument("--chainlink-feed-address", default=CHAINLINK_BTC_USD_FEED)
    ap.add_argument("--interval-sec", type=float, default=1.0, help="Polling interval in seconds.")
    ap.add_argument("--duration-sec", type=float, default=3600.0, help="Total run duration in seconds.")
    ap.add_argument("--book-depth-levels", type=int, default=5, help="Depth levels to compute imbalance.")
    ap.add_argument("--timeout-sec", type=float, default=8.0)
    ap.add_argument("--max-errors", type=int, default=20)
    args = ap.parse_args()

    started = time.time()
    n_rows = 0
    n_err = 0
    print(
        json.dumps(
            {
                "collector": "polymarket_chainlink_ticks",
                "out_csv": str(args.out_csv),
                "token_id": args.token_id,
                "interval_sec": float(args.interval_sec),
                "duration_sec": float(args.duration_sec),
            }
        )
    )

    while True:
        now = time.time()
        if now - started > float(args.duration_sec):
            break
        t0 = time.time()
        try:
            bin_px = fetch_binance_btc(float(args.timeout_sec))
            cb_px = fetch_coinbase_btc(float(args.timeout_sec))
            cl_px, cl_updated = fetch_chainlink_btc_usd(
                args.rpc_url, args.chainlink_feed_address, float(args.timeout_sec)
            )
            mkt_mid, imb = fetch_polymarket_mid_and_imbalance(
                args.token_id, int(args.book_depth_levels), float(args.timeout_sec)
            )
            row = TickRow(
                ts=_utc_now_iso(),
                binance_px=float(bin_px),
                coinbase_px=float(cb_px),
                chainlink_px=float(cl_px),
                market_mid=float(mkt_mid),
                orderbook_imbalance=float(imb),
            )
            _append_row(args.out_csv, row)
            n_rows += 1
            print(
                f"row={n_rows} ts={row.ts} binance={row.binance_px:.2f} coinbase={row.coinbase_px:.2f} "
                f"chainlink={row.chainlink_px:.2f} cl_updated={cl_updated} mid={row.market_mid:.4f} imb={row.orderbook_imbalance:.4f}"
            )
        except Exception as e:  # keep loop alive; fail only after repeated errors
            n_err += 1
            print(f"warn error_count={n_err}: {e}")
            if n_err >= int(args.max_errors):
                raise RuntimeError(f"too many errors ({n_err}); aborting") from e

        elapsed = time.time() - t0
        sleep_for = max(float(args.interval_sec) - elapsed, 0.0)
        if sleep_for > 0:
            time.sleep(sleep_for)

    print(json.dumps({"status": "done", "rows": n_rows, "errors": n_err, "out_csv": str(args.out_csv)}))


if __name__ == "__main__":
    main()
