"""
strategy.py — evaluate a symbol against all filters in rules.json.
All market data comes from IBKR (daily bars + live ticker).
"""
import json
import math
from datetime import datetime, time
from pathlib import Path
from zoneinfo import ZoneInfo

ET    = ZoneInfo("America/New_York")
RULES = json.loads(Path("rules.json").read_text())


def _safe(val) -> float:
    if val is None:
        return float("nan")
    try:
        f = float(val)
        return f if not math.isnan(f) else float("nan")
    except Exception:
        return float("nan")


def _ok(val) -> bool:
    v = _safe(val)
    return not math.isnan(v) and v > 0


def evaluate(symbol: str, ib) -> dict:
    reasons_fail = []

    # ── 1. Position dedupe ──────────────────────────────────────────────────
    for pos in ib.positions():
        if pos.contract.symbol == symbol and pos.position > 0:
            return {"pass": False, "reasons": ["already_in_position"], "price": 0.0}

    # ── 2. Time window ──────────────────────────────────────────────────────
    now_et   = datetime.now(ET).time()
    earliest = time.fromisoformat(RULES["time_filter"]["earliest_entry_et"])
    latest   = time.fromisoformat(RULES["time_filter"]["latest_entry_et"])
    if not (earliest <= now_et <= latest):
        return {"pass": False, "reasons": ["outside_entry_window"], "price": 0.0}

    # ── 3. Historical daily bars (IBKR) ─────────────────────────────────────
    from src.ibkr_client import IBKRClient as _C
    # ib is the raw IB() object passed in from bot.py; wrap it temporarily
    # by importing the helper directly
    from ib_async import Stock as _Stock
    _contract = _Stock(symbol, "SMART", "USD")
    ib.qualifyContracts(_contract)

    daily = ib.reqHistoricalData(
        _contract,
        endDateTime="",
        durationStr="205 D",
        barSizeSetting="1 day",
        whatToShow="TRADES",
        useRTH=True,
        formatDate=1,
    )
    if not daily or len(daily) < 2:
        return {"pass": False, "reasons": ["ibkr_no_daily_history"], "price": 0.0}

    prior_close = _safe(daily[-2].close)
    prior_high  = _safe(daily[-2].high)
    today_open  = _safe(daily[-1].open)

    closes  = [b.close for b in daily if b.close and b.close > 0]
    sma200  = float(sum(closes[-200:]) / len(closes[-200:])) if len(closes) >= 200 else float("nan")

    rvol_days = RULES["intraday_filters"].get("I3_rvol_lookback_days", 14)
    rvol_bars = daily[-(rvol_days + 1):]
    avg_vol   = float(sum(b.volume for b in rvol_bars[:-1]) / max(len(rvol_bars) - 1, 1))
    today_vol = float(rvol_bars[-1].volume) if rvol_bars else 0.0
    rvol      = (today_vol / avg_vol) if avg_vol > 0 else 0.0

    # ── 4. Live price + intraday HOD/LOD (IBKR) ────────────────────────────
    _ticker = ib.reqMktData(_contract, "", False, False)
    ib.sleep(3)
    price = _safe(_ticker.last) if _ok(_ticker.last) else _safe(_ticker.close)
    hod   = _safe(_ticker.high)
    lod   = _safe(_ticker.low)
    ib.cancelMktData(_contract)

    # Fallback: 1m bars for HOD/LOD if live tick not available
    if not _ok(hod) or not _ok(lod):
        intra_bars = ib.reqHistoricalData(
            _contract,
            endDateTime="",
            durationStr="1 D",
            barSizeSetting="1 min",
            whatToShow="TRADES",
            useRTH=False,
            formatDate=1,
        )
        if intra_bars:
            hod = float(max(b.high  for b in intra_bars))
            lod = float(min(b.low   for b in intra_bars))
            if not _ok(price):
                price = _safe(intra_bars[-1].close)

    if not _ok(price):
        return {"pass": False, "reasons": ["price_unavailable"], "price": 0.0}

    # ── 5. Price range filter ───────────────────────────────────────────────
    min_p = RULES["universe_filters"].get("min_price_usd", 1.0)
    max_p = RULES["universe_filters"].get("max_price_usd", 20.0)
    if not (min_p <= price <= max_p):
        return {"pass": False, "reasons": [f"price_out_of_range_{price:.2f}"], "price": price}

    # ── 6. Daily filters ────────────────────────────────────────────────────
    if RULES["daily_filters"].get("D1_above_prior_day_high") and _ok(prior_high):
        if price <= prior_high:
            reasons_fail.append("D1_not_above_prior_high")

    if RULES["daily_filters"].get("D2_prior_close_above_sma200") and _ok(sma200):
        if prior_close <= sma200:
            reasons_fail.append("D2_prior_close_below_sma200")

    min_gap = RULES["daily_filters"].get("D3_min_gap_pct_from_prior_close", 3.0)
    if _ok(prior_close) and _ok(today_open):
        gap_pct = (today_open - prior_close) / prior_close * 100
        if gap_pct < min_gap:
            reasons_fail.append(f"D3_gap_{gap_pct:.1f}pct_lt_{min_gap}pct")
    else:
        reasons_fail.append("D3_gap_data_missing")

    # ── 7. Intraday filters ─────────────────────────────────────────────────
    if RULES["intraday_filters"].get("I1_above_premarket_high") and _ok(today_open):
        if price <= today_open:
            reasons_fail.append("I1_not_above_premarket_high")

    if RULES["intraday_filters"].get("I2_above_today_hod") and _ok(hod):
        if price < hod:
            reasons_fail.append("I2_not_at_hod")

    min_rvol = RULES["intraday_filters"].get("I3_rvol_min", 2.0)
    if rvol < min_rvol:
        reasons_fail.append(f"I3_rvol_{rvol:.2f}x_lt_{min_rvol}x")

    if reasons_fail:
        return {"pass": False, "reasons": reasons_fail, "price": price}

    reasons_pass = [
        "D1_above_prior_high",
        "D2_prior_close_above_sma200",
        "D3_gap_ok",
        "I1_above_premarket_high",
        "I2_at_hod",
        f"I3_rvol_{rvol:.1f}x",
        "time_ok",
        "no_existing_position",
    ]

    log_dir = Path("logs")
    log_dir.mkdir(exist_ok=True)
    with open(log_dir / "safety_log.jsonl", "a") as f:
        f.write(json.dumps({
            "symbol":      symbol,
            "pass":        True,
            "price":       price,
            "hod":         hod   if _ok(hod)   else None,
            "lod":         lod   if _ok(lod)   else None,
            "rvol":        round(rvol, 2),
            "sma200":      round(sma200, 2) if _ok(sma200) else None,
            "prior_close": prior_close,
            "prior_high":  prior_high,
        }) + "\n")

    return {
        "pass":    True,
        "reasons": reasons_pass,
        "price":   float(price),
        "lod":     float(lod) if _ok(lod) else float(price) * 0.99,
        "hod":     float(hod) if _ok(hod) else float(price),
        "rvol":    round(rvol, 2),
    }
