"""
prediction_engine.py

[SIMULATION-ONLY 2026-09-05] Combines everything built so far
(stream_features, trend_engine, structure_engine, optionally
regime_engine + setup_engine) into a single 0-100 PREDICTION_SCORE, a
prediction_slope (improving/deteriorating vs the caller-held prior
reading), and a deterministic next-move probability distribution over
five states (CONTINUATION_UP, PULLBACK_THEN_UP, SIDEWAYS, PULLBACK_DOWN,
REVERSAL_DOWN) for 1m/3m/5m horizons. Not consumed anywhere in the live
bot yet -- see config.json's prediction_engine._note.

DETERMINISTIC, NOT ML -- per the project's original design brief's
explicit instruction ("first build deterministic prediction engine...
then the historical dataset can later be used to train/calibrate").
The next-move probabilities are a documented, inspectable heuristic
(see _next_move_distribution()) built from a single "bullish_bias"
figure -- NOT a statistically calibrated forecast. Nothing here has
been validated against real forward outcomes yet; that calibration is
exactly what the (not-yet-built) learning-dataset/outcome-recording
piece from the design brief's section 33 is for.

EXTENSION/EXHAUSTION PENALTIES computed here, not in a separate
extension_engine.py -- same reasoning as stream_features.py's pressure
addition: the design brief describes these as their own "engine"
conceptually (section 14) but the recommended-module list has no
separate file for them, and they're specifically about THIS module's
job (judging entry quality), not a generic stream feature.

RELATIVE STRENGTH -- honestly NOT implemented. It needs a cross-symbol
or market/sector comparison this module has no data source for. Rather
than fabricate a number, its weight is redistributed across the other
components and it's listed in `missing` on every reading -- per the
project's explicit failure-safety rule: don't manufacture values for
data that isn't available.
"""

from dataclasses import dataclass, field

from config_loader import get_config

NEXT_MOVE_STATES = ["CONTINUATION_UP", "PULLBACK_THEN_UP", "SIDEWAYS", "PULLBACK_DOWN", "REVERSAL_DOWN"]
HORIZONS = ["1m", "3m", "5m"]

DEFAULT_CONFIG = {
    "weights": {
        "trend": 15, "structure": 15, "momentum": 12, "acceleration": 10,
        "vwap": 10, "volume": 12, "pressure": 10, "relative_strength": 6,
        "breakout_quality": 5, "risk_quality": 5,
    },
    "momentum_scale": 2.0,          # trend_slope magnitude that maps to a "fully positive" momentum component
    "acceleration_scale": 3.0,
    "max_vwap_atr": 3.0,
    "max_ema9_atr": 3.0,
    "extension_penalty_weight": 20,
    "exhaustion_min_bars_since_pullback": 6,
    "exhaustion_penalty_weight": 15,
    "max_spread_pct": 1.0,
    "spread_penalty_weight": 10,
    "chop_penalty_weight": 15,
    "min_relative_volume": 0.5,
    "poor_liquidity_penalty_weight": 10,
    "horizon_confidence_decay": {"1m": 1.0, "3m": 0.8, "5m": 0.6},
    "expected_move_confidence_factor": 0.5,
}


@dataclass
class PredictionReading:
    symbol: str
    score: float = 0.0
    slope: float = 0.0
    breakdown: dict = field(default_factory=dict)
    penalties: dict = field(default_factory=dict)
    next_move: dict = field(default_factory=dict)
    bullish_probability: dict = field(default_factory=dict)
    expected_move_pct: dict = field(default_factory=dict)
    insufficient_data: bool = False
    missing: list = field(default_factory=list)


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


def _vwap_component(features) -> float:
    """Exact lookup table from the project's original design brief's
    VWAP-engine section: above+rising=strong, above+flat=moderate,
    below+rising=possible reclaim, below+falling=bearish."""
    above = (features.vwap.get("price_vs_vwap_pct") or 0) > 0
    rising = (features.vwap.get("slope") or 0) > 0
    if above and rising:
        return 1.0
    if above and not rising:
        return 0.7
    if not above and rising:
        return 0.4
    return 0.0


def _score_components(features, trend, structure, setups, cfg: dict) -> dict:
    direction_credit = {"UP": 1.0, "NEUTRAL": 0.3, "DOWN": 0.0}[trend.direction]
    trend_component = (trend.score / 100.0) * direction_credit

    structure_credit = {"BULLISH_STRUCTURE": 1.0, "TRANSITION_STRUCTURE": 0.3,
                         "RANGE_STRUCTURE": 0.3, "BEARISH_STRUCTURE": 0.0}[structure.structure]
    structure_component = (structure.score / 100.0) * structure_credit

    momentum_component = _clamp01(0.5 + trend.trend_slope / (2 * cfg["momentum_scale"]))
    accel = features.acceleration.get("momentum_acceleration") or 0.0
    acceleration_component = _clamp01(0.5 + accel / (2 * cfg["acceleration_scale"]))

    vwap_component = _vwap_component(features)

    vol_accel = features.volume.get("volume_acceleration")
    volume_component = _clamp01((vol_accel - 0.5) / 1.5) if vol_accel is not None else 0.5

    pressure_score = features.pressure.get("score")
    pressure_component = _clamp01((pressure_score + 100) / 200.0) if pressure_score is not None else 0.5

    breakout_component = 0.0
    if setups:
        best = max((r for r in setups.values() if r.valid), key=lambda r: r.confidence, default=None)
        if best is not None:
            breakout_component = best.confidence / 100.0

    spread_pct = features.quote.get("spread_pct")
    spread_quality = _clamp01(1 - (spread_pct / cfg["max_spread_pct"])) if spread_pct is not None else 0.5
    vwap_atr = features.volatility.get("vwap_distance_atr")
    extension_quality = _clamp01(1 - abs(vwap_atr) / cfg["max_vwap_atr"]) if vwap_atr is not None else 0.5
    risk_component = (spread_quality + extension_quality) / 2.0

    return {
        "trend": trend_component, "structure": structure_component,
        "momentum": momentum_component, "acceleration": acceleration_component,
        "vwap": vwap_component, "volume": volume_component, "pressure": pressure_component,
        "breakout_quality": breakout_component, "risk_quality": risk_component,
        # relative_strength deliberately absent -- see module docstring
    }


def _compute_penalties(features, structure, regime_label: str, cfg: dict) -> dict:
    """
    [BUGFIX 2026-09-05] extension_fraction below is deliberately left
    UNCAPPED for the penalty calculation (only clamped to [0,1] later,
    separately, for uses that need a blend weight -- see
    compute_prediction()'s extension_fraction_capped). Capping it at 1.0
    here was a real bug: momentum_component/acceleration_component in
    _score_components() are driven by the SAME runaway price move that
    creates extreme extension, and nothing bounds how high THEY can
    read. Verified live in this module's own testing: a moderately
    extended healthy-continuation case (price_vs_vwap_atr=6.2, extension
    already >1.0x threshold, penalty capped at 20) scored 57.18, while a
    deliberately MORE extended case (price_vs_vwap_atr=10.5) scored
    HIGHER at 66.81 -- because momentum/acceleration saturated at their
    max credit from the same extreme move while the capped penalty
    couldn't grow to offset it. This is exactly backwards from section
    39's explicit requirement ("bullish but overextended" must score
    LOWER for entry purposes than a comparable non-extended setup).
    Letting the fraction scale without an artificial ceiling means an
    increasingly extreme move keeps getting an increasingly large
    penalty, so it can never "outrun" the components it inflates.
    """
    penalties = {}

    vwap_atr = features.volatility.get("vwap_distance_atr")
    ema_atr = features.ema.get("price_vs_ema9_atr")
    over_vwap = max(0.0, (abs(vwap_atr) - cfg["max_vwap_atr"]) / cfg["max_vwap_atr"]) if vwap_atr is not None else 0.0
    over_ema = max(0.0, (abs(ema_atr) - cfg["max_ema9_atr"]) / cfg["max_ema9_atr"]) if ema_atr is not None else 0.0
    extension_fraction = max(over_vwap, over_ema)  # deliberately uncapped -- see docstring above
    penalties["extension"] = round(extension_fraction * cfg["extension_penalty_weight"], 2)

    bars_since_low = None
    lows = [s for s in structure.swings if s.kind == "LOW"]
    if lows:
        bars_since_low = len(structure.swings) - 1 - structure.swings.index(lows[-1])
    if bars_since_low is not None and bars_since_low > cfg["exhaustion_min_bars_since_pullback"]:
        excess = bars_since_low - cfg["exhaustion_min_bars_since_pullback"]
        penalties["exhaustion"] = round(min(1.0, excess / cfg["exhaustion_min_bars_since_pullback"])
                                         * cfg["exhaustion_penalty_weight"], 2)
    else:
        penalties["exhaustion"] = 0.0

    spread_pct = features.quote.get("spread_pct")
    if spread_pct is not None and spread_pct > cfg["max_spread_pct"]:
        over = min(1.0, (spread_pct - cfg["max_spread_pct"]) / cfg["max_spread_pct"])
        penalties["spread"] = round(over * cfg["spread_penalty_weight"], 2)
    else:
        penalties["spread"] = 0.0

    penalties["chop"] = cfg["chop_penalty_weight"] if regime_label == "CHOP" else 0.0

    rvol = features.volume.get("relative_volume")
    if rvol is not None and rvol < cfg["min_relative_volume"]:
        shortfall = _clamp01((cfg["min_relative_volume"] - rvol) / cfg["min_relative_volume"])
        penalties["poor_liquidity"] = round(shortfall * cfg["poor_liquidity_penalty_weight"], 2)
    else:
        penalties["poor_liquidity"] = 0.0

    return penalties, extension_fraction


def _bullish_bias(components: dict, trend) -> float:
    """A single -1..1 figure combining the direction-aware components
    already computed for the score -- feeds the next-move heuristic
    below. Deliberately reuses components already computed for the
    score rather than a separate calculation, so the two stay
    consistent with each other."""
    direction_sign = {"UP": 1.0, "NEUTRAL": 0.0, "DOWN": -1.0}[trend.direction]
    bias = (direction_sign * (trend.score / 100.0) * 0.4
            + (components["momentum"] - 0.5) * 2 * 0.3
            + (components["pressure"] - 0.5) * 2 * 0.3)
    return max(-1.0, min(1.0, bias))


def _next_move_distribution(bullish_bias: float, extension_fraction: float, decay: float) -> dict:
    """
    Deterministic heuristic allocation across the 5 states -- see module
    docstring. Higher |bullish_bias| pushes weight toward the matching
    directional states; extension_fraction shifts weight AWAY from clean
    continuation and toward pullback/reversal even when bias is
    positive (the "bullish but overextended" case). decay (<=1) pulls
    the whole distribution toward SIDEWAYS at longer horizons, reflecting
    genuinely lower confidence further out -- not a claim of measured
    accuracy decay, just an honest "we know less the further out we look."
    """
    b = bullish_bias * decay
    continuation_up = max(0.0, b) * (1 - extension_fraction)
    pullback_then_up = max(0.0, b) * (0.3 + 0.7 * extension_fraction) * 0.6
    reversal_down = max(0.0, -b) + max(0.0, b) * extension_fraction * 0.3
    pullback_down = max(0.0, -b) * 0.5
    sideways = max(0.05, 1 - abs(b))

    raw = {"CONTINUATION_UP": continuation_up, "PULLBACK_THEN_UP": pullback_then_up,
           "SIDEWAYS": sideways, "PULLBACK_DOWN": pullback_down, "REVERSAL_DOWN": reversal_down}
    total = sum(raw.values()) or 1.0
    return {k: round(v / total, 3) for k, v in raw.items()}


def compute_prediction(symbol: str, features, trend, structure,
                        regime=None, setups: dict = None,
                        prior: "PredictionReading" = None, cfg: dict = None) -> PredictionReading:
    cfg = {**DEFAULT_CONFIG, **(cfg or get_config().get("prediction_engine", {}))}
    weights = {**DEFAULT_CONFIG["weights"], **(cfg.get("weights") or {})}

    missing = ["relative_strength (not implemented -- needs cross-symbol/market comparison "
               "data this module doesn't have)"]
    if (getattr(features, "insufficient_data", False) or getattr(trend, "insufficient_data", False)
            or getattr(structure, "insufficient_data", False)):
        return PredictionReading(symbol=symbol, insufficient_data=True,
                                  missing=missing + ["prediction -- upstream reading was insufficient"])

    components = _score_components(features, trend, structure, setups, cfg)
    weight_used = {k: v for k, v in weights.items() if k != "relative_strength"}
    weight_total = sum(weight_used.values())
    breakdown = {k: round(components[k] * weight_used[k], 2) for k in weight_used}
    base_score = sum(breakdown.values()) / weight_total * 100.0

    regime_label = getattr(regime, "regime", None) if regime is not None else None
    penalties, extension_fraction = _compute_penalties(features, structure, regime_label, cfg)
    score = round(max(0.0, min(100.0, base_score - sum(penalties.values()))), 2)

    slope = 0.0
    if prior is not None and not prior.insufficient_data:
        slope = round(score - prior.score, 2)

    bias = _bullish_bias(components, trend)
    extension_fraction_capped = _clamp01(extension_fraction)  # blend weight use only -- penalty above uses the uncapped value
    decay_cfg = cfg["horizon_confidence_decay"]
    next_move, bullish_probability, expected_move_pct = {}, {}, {}
    base_rate_pct_per_min = abs(features.slope.get("slope_1m") or 0.0)
    horizon_minutes = {"1m": 1, "3m": 3, "5m": 5}
    for h in HORIZONS:
        dist = _next_move_distribution(bias, extension_fraction_capped, decay_cfg.get(h, 1.0))
        next_move[h] = dist
        bullish_probability[h] = round(dist["CONTINUATION_UP"] + dist["PULLBACK_THEN_UP"], 3)
        expected_move_pct[h] = round(
            bias * base_rate_pct_per_min * horizon_minutes[h] * cfg["expected_move_confidence_factor"], 3)

    return PredictionReading(
        symbol=symbol, score=score, slope=slope, breakdown=breakdown, penalties=penalties,
        next_move=next_move, bullish_probability=bullish_probability,
        expected_move_pct=expected_move_pct, insufficient_data=False, missing=missing,
    )
