"""
regime_detector.py

[SIMULATION-ONLY 2026-09-04] Classifies a symbol's session into one of
seven shape-of-the-day "regimes" so entry/exit strategies can eventually
be tuned per-regime instead of the single confirmation rule set
entry_engine.py currently applies to every candidate. NOT yet wired into
monitor.py / entry_engine.py / intraday_health.py -- see config.json's
regime_detection._note. For now this is consumed only by
simulate_regimes.py against data/snapshots/<date>/.

Regimes (best-effort description, not mutually exclusive in theory but
classify_regime() picks exactly one per session via priority order):

    REGIME_1_NORMAL_MOMENTUM
        No unusual gap. Orderly trend for the session (net direction
        matches the regression fit, no failed breakout, no big dip).

    REGIME_2_GAP_MOMENTUM
        Opened with a real gap vs. prior close, and the session
        continued trending in the gap's direction rather than fading it.

    REGIME_3_CATALYST_GAP_HIGH_VOLATILITY
        Large gap AND a wide intraday range -- the profile of a
        news-driven name, whichever direction it nets out.

    REGIME_4_RANGE_CHOP
        Small session range, no persistent direction. The whipsaw
        pattern several of 2026-09-04's real losers (ONDS/NOK/SMR/RDW)
        showed: price reversed almost immediately after any entry.

    REGIME_5_FAILED_BREAKOUT
        Price broke and held above the OPENING-RANGE high at some point
        (a resistance level established over breakout_lookback_bars),
        then gave most or all of that breakout back by session end --
        RDW's 2026-09-04 pattern (broke $10.64, gave back 195% of the
        move). Requires an actual defined resistance level to have been
        broken; see REGIME_7 below for the faster, gap-less spike case
        this definition structurally cannot see.

    REGIME_6_RECOVERY_RECLAIM
        A meaningful early dip from the open, followed by reclaiming and
        holding back above the open -- a V-shape or higher-low-after-
        weakness pattern.

    REGIME_7_INTRADAY_SPIKE_NO_GAP_WHIPSAW
        No real overnight gap, but a violent spike within the first few
        minutes of the open that mostly reverses shortly after -- CHPT's
        actual 2026-09-04 pattern (opened flat, spiked from $9.28 to
        $10.20 by minute 2, the bot entered at $9.83 already riding the
        reversal down, price kept falling to $9.32). This is DIFFERENT
        from REGIME_5: the spike happens too fast and too early to ever
        get treated as "the opening range" in the first place -- with
        breakout_lookback_bars=5, CHPT's minute-2 spike is INSIDE the
        window used to compute the opening-range high, so it's baked
        into the resistance level itself instead of breaking it. This
        regime checks the raw first N minutes directly instead of
        relying on a resistance break/fail.

Pure functions only -- no state, no file I/O, no config mutation. Bars use
the same {"t","o","h","l","c","v"} shape as indicators.py, with "t" a
tz-aware datetime (so session-window filtering works regardless of
whether pre/post-market bars are mixed in).
"""

from dataclasses import dataclass, field
from datetime import datetime, time, timedelta

from indicators import (
    atr, breakout_confirmed, regression_fit_pct_change,
)

REGIME_NORMAL_MOMENTUM = "REGIME_1_NORMAL_MOMENTUM"
REGIME_GAP_MOMENTUM = "REGIME_2_GAP_MOMENTUM"
REGIME_CATALYST_GAP_HIGH_VOL = "REGIME_3_CATALYST_GAP_HIGH_VOLATILITY"
REGIME_RANGE_CHOP = "REGIME_4_RANGE_CHOP"
REGIME_FAILED_BREAKOUT = "REGIME_5_FAILED_BREAKOUT"
REGIME_RECOVERY_RECLAIM = "REGIME_6_RECOVERY_RECLAIM"
REGIME_WHIPSAW = "REGIME_7_INTRADAY_SPIKE_NO_GAP_WHIPSAW"

ALL_REGIMES = [
    REGIME_NORMAL_MOMENTUM, REGIME_GAP_MOMENTUM, REGIME_CATALYST_GAP_HIGH_VOL,
    REGIME_RANGE_CHOP, REGIME_FAILED_BREAKOUT, REGIME_RECOVERY_RECLAIM,
    REGIME_WHIPSAW,
]

DEFAULT_CONFIG = {
    "session_start": "09:30",
    "session_end": "16:00",
    "gap_momentum_min_pct": 3.0,
    "catalyst_gap_min_pct": 8.0,
    "catalyst_range_min_pct": 15.0,
    "chop_max_range_pct": 6.0,
    "chop_max_net_change_pct": 2.0,
    "breakout_lookback_bars": 5,
    "breakout_buffer_pct": 0.3,
    "breakout_hold_bars": 2,
    "failed_breakout_retrace_pct": 50.0,
    "reclaim_min_dip_pct": 3.0,
    "reclaim_hold_bars": 5,
    "atr_period": 14,
    "min_bars_required": 10,
    "whipsaw_spike_lookback_bars": 10,
    "whipsaw_min_spike_pct": 4.0,
    "whipsaw_reversal_window_bars": 10,
    "whipsaw_min_retrace_pct": 60.0,
}


@dataclass
class RegimeReading:
    symbol: str
    regime: str
    confidence: str          # "high" / "medium" / "low" -- see _confidence()
    gap_pct: float
    range_pct: float
    net_change_pct: float
    trend_fit_pct: float
    atr_pct: float
    breakout: dict = field(default_factory=dict)
    reclaim: dict = field(default_factory=dict)
    whipsaw: dict = field(default_factory=dict)
    insufficient_data: bool = False


def _parse_hhmm(s: str) -> time:
    h, m = s.split(":")
    return time(int(h), int(m))


def session_slice(bars: list, session_start: str, session_end: str) -> list:
    """Bars whose timestamp falls within [session_start, session_end] wall
    time, inclusive -- drops pre/post-market bars so gap/range/breakout
    math is computed on the regular session only. Falls back to the full
    bar list if nothing falls in that window (e.g. a partial/odd feed)
    rather than returning empty and forcing every caller to guard against
    that themselves."""
    start_t = _parse_hhmm(session_start)
    end_t = _parse_hhmm(session_end)
    filtered = [b for b in bars if start_t <= b["t"].time() <= end_t]
    return filtered if filtered else list(bars)


def opening_window(bars: list, session_start: str = "09:30", window_minutes: int = 90) -> list:
    """
    Restricts bars to [session_start, session_start + window_minutes] --
    the decision-relevant window for a bot that only ever enters in the
    opening range (2026-09-04's real trades: 12 of 13 entries landed in
    the first 27 minutes, the 13th at +61 min). Classifying on the full
    6.5-hour session instead of this window hides exactly the pattern
    that matters: e.g. CHPT round-tripped a spike-to-$10.20-and-dump
    within its first 4 minutes but closed the full day net positive, so
    a full-session classification reads it as ordinary uptrend
    momentum and completely misses the failed-breakout shape the bot's
    actual entry got caught in. Use this to slice bars BEFORE calling
    classify_regime() when the question is "what regime was this
    symbol in when the bot could have entered," as opposed to
    "what kind of day did this symbol have overall."
    """
    start_t = _parse_hhmm(session_start)
    end_t = (datetime.combine(datetime.min, start_t) + timedelta(minutes=window_minutes)).time()
    filtered = [b for b in bars if start_t <= b["t"].time() <= end_t]
    return filtered if filtered else list(bars)


def _detect_breakout_and_failure(session_bars: list, cfg: dict) -> dict:
    """
    Did the session break above its own opening-range high at some point,
    and if so, did it fail (give most of that breakout back by the last
    bar)? See breakout_confirmed() in indicators.py -- same "held for N
    bars, not just spiked through" definition used by the live confirmation
    logic, applied here retrospectively over the whole session instead of
    at a single instant.
    """
    lookback = min(cfg["breakout_lookback_bars"], max(1, len(session_bars) // 3))
    hold_bars = cfg["breakout_hold_bars"]
    if len(session_bars) < lookback + hold_bars + 1:
        return {"broke_out": False, "failed": False, "opening_range_high": None}

    opening_range_high = max(b["h"] for b in session_bars[:lookback])
    rest = session_bars[lookback:]

    breakout_idx = None
    for i in range(len(rest) - hold_bars + 1):
        window = rest[i:i + hold_bars]
        if breakout_confirmed(window, opening_range_high, cfg["breakout_buffer_pct"], hold_bars):
            breakout_idx = i
            break

    if breakout_idx is None:
        return {"broke_out": False, "failed": False, "opening_range_high": round(opening_range_high, 4)}

    after_breakout = rest[breakout_idx:]
    breakout_high = max(b["h"] for b in after_breakout)
    final_close = session_bars[-1]["c"]

    move = breakout_high - opening_range_high
    retrace_pct = ((breakout_high - final_close) / move * 100.0) if move > 0 else 0.0
    failed = final_close < opening_range_high and retrace_pct >= cfg["failed_breakout_retrace_pct"]

    return {
        "broke_out": True,
        "failed": failed,
        "opening_range_high": round(opening_range_high, 4),
        "breakout_high": round(breakout_high, 4),
        "retrace_pct": round(retrace_pct, 1),
    }


def _detect_reclaim(session_bars: list, cfg: dict) -> dict:
    """
    A meaningful dip from the open followed by reclaiming and HOLDING
    back above the open into the close (checked over the last
    reclaim_hold_bars closes, not just a single tag-and-go print).
    """
    open_price = session_bars[0]["o"]
    if open_price <= 0:
        return {"detected": False, "dip_pct": 0.0}

    low_idx = min(range(len(session_bars)), key=lambda i: session_bars[i]["l"])
    day_low = session_bars[low_idx]["l"]
    dip_pct = (open_price - day_low) / open_price * 100.0

    if dip_pct < cfg["reclaim_min_dip_pct"]:
        return {"detected": False, "dip_pct": round(dip_pct, 2)}

    after_low = session_bars[low_idx:]
    hold_bars = cfg["reclaim_hold_bars"]
    if len(after_low) < hold_bars:
        return {"detected": False, "dip_pct": round(dip_pct, 2), "day_low": round(day_low, 4)}

    reclaimed = all(b["c"] >= open_price for b in after_low[-hold_bars:])
    return {
        "detected": reclaimed,
        "dip_pct": round(dip_pct, 2),
        "day_low": round(day_low, 4),
    }


def _detect_whipsaw(session_bars: list, cfg: dict) -> dict:
    """
    A violent spike within the first whipsaw_spike_lookback_bars minutes
    of the open, followed by giving back at least whipsaw_min_retrace_pct
    of that spike within the next whipsaw_reversal_window_bars -- checked
    directly against the open, NOT against a resistance/breakout level
    (see REGIME_7's docstring for why _detect_breakout_and_failure()
    structurally cannot see this pattern when the spike happens inside
    its own opening-range reference window).
    """
    open_price = session_bars[0]["o"]
    if open_price <= 0:
        return {"detected": False, "spike_pct": 0.0, "retrace_pct": 0.0}

    spike_window = session_bars[:cfg["whipsaw_spike_lookback_bars"]]
    peak_idx = max(range(len(spike_window)), key=lambda i: spike_window[i]["h"])
    peak_price = spike_window[peak_idx]["h"]
    spike_pct = (peak_price - open_price) / open_price * 100.0

    if spike_pct < cfg["whipsaw_min_spike_pct"]:
        return {"detected": False, "spike_pct": round(spike_pct, 2), "retrace_pct": 0.0}

    reversal_window = session_bars[peak_idx:peak_idx + cfg["whipsaw_reversal_window_bars"]]
    if len(reversal_window) < 2:
        return {"detected": False, "spike_pct": round(spike_pct, 2), "retrace_pct": 0.0}

    reversal_low = min(b["l"] for b in reversal_window)
    move = peak_price - open_price
    retrace_pct = (peak_price - reversal_low) / move * 100.0 if move > 0 else 0.0

    detected = retrace_pct >= cfg["whipsaw_min_retrace_pct"]
    return {
        "detected": detected,
        "spike_pct": round(spike_pct, 2),
        "retrace_pct": round(retrace_pct, 1),
        "peak_price": round(peak_price, 4),
        "reversal_low": round(reversal_low, 4),
    }


def _confidence(gap_pct: float, range_pct: float, trend_fit_pct: float, cfg: dict) -> str:
    """Rough distance-from-threshold heuristic, not a statistical measure
    -- just flags cases that landed near a boundary so a human reviewing
    simulation output knows which labels to double-check by eye rather
    than trust blindly."""
    margins = [
        abs(abs(gap_pct) - cfg["gap_momentum_min_pct"]),
        abs(abs(gap_pct) - cfg["catalyst_gap_min_pct"]),
        abs(range_pct - cfg["chop_max_range_pct"]),
        abs(range_pct - cfg["catalyst_range_min_pct"]),
    ]
    closest = min(margins)
    if closest < 0.5:
        return "low"
    if closest < 1.5:
        return "medium"
    return "high"


def classify_regime(symbol: str, bars: list, prev_close: float, cfg: dict = None) -> RegimeReading:
    """
    Pure evaluation -- given one symbol's bars (any mix of pre/regular/
    post-market, any length) and its previous session's close, returns a
    single regime label for the session plus the metrics that drove it.

    Can be called with a partial day's bars (e.g. bars up to "now") for
    a live/rolling classification later; simulate_regimes.py currently
    calls it once per symbol with the full day's bars.
    """
    cfg = {**DEFAULT_CONFIG, **(cfg or {})}
    session_bars = session_slice(bars, cfg["session_start"], cfg["session_end"])

    if len(session_bars) < cfg["min_bars_required"] or not prev_close:
        return RegimeReading(
            symbol=symbol, regime=REGIME_RANGE_CHOP, confidence="low",
            gap_pct=0.0, range_pct=0.0, net_change_pct=0.0, trend_fit_pct=0.0,
            atr_pct=0.0, insufficient_data=True,
        )

    open_price = session_bars[0]["o"]
    final_close = session_bars[-1]["c"]
    day_high = max(b["h"] for b in session_bars)
    day_low = min(b["l"] for b in session_bars)
    closes = [b["c"] for b in session_bars]

    gap_pct = (open_price - prev_close) / prev_close * 100.0
    range_pct = (day_high - day_low) / prev_close * 100.0
    net_change_pct = (final_close - open_price) / open_price * 100.0 if open_price else 0.0
    trend_fit_pct = regression_fit_pct_change(closes)
    a = atr(session_bars, period=min(cfg["atr_period"], max(2, len(session_bars) - 1)))
    atr_pct = (a / final_close * 100.0) if final_close else 0.0

    breakout = _detect_breakout_and_failure(session_bars, cfg)
    reclaim = _detect_reclaim(session_bars, cfg)
    whipsaw = _detect_whipsaw(session_bars, cfg)

    same_direction_continuation = (
        (gap_pct > 0 and net_change_pct >= 0 and trend_fit_pct >= 0)
        or (gap_pct < 0 and net_change_pct <= 0 and trend_fit_pct <= 0)
    )
    no_real_gap = abs(gap_pct) < cfg["gap_momentum_min_pct"]

    if abs(gap_pct) >= cfg["catalyst_gap_min_pct"] and range_pct >= cfg["catalyst_range_min_pct"]:
        regime = REGIME_CATALYST_GAP_HIGH_VOL
    elif no_real_gap and whipsaw.get("detected"):
        regime = REGIME_WHIPSAW
    elif breakout.get("failed"):
        regime = REGIME_FAILED_BREAKOUT
    elif abs(gap_pct) >= cfg["gap_momentum_min_pct"] and same_direction_continuation:
        regime = REGIME_GAP_MOMENTUM
    elif reclaim.get("detected"):
        regime = REGIME_RECOVERY_RECLAIM
    elif range_pct <= cfg["chop_max_range_pct"] and abs(net_change_pct) <= cfg["chop_max_net_change_pct"]:
        regime = REGIME_RANGE_CHOP
    elif abs(trend_fit_pct) > 0:
        regime = REGIME_NORMAL_MOMENTUM
    else:
        regime = REGIME_RANGE_CHOP

    return RegimeReading(
        symbol=symbol,
        regime=regime,
        confidence=_confidence(gap_pct, range_pct, trend_fit_pct, cfg),
        gap_pct=round(gap_pct, 2),
        range_pct=round(range_pct, 2),
        net_change_pct=round(net_change_pct, 2),
        trend_fit_pct=round(trend_fit_pct, 2),
        atr_pct=round(atr_pct, 2),
        breakout=breakout,
        reclaim=reclaim,
        whipsaw=whipsaw,
    )


def bars_from_price_rows(rows: list) -> list:
    """
    Adapter: turns the {timestamp_et, open, high, low, close, volume}
    rows written by data/snapshots/<date>/prices/<SYMBOL>.csv into the
    {"t","o","h","l","c","v"} bar dicts every function in this module
    (and indicators.py) expects. `timestamp_et` values are naive
    "YYYY-MM-DD HH:MM:SS" strings already in ET wall time (see
    build_snapshot.py) -- kept naive here since session_slice() only
    ever compares wall-clock time-of-day, not absolute instants.
    """
    bars = []
    for row in rows:
        bars.append({
            "t": datetime.strptime(row["timestamp_et"], "%Y-%m-%d %H:%M:%S"),
            "o": float(row["open"]),
            "h": float(row["high"]),
            "l": float(row["low"]),
            "c": float(row["close"]),
            "v": float(row["volume"]) if row["volume"] not in ("", None) else 0.0,
        })
    return bars
