"""
Stage 0: OKX 数据管道
拉取 BTC/USDT 永续合约 15分钟K线 → parquet存储
"""
import os, time, hmac, hashlib, base64, json
import requests
import pandas as pd
import numpy as np

# ── OKX 认证 ──
API_KEY = os.getenv("OKX_API_KEY", "")
SECRET = os.getenv("OKX_SECRET_KEY", "")
PASSPHRASE = os.getenv("OKX_PASSPHRASE", "")

def okx_get(path, params=None):
    ts = time.strftime("%Y-%m-%dT%H:%M:%S.000Z", time.gmtime())
    qs = "?" + "&".join(f"{k}={v}" for k,v in (params or {}).items()) if params else ""
    sign_str = ts + "GET" + path + qs + ""  # GET body = empty
    sig = base64.b64encode(hmac.new(SECRET.encode(), sign_str.encode(), hashlib.sha256).digest()).decode()
    
    resp = requests.get(
        f"https://www.okx.com{path}{qs}",
        headers={
            "OK-ACCESS-KEY": API_KEY,
            "OK-ACCESS-SIGN": sig,
            "OK-ACCESS-TIMESTAMP": ts,
            "OK-ACCESS-PASSPHRASE": PASSPHRASE,
        },
        timeout=15
    )
    data = resp.json()
    if data.get("code") != "0":
        raise Exception(f"OKX error {data.get('code')}: {data.get('msg')}")
    return data["data"]


# ── 拉取 K 线（分页） ──
def fetch_all_candles(inst_id="BTC-USDT-SWAP", bar="15m", limit=300, max_pages=20):
    """分页拉取历史K线。OKX单次最多300根，每页往回翻"""
    all_rows = []
    after_ts = None
    
    for page in range(max_pages):
        params = {"instId": inst_id, "bar": bar, "limit": str(limit)}
        if after_ts:
            params["after"] = str(after_ts)
        
        data = okx_get("/api/v5/market/candles", params)
        if not data or len(data) < 2:
            break
        
        all_rows.extend(data)
        oldest_ts = int(data[-1][0])  # 最旧一根的时间戳(ms)
        after_ts = oldest_ts - 1
        print(f"  page {page+1}: {len(data)} candles, oldest={pd.to_datetime(oldest_ts, unit='ms')}")
        time.sleep(0.3)  # 限速
    
    return all_rows


# ── 主流程 ──
print("📡 Fetching BTC-USDT-SWAP 15m candles...")
raw = fetch_all_candles()

# 解析为 DataFrame
cols = ["ts","open","high","low","close","vol","vol_ccy","vol_ccy_quote","confirm"]
df = pd.DataFrame(raw, columns=cols)
df["ts"] = pd.to_datetime(df["ts"].astype(np.int64), unit="ms")
for c in ["open","high","low","close","vol","vol_ccy"]:
    df[c] = df[c].astype(float)
df = df.sort_values("ts").reset_index(drop=True)

# 保存
os.makedirs("data", exist_ok=True)
path = "data/btc_15m.parquet"
df.to_parquet(path, index=False)

print(f"\n✅ 保存到 {path}")
print(f"   {len(df)} 根K线")
print(f"   时间范围: {df['ts'].min()} → {df['ts'].max()}")
print(f"   O-H-L-C-V 完整: {df.notna().all().all()}")
