"""
HYPERION ENGINE — Phase 2.1 worker
Multi-pair paper trading with full activity logging: every scan, regime call,
signal, entry, rejection, close, breaker, command, and error is written to
hyperion_activity for the dashboard's live feed.
"""
from __future__ import annotations
import os, time, traceback
from datetime import datetime, timezone

import requests
from dotenv import load_dotenv

from hl_client import HLInfo
from strategies import best_signal, classify_regime
from paper import PaperBook
from activity import ActivityLog

load_dotenv()

SUPABASE_URL = os.environ["SUPABASE_URL"].rstrip("/")
SERVICE_KEY = os.environ["SUPABASE_SERVICE_KEY"]
HL_MAINNET = os.environ.get("HL_MAINNET", "true").lower() == "true"

TICK_SECONDS = int(os.environ.get("TICK_SECONDS", "30"))
FILLS_SYNC_EVERY = 4
CANDLE_REFRESH_EVERY = 4
CANDLE_INTERVAL = "15m"
CANDLE_LOOKBACK_H = 96
PAPER_START_BALANCE = float(os.environ.get("PAPER_START_BALANCE", "1000"))

HEADERS = {"apikey": SERVICE_KEY, "Authorization": f"Bearer {SERVICE_KEY}",
           "Content-Type": "application/json"}

hl = HLInfo(mainnet=HL_MAINNET)


# ---------------- supabase ----------------
def sb(table): return f"{SUPABASE_URL}/rest/v1/{table}"

def sb_get(table, params):
    r = requests.get(sb(table), headers=HEADERS, params=params, timeout=15)
    r.raise_for_status(); return r.json()

def sb_insert(table, rows, ignore_dupes=False, ret=False):
    h = dict(HEADERS)
    if ignore_dupes: h["Prefer"] = "resolution=ignore-duplicates"
    if ret: h["Prefer"] = "return=representation"
    r = requests.post(sb(table), headers=h, json=rows, timeout=30)
    if r.status_code >= 400:
        raise RuntimeError(f"{table} insert {r.status_code}: {r.text[:200]}")
    return r.json() if ret else None

def sb_update(table, params, patch):
    r = requests.patch(sb(table), headers=HEADERS, params=params, json=patch, timeout=15)
    r.raise_for_status()

def sb_delete(table, params):
    r = requests.delete(sb(table), headers=HEADERS, params=params, timeout=15)
    r.raise_for_status()


log = ActivityLog(sb_insert, sb_delete)


def get_config():
    rows = sb_get("hyperion_config_versions", {"active": "eq.true", "select": "id,config", "limit": "1"})
    return (rows[0]["id"], rows[0]["config"]) if rows else (None, {})

def get_control():
    rows = sb_get("hyperion_control", {"id": "eq.1", "select": "*"})
    return rows[0] if rows else {}

def heartbeat(mode):
    sb_update("hyperion_control", {"id": "eq.1"},
              {"engine_heartbeat": datetime.now(timezone.utc).isoformat(), "engine_mode": mode})


# ---------------- market data ----------------
class MarketData:
    def __init__(self):
        self.candles: dict[str, list[dict]] = {}
        self.mids: dict[str, float] = {}

    def refresh_mids(self, dexes: set[str]):
        mids = {}
        try: mids.update(hl.all_mids())
        except Exception as e: log.warn("mids", f"main mids fetch failed: {e}")
        for d in dexes:
            if d:
                try: mids.update(hl.all_mids(d))
                except Exception as e: log.warn("mids", f"mids fetch failed dex={d}: {e}")
        self.mids = {k: float(v) for k, v in mids.items()}

    def refresh_candles(self, coins: list[str]):
        end = int(time.time() * 1000)
        start = end - CANDLE_LOOKBACK_H * 3600_000
        ok = 0
        for coin in coins:
            try:
                cs = hl.candles(coin, CANDLE_INTERVAL, start, end)
                self.candles[coin] = [
                    {"t": int(c["t"]), "o": float(c["o"]), "h": float(c["h"]),
                     "l": float(c["l"]), "c": float(c["c"]), "v": float(c.get("v", 0))}
                    for c in cs]
                ok += 1
            except Exception as e:
                log.warn("candles", f"candles {coin} failed: {e}")
        log.info("candles", f"candles refreshed for {ok}/{len(coins)} pairs ({CANDLE_INTERVAL}, {CANDLE_LOOKBACK_H}h)")


# ---------------- breakers ----------------
class Breakers:
    def __init__(self):
        self.consec_losses = 0
        self.day = None
        self.day_pnl = 0.0
        self.peak_equity = None

    def register_pnls(self, pnls):
        today = datetime.now(timezone.utc).date()
        if self.day != today:
            self.day, self.day_pnl = today, 0.0
        for p in pnls:
            self.day_pnl += p
            self.consec_losses = self.consec_losses + 1 if p < 0 else 0

    def check(self, cfg, equity):
        bk = cfg.get("breakers", {})
        if self.peak_equity is None or equity > self.peak_equity:
            self.peak_equity = equity
        if self.consec_losses >= int(bk.get("consec_losses", 4)):
            return f"consec_losses:{self.consec_losses}"
        if self.day_pnl <= -abs(float(cfg.get("daily_loss_cap_usd", 50))):
            return f"daily_loss:{self.day_pnl:.2f}"
        dd_pct = float(bk.get("drawdown_pct", 10))
        if self.peak_equity and equity < self.peak_equity * (1 - dd_pct / 100):
            return f"drawdown:{(1 - equity / self.peak_equity) * 100:.1f}%"
        return None


def trip_breaker(reason):
    breaker, _, detail = reason.partition(":")
    sb_insert("hyperion_breaker_events", {"breaker": breaker, "detail": {"value": detail}})
    sb_update("hyperion_control", {"id": "eq.1"},
              {"bot_enabled": False, "updated_at": datetime.now(timezone.utc).isoformat()})
    log.error("breaker", f"BREAKER TRIPPED: {reason} — bot disabled", {"reason": reason})


# ---------------- manual commands ----------------
def process_commands(book, md):
    rows = sb_get("hyperion_breaker_events",
                  {"breaker": "eq.manual_command", "resumed_at": "is.null",
                   "select": "id,detail", "order": "id.asc", "limit": "20"})
    for r in rows:
        cmd = r.get("detail") or {}
        done = False
        try:
            if cmd.get("action") == "close":
                done = book.close_coin(cmd.get("coin", ""), md.mids)
            elif cmd.get("action") == "set_sltp":
                done = book.set_sltp(cmd.get("coin", ""), cmd.get("sl"), cmd.get("tp"))
        except Exception as e:
            log.error("cmd", f"command error: {e}", {"cmd": cmd})
        sb_update("hyperion_breaker_events", {"id": f"eq.{r['id']}"},
                  {"resumed_at": datetime.now(timezone.utc).isoformat(),
                   "detail": {**cmd, "executed": done}})
        log.trade("cmd", f"manual {cmd.get('action')} {cmd.get('coin')} -> "
                         f"{'executed' if done else 'no matching position'}", {"cmd": cmd, "executed": done})


def paper_equity(book, md):
    rows = sb_get("hyperion_trades",
                  {"mode": "eq.paper", "status": "eq.closed",
                   "select": "realized_pnl", "limit": "100000"})
    realized = sum(float(r["realized_pnl"] or 0) for r in rows)
    return PAPER_START_BALANCE + realized + book.upnl(md.mids)


def sync_fills(cfg):
    addr = cfg.get("wallet_address")
    if not addr:
        return
    last = sb_get("hyperion_fills", {"select": "ts", "order": "ts.desc", "limit": "1"})
    if last:
        start_ms = int(datetime.fromisoformat(last[0]["ts"].replace("Z", "+00:00")).timestamp() * 1000) + 1
    else:
        start_ms = int(time.time() * 1000) - 30 * 86400_000
    fills = hl.user_fills_paginated(addr, start_ms)
    if not fills:
        return
    rows = [{"hl_tid": str(f.get("tid")), "coin": f.get("coin"),
             "px": float(f.get("px", 0)), "sz": float(f.get("sz", 0)),
             "side": f.get("side", ""), "fee": float(f.get("fee", 0) or 0),
             "closed_pnl": float(f.get("closedPnl", 0) or 0),
             "ts": datetime.fromtimestamp(int(f["time"]) / 1000, tz=timezone.utc).isoformat(),
             "oid": str(f.get("oid", "")), "raw": f} for f in fills]
    sb_insert("hyperion_fills", rows, ignore_dupes=True)
    log.info("fills", f"synced {len(rows)} live wallet fills")


# ---------------- strategy tick ----------------
def strategy_tick(cfg, book, md, brk):
    assets = cfg.get("enabled_assets") or []
    if not assets:
        log.warn("scan", "bot running but no assets enabled in config")
        return

    # 1. manage open positions
    pnls = book.manage(md.mids)
    brk.register_pnls(pnls)

    # 2. breakers
    eq = paper_equity(book, md)
    reason = brk.check(cfg, eq)
    if reason:
        trip_breaker(reason)
        return

    # 3. entry scan — log everything
    max_pos = int(cfg.get("max_concurrent_positions", 3))
    scan_detail = {}
    candidates = []
    for a in assets:
        coin = a["coin"]
        if book.has(coin):
            scan_detail[coin] = {"skip": "position open"}
            continue
        candles = md.candles.get(coin)
        px = md.mids.get(coin)
        if not candles or len(candles) < 60:
            scan_detail[coin] = {"skip": f"insufficient candles ({len(candles or [])}/60)"}
            continue
        if px is None:
            scan_detail[coin] = {"skip": "no mid price"}
            continue
        regime = classify_regime(candles)
        res = best_signal(candles, {"mid": px})
        if res:
            name, sig, _ = res
            scan_detail[coin] = {"regime": regime, "signal": name, "side": sig["side"],
                                 "strength": round(sig["strength"], 3), "reason": sig["reason"]}
            candidates.append((sig["strength"], coin, a.get("dex", ""), name, sig, regime, px))
        else:
            scan_detail[coin] = {"regime": regime, "signal": None}

    n_sig = len(candidates)
    slots = max_pos - book.count()
    log.scan("scan",
             f"scanned {len(assets)} pairs | {n_sig} signal{'s' if n_sig != 1 else ''} | "
             f"{book.count()}/{max_pos} positions | equity {eq:.2f}",
             {"pairs": scan_detail, "slots_free": slots})

    if slots <= 0:
        if n_sig:
            log.info("rank", f"{n_sig} signals but no free slots (max_concurrent={max_pos})")
        return

    if cfg.get("ranking_mode") == "cheapest_first":
        candidates.sort(key=lambda c: c[6])
    else:
        candidates.sort(key=lambda c: -c[0])

    if candidates and len(candidates) > slots:
        skipped = [c[1] for c in candidates[slots:]]
        log.info("rank", f"ranked {len(candidates)} candidates ({cfg.get('ranking_mode','signal_strength')}), "
                         f"taking {slots}, skipping: {', '.join(skipped)}")

    lev = float(cfg.get("default_leverage", 3))
    alloc = float(cfg.get("per_pair_alloc_pct", 25)) / 100
    sl_pct = float(cfg.get("sl_pct", 2)) / 100
    tp_pct = float(cfg.get("tp_pct", 4)) / 100

    for strength, coin, dex, name, sig, regime, px in candidates[:slots]:
        notional = eq * alloc * lev
        size = notional / px
        if sig["side"] == "long":
            sl, tp = px * (1 - sl_pct), px * (1 + tp_pct)
        else:
            sl, tp = px * (1 + sl_pct), px * (1 - tp_pct)
        book.enter(coin, dex, sig["side"], size, px, sl, tp, name, regime, sig["reason"])
        log.trade("enter",
                  f"ENTER {sig['side'].upper()} {coin} sz {size:.6g} @ {px:.6g} "
                  f"| {name}/{regime} | str {strength:.2f} | sl {sl:.6g} tp {tp:.6g} "
                  f"| notional {notional:.2f}",
                  {"coin": coin, "side": sig["side"], "strategy": name, "regime": regime,
                   "strength": strength, "reason": sig["reason"], "notional": notional})


# ---------------- main ----------------
def main():
    net = "MAINNET" if HL_MAINNET else "TESTNET"
    log.info("boot", f"HYPERION ENGINE P2.1 — {net} — tick {TICK_SECONDS}s — paper start {PAPER_START_BALANCE}")
    md = MarketData()
    book = PaperBook(sb_get, sb_insert, sb_update,
                     log=lambda m: log.trade("book", m))
    brk = Breakers()
    log.info("boot", f"reloaded {book.count()} open paper positions from db")
    log.flush()
    tick = 0
    while True:
        t0 = time.time()
        try:
            _, cfg = get_config()
            ctl = get_control()
            mode = cfg.get("mode", "paper")
            heartbeat(mode)

            assets = cfg.get("enabled_assets") or []
            dexes = {a.get("dex", "") for a in assets}
            coins = [a["coin"] for a in assets]

            if ctl.get("kill_switch"):
                md.refresh_mids(dexes)
                closed = 0
                for tid in list(book.open.keys()):
                    r = book.open[tid]
                    if r["coin"] in md.mids:
                        book.close(tid, md.mids[r["coin"]], "kill_switch")
                        closed += 1
                if closed:
                    log.error("kill", f"KILL SWITCH: flattened {closed} paper positions")
                elif tick % 10 == 0:
                    log.warn("kill", "kill switch engaged — engine holding flat")
            else:
                md.refresh_mids(dexes)
                if tick % CANDLE_REFRESH_EVERY == 0 and coins:
                    md.refresh_candles(coins)
                process_commands(book, md)

                if mode == "paper" and ctl.get("bot_enabled"):
                    strategy_tick(cfg, book, md, brk)
                elif tick % 10 == 0:
                    why = "bot disabled in config" if mode == "paper" else f"mode={mode} (live executor is Phase 3)"
                    log.info("idle", f"idle — {why} | {len(md.mids)} mids | {book.count()} positions held")

                if mode == "paper":
                    eq = paper_equity(book, md)
                    sb_insert("hyperion_equity", {"mode": "paper", "account_value": eq,
                                                  "margin_used": 0, "upnl": book.upnl(md.mids)})
                elif cfg.get("wallet_address"):
                    s = hl.account_summary(cfg["wallet_address"])
                    sb_insert("hyperion_equity", {"mode": "live",
                                                  "account_value": s["total_account_value"],
                                                  "margin_used": s["total_margin_used"],
                                                  "upnl": s["total_upnl"]})
                if tick % FILLS_SYNC_EVERY == 0:
                    sync_fills(cfg)
        except KeyboardInterrupt:
            log.info("boot", "engine stopped by user"); log.flush(); return
        except Exception as e:
            log.error("err", f"tick error: {e}", {"traceback": traceback.format_exc()[-1500:]})
        log.flush()
        tick += 1
        time.sleep(max(1, TICK_SECONDS - (time.time() - t0)))


if __name__ == "__main__":
    main()
