"""
relative_strength_engine.py

[Phase 3 -- new] Compares a candidate's own return against the broad market
(SPY/QQQ) and, where a sector/theme mapping exists, against its peer group --
per the project's explicit requirement that an isolated "ticker is green" is
weaker evidence than the same move confirmed by sector-wide participation
(SMR strongly positive + OKLO/URA/nuclear peers positive + SPY/QQQ flat is a
materially stronger signal than SMR green in isolation).

Pure function, no state, no network calls. This module is handed bars for
the symbol AND for SPY/QQQ/peers by the CALLER -- monitor.py owns whatever
market-data subscriptions are needed to get those bars, the same separation
of concerns every other engine here already follows (structure_engine/
trend_engine don't fetch their own bars either, they're handed them).

SECTOR/PEER MAPPING IS A STATIC, CONFIGURED MAP
(config.relative_strength_engine.sector_map) -- there is no dynamic sector-
classification data source anywhere in this codebase, and fabricating one
would violate this project's explicit no-manufactured-data rule (see
prediction_engine.py's own relative_strength docstring for the same
principle applied to the individual-symbol case). A symbol with no entry in
the map simply skips the sector layer -- rs_score vs SPY/QQQ still applies
in full. Per the project's explicit instruction, sector confirmation is an
ADDITIONAL confidence multiplier, never a requirement: an unmapped symbol,
or one whose peers didn't confirm, is not penalized for it, only left
without the bonus.
"""

from dataclasses import dataclass, field

from config_loader import get_config

DEFAULT_CONFIG = {
    "windows": {"1m": 1, "3m": 3, "5m": 5, "15m": 15},
    "window_weights": {"1m": 0.15, "3m": 0.25, "5m": 0.30, "15m": 0.20, "premarket": 0.10},
    "rs_scale_pct": 2.0,  # +/- this many pct of outperformance vs market maps to a full 0/100 score
    "sector_breadth_min_for_bonus": 0.6,
    "sector_confirmation_bonus": 0.15,
    "sector_peer_window": "5m",
    "sector_map": {},
    # example shape, populated per-symbol in config.json:
    # "sector_map": {"SMR": {"sector": "nuclear", "peers": ["OKLO", "LEU", "CCJ"], "sector_etf": "URA"}}
    "market_regime_window": "5m",
    "market_regime_thresholds": {"strong": 0.5, "mild": 0.15},  # composite SPY/QQQ %-return breakpoints
}

STRONG_RISK_ON = "STRONG_RISK_ON"
RISK_ON = "RISK_ON"
NEUTRAL = "NEUTRAL"
RISK_OFF = "RISK_OFF"
STRONG_RISK_OFF = "STRONG_RISK_OFF"


@dataclass
class MarketRegimeReading:
    regime: str = NEUTRAL
    spy_return_pct: float = None
    qqq_return_pct: float = None
    composite_return_pct: float = None
    insufficient_data: bool = False


def classify_market_regime(spy_bars: list, qqq_bars: list, cfg: dict = None) -> MarketRegimeReading:
    """
    Broad-market context (section 19): STRONG_RISK_ON..STRONG_RISK_OFF from
    SPY/QQQ's own recent composite return. Deliberately NOT a hard filter
    anywhere downstream -- per the project's explicit instruction, a strong
    symbol outperforming a weak market should score BETTER for its relative
    strength, not be rejected for the market being weak. See winner_score.py
    for how this is actually used (a small, non-veto weight).
    """
    cfg = {**DEFAULT_CONFIG, **(cfg or get_config().get("relative_strength_engine", {}))}
    n = cfg["windows"].get(cfg["market_regime_window"], 5)
    spy_ret = _window_return_pct(spy_bars, n)
    qqq_ret = _window_return_pct(qqq_bars, n)
    if spy_ret is None or qqq_ret is None:
        return MarketRegimeReading(insufficient_data=True)

    composite = (spy_ret + qqq_ret) / 2.0
    t = cfg["market_regime_thresholds"]
    if composite >= t["strong"]:
        regime = STRONG_RISK_ON
    elif composite >= t["mild"]:
        regime = RISK_ON
    elif composite <= -t["strong"]:
        regime = STRONG_RISK_OFF
    elif composite <= -t["mild"]:
        regime = RISK_OFF
    else:
        regime = NEUTRAL

    return MarketRegimeReading(regime=regime, spy_return_pct=round(spy_ret, 3),
                                qqq_return_pct=round(qqq_ret, 3), composite_return_pct=round(composite, 3))


@dataclass
class RelativeStrengthReading:
    symbol: str
    rs_score: float = 0.0
    final_score: float = 0.0
    rs_vs_spy: dict = field(default_factory=dict)
    rs_vs_qqq: dict = field(default_factory=dict)
    sector: str = None
    sector_breadth: float = None
    sector_momentum: float = None
    sector_volume_confirmation: float = None
    insufficient_data: bool = False
    missing: list = field(default_factory=list)


def _clamp01(x: float) -> float:
    return max(0.0, min(1.0, x))


def _window_return_pct(bars: list, n_bars: int):
    if not bars or len(bars) <= n_bars:
        return None
    now, then = bars[-1]["c"], bars[-1 - n_bars]["c"]
    if not then:
        return None
    return (now - then) / then * 100.0


def _volume_confirming(bars: list) -> bool:
    """Same recent-half-vs-early-half shape as indicators.volume_acceleration,
    reused here rather than re-derived under a different name -- a peer
    "confirms" when its own volume is not fading."""
    vols = [b["v"] for b in bars[-10:]] if len(bars) >= 10 else [b["v"] for b in bars]
    if len(vols) < 4:
        return None
    mid = len(vols) // 2
    early = sum(vols[:mid]) / mid or 1e-9
    recent = sum(vols[mid:]) / (len(vols) - mid)
    return (recent / early) >= 1.0


def _sector_layer(symbol: str, cfg: dict, peer_bars: dict):
    entry = cfg.get("sector_map", {}).get(symbol)
    if not entry or not peer_bars:
        return None, None, None, None

    sector = entry.get("sector")
    n_bars = cfg["windows"].get(cfg["sector_peer_window"], 5)
    peer_returns, peer_vol_flags = [], []
    for peer, bars in peer_bars.items():
        ret = _window_return_pct(bars, n_bars)
        if ret is not None:
            peer_returns.append(ret)
        vol_ok = _volume_confirming(bars)
        if vol_ok is not None:
            peer_vol_flags.append(vol_ok)

    if not peer_returns:
        return sector, None, None, None

    breadth = round(sum(1 for r in peer_returns if r > 0) / len(peer_returns), 3)
    momentum = round(sum(peer_returns) / len(peer_returns), 3)
    vol_confirmation = round(sum(peer_vol_flags) / len(peer_vol_flags), 3) if peer_vol_flags else None
    return sector, breadth, momentum, vol_confirmation


def compute_relative_strength(symbol: str, bars: list, spy_bars: list, qqq_bars: list,
                               premarket_return_pct: float = None,
                               premarket_spy_return_pct: float = None,
                               premarket_qqq_return_pct: float = None,
                               peer_bars: dict = None, cfg: dict = None) -> RelativeStrengthReading:
    """
    bars / spy_bars / qqq_bars: regular-session 1-min bars, oldest first, for
        the candidate / SPY / QQQ respectively -- same shape stream.py hands
        every other engine.
    premarket_*_return_pct: optional plain floats (this symbol's own, SPY's,
        QQQ's premarket % change) for the "premarket" window -- plain floats
        rather than raw bars because premarket_scanner.py already computes
        the candidate's own premarket move (see scorer.price_momentum);
        recomputing that from raw bars a second time here would duplicate
        it. Omit any of the three to skip this window; its weight is simply
        redistributed across the others.
    peer_bars: optional {peer_symbol: bars} for this symbol's sector map
        entry's `peers` list. Omit entirely to skip the sector layer -- rs
        vs SPY/QQQ is unaffected either way.
    """
    cfg = {**DEFAULT_CONFIG, **(cfg or get_config().get("relative_strength_engine", {}))}
    weights = dict(cfg["window_weights"])

    if len(bars) < 2 or len(spy_bars) < 2 or len(qqq_bars) < 2:
        return RelativeStrengthReading(symbol=symbol, insufficient_data=True,
                                        missing=["need at least 2 bars each for symbol/SPY/QQQ"])

    rs_vs_spy, rs_vs_qqq, window_scores = {}, {}, {}
    for label, n in cfg["windows"].items():
        sym_ret = _window_return_pct(bars, n)
        spy_ret = _window_return_pct(spy_bars, n)
        qqq_ret = _window_return_pct(qqq_bars, n)
        if sym_ret is None or spy_ret is None or qqq_ret is None:
            weights.pop(label, None)
            continue
        rs_vs_spy[label] = round(sym_ret - spy_ret, 3)
        rs_vs_qqq[label] = round(sym_ret - qqq_ret, 3)
        avg_vs_market = (rs_vs_spy[label] + rs_vs_qqq[label]) / 2.0
        window_scores[label] = _clamp01((avg_vs_market + cfg["rs_scale_pct"]) / (2 * cfg["rs_scale_pct"])) * 100.0

    have_premarket = None not in (premarket_return_pct, premarket_spy_return_pct, premarket_qqq_return_pct)
    if have_premarket:
        rs_vs_spy["premarket"] = round(premarket_return_pct - premarket_spy_return_pct, 3)
        rs_vs_qqq["premarket"] = round(premarket_return_pct - premarket_qqq_return_pct, 3)
        avg_vs_market = (rs_vs_spy["premarket"] + rs_vs_qqq["premarket"]) / 2.0
        window_scores["premarket"] = _clamp01(
            (avg_vs_market + cfg["rs_scale_pct"]) / (2 * cfg["rs_scale_pct"])) * 100.0
    else:
        weights.pop("premarket", None)

    if not window_scores:
        return RelativeStrengthReading(symbol=symbol, insufficient_data=True,
                                        missing=["no window had enough bars for a return comparison"])

    weight_total = sum(weights.get(w, 0) for w in window_scores)
    rs_score = round(sum(window_scores[w] * weights.get(w, 0) for w in window_scores) / weight_total, 2) \
        if weight_total else 0.0

    sector, breadth, momentum, vol_conf = _sector_layer(symbol, cfg, peer_bars or {})
    final_score = rs_score
    if breadth is not None and breadth >= cfg["sector_breadth_min_for_bonus"]:
        final_score = min(100.0, rs_score * (1 + cfg["sector_confirmation_bonus"]))

    return RelativeStrengthReading(
        symbol=symbol, rs_score=rs_score, final_score=round(final_score, 2),
        rs_vs_spy=rs_vs_spy, rs_vs_qqq=rs_vs_qqq,
        sector=sector, sector_breadth=breadth, sector_momentum=momentum,
        sector_volume_confirmation=vol_conf, insufficient_data=False,
    )
