"""
outcome_engine.py

[SIMULATION-ONLY 2026-09-05] The learning-dataset / outcome-recording
piece from the project's original design brief's section 33: records
every prediction_engine.py reading made for a symbol, then -- once
enough forward bars exist -- resolves what ACTUALLY happened at each
horizon (1m/3m/5m) and reports how well the deterministic heuristic's
next-move probabilities and expected-move estimates matched reality.

Not consumed anywhere in the live bot yet -- see config.json's
outcome_engine._note.

STILL NOT ML. This does not train, fit, or adjust any weight in
prediction_engine.py -- per the design brief's own two-phase instruction
("first build deterministic prediction engine ... then the historical
dataset can later be used to train/calibrate"), this module is the
DATASET half of that sentence, not the TRAIN half. It writes one record
per (symbol, as_of) prediction, resolves it against real forward bars,
and computes calibration statistics (hit rate, Brier score, expected-
move error) a future calibration step would consume -- it never changes
prediction_engine.py's behavior itself.

REALIZED-OUTCOME CLASSIFICATION mirrors prediction_engine.NEXT_MOVE_STATES
exactly, using the same kind of threshold logic a person would apply
looking at a completed forward price path (classify_realized_outcome()
below), deterministic and documented, not a different scheme that just
happens to share label names.

No state, no file I/O, no network calls -- same pure-function contract
as every other engine built this session. The caller (a backtest/
orchestrator script today; eventually monitor.py) owns persistence and
owns holding pending records until enough forward bars exist to resolve
them.
"""

from dataclasses import dataclass, field, asdict
from datetime import datetime, timezone

from config_loader import get_config

CONTINUATION_UP = "CONTINUATION_UP"
PULLBACK_THEN_UP = "PULLBACK_THEN_UP"
SIDEWAYS = "SIDEWAYS"
PULLBACK_DOWN = "PULLBACK_DOWN"
REVERSAL_DOWN = "REVERSAL_DOWN"
BULLISH_STATES = (CONTINUATION_UP, PULLBACK_THEN_UP)

DEFAULT_CONFIG = {
    "sideways_band_pct": 0.15,         # |net move| within this band at a horizon -> SIDEWAYS, regardless of path
    "pullback_min_retrace_pct": 0.10,  # min adverse excursion before an up move counts as "pullback-then-up"
    "reversal_min_down_pct": 0.15,     # net down move beyond this -> REVERSAL_DOWN rather than a mild PULLBACK_DOWN
    "calibration_bins": 5,
}


@dataclass
class PredictionRecord:
    symbol: str
    as_of: str                      # ISO timestamp -- the prediction's own reference instant
    reference_price: float
    predicted_score: float
    predicted_next_move: dict       # {horizon: {state: prob}}
    predicted_bullish_probability: dict   # {horizon: prob}
    predicted_expected_move_pct: dict     # {horizon: pct}
    context: dict = field(default_factory=dict)    # trend/structure/regime labels, for slicing a report later
    realized: dict = field(default_factory=dict)   # filled in by resolve_outcomes(), keyed by horizon
    resolved: bool = False

    def to_dict(self) -> dict:
        return asdict(self)


def record_prediction(symbol: str, as_of, reference_price: float, prediction,
                       context: dict = None) -> PredictionRecord:
    """
    Builds a pending PredictionRecord from a prediction_engine.PredictionReading.
    Returns None if the prediction itself was insufficient_data -- there is
    nothing to calibrate against a reading the engine already declined to make.
    Persistence (appending record.to_dict() to a JSONL dataset file, or
    equivalent) is the caller's job, matching every other module's
    "state ownership lives with the caller" contract.
    """
    if getattr(prediction, "insufficient_data", False):
        return None
    return PredictionRecord(
        symbol=symbol,
        as_of=as_of.isoformat() if isinstance(as_of, datetime) else str(as_of),
        reference_price=reference_price,
        predicted_score=prediction.score,
        predicted_next_move=prediction.next_move,
        predicted_bullish_probability=prediction.bullish_probability,
        predicted_expected_move_pct=prediction.expected_move_pct,
        context=context or {},
    )


def classify_realized_outcome(reference_price: float, path_bars: list, cfg: dict = None) -> dict:
    """
    path_bars: bars strictly AFTER the reference instant, up to and
    including the horizon bar, oldest first -- the forward path being
    judged. Returns a dict with the realized state plus the raw
    measurements a calibration report needs (net move, and the best/worst
    excursion along the way, so a near-miss is visible, not just the
    final bucket).

    Fixed priority order, exactly one state returned -- same "worked
    example" discipline as entry_score.py's ABORT-not-average behavior:

      1. SIDEWAYS         -- net move at the horizon stayed within
                              sideways_band_pct either direction,
                              regardless of what happened in between.
      2. REVERSAL_DOWN     -- net move down beyond reversal_min_down_pct
                              (a real, held decline).
      3. PULLBACK_DOWN      -- net move down, but a smaller give-back
                              than that.
      4. PULLBACK_THEN_UP   -- net move at the horizon is UP, but the
                              path dipped at least pullback_min_retrace_pct
                              BELOW the reference price before recovering.
      5. CONTINUATION_UP    -- net move is up and never gave back more
                              than pullback_min_retrace_pct along the way.

    insufficient path_bars (empty) reports SIDEWAYS with is_sufficient=False
    rather than fabricating a directional outcome from no data.
    """
    cfg = {**DEFAULT_CONFIG, **(cfg or get_config().get("outcome_engine", {}))}

    if not path_bars or not reference_price:
        return {"state": SIDEWAYS, "net_move_pct": 0.0, "max_favorable_pct": 0.0,
                "max_adverse_pct": 0.0, "is_sufficient": False}

    final_price = path_bars[-1]["c"]
    net_move_pct = (final_price - reference_price) / reference_price * 100.0
    highs = [(b["h"] - reference_price) / reference_price * 100.0 for b in path_bars]
    lows = [(b["l"] - reference_price) / reference_price * 100.0 for b in path_bars]
    max_favorable_pct = max(highs)
    max_adverse_pct = min(lows)

    if abs(net_move_pct) <= cfg["sideways_band_pct"]:
        state = SIDEWAYS
    elif net_move_pct < 0:
        state = REVERSAL_DOWN if net_move_pct <= -cfg["reversal_min_down_pct"] else PULLBACK_DOWN
    elif max_adverse_pct <= -cfg["pullback_min_retrace_pct"]:
        state = PULLBACK_THEN_UP
    else:
        state = CONTINUATION_UP

    return {"state": state, "net_move_pct": round(net_move_pct, 4),
            "max_favorable_pct": round(max_favorable_pct, 4),
            "max_adverse_pct": round(max_adverse_pct, 4), "is_sufficient": True}


def resolve_outcomes(record: PredictionRecord, forward_bars_by_horizon: dict, cfg: dict = None) -> PredictionRecord:
    """
    forward_bars_by_horizon: {"1m": [...], "3m": [...], "5m": [...]} --
    for each horizon, the real bars strictly after record.as_of up to
    that horizon's close (caller slices these from its own bar buffer;
    this function does no time arithmetic itself, so it works identically
    whether the caller is a live poll loop or an offline replay). A
    horizon missing from the dict (not enough forward data existed yet,
    e.g. the session ended) is simply left unresolved for that horizon --
    partial resolution is normal, not an error.
    """
    for horizon, bars in forward_bars_by_horizon.items():
        outcome = classify_realized_outcome(record.reference_price, bars, cfg)
        predicted_dist = record.predicted_next_move.get(horizon, {})
        predicted_bullish = record.predicted_bullish_probability.get(horizon)
        predicted_move = record.predicted_expected_move_pct.get(horizon)

        argmax_state = max(predicted_dist, key=predicted_dist.get) if predicted_dist else None
        realized_bullish = 1.0 if outcome["state"] in BULLISH_STATES else 0.0
        brier = (predicted_bullish - realized_bullish) ** 2 if predicted_bullish is not None else None
        move_error_pct = (abs(predicted_move - outcome["net_move_pct"])
                           if (predicted_move is not None and outcome["is_sufficient"]) else None)

        record.realized[horizon] = {
            **outcome,
            "predicted_argmax_state": argmax_state,
            "argmax_correct": (argmax_state == outcome["state"]) if argmax_state else None,
            "predicted_bullish_probability": predicted_bullish,
            "brier_component": brier,
            "predicted_expected_move_pct": predicted_move,
            "move_error_pct": move_error_pct,
        }
    record.resolved = len(record.realized) > 0
    return record


def summarize_calibration(records: list, horizons: list = ("1m", "3m", "5m"), cfg: dict = None) -> dict:
    """
    Aggregates resolved PredictionRecords into per-horizon calibration
    stats: argmax hit rate (how often the single most-likely predicted
    state was the one that actually happened -- a tough bar, since there
    are 5 states), mean Brier component for the bullish/not-bullish call
    (0 = perfect, 0.25 = no better than always guessing 50%, 1 = always
    confidently wrong), mean absolute expected-move error, and a
    reliability table (bucket predicted_bullish_probability into N bins,
    report each bin's actual realized-bullish rate -- a well-calibrated
    engine's bins should track the diagonal).
    """
    cfg = {**DEFAULT_CONFIG, **(cfg or get_config().get("outcome_engine", {}))}
    n_bins = cfg["calibration_bins"]
    out = {}

    for h in horizons:
        readings = [r.realized[h] for r in records if r.resolved and h in r.realized
                    and r.realized[h]["is_sufficient"]]
        if not readings:
            out[h] = {"n": 0}
            continue

        n = len(readings)
        argmax_hits = sum(1 for x in readings if x["argmax_correct"])
        briers = [x["brier_component"] for x in readings if x["brier_component"] is not None]
        move_errors = [x["move_error_pct"] for x in readings if x["move_error_pct"] is not None]

        bins = [[] for _ in range(n_bins)]
        for x in readings:
            p = x["predicted_bullish_probability"]
            if p is None:
                continue
            idx = min(n_bins - 1, int(p * n_bins))
            bins[idx].append(x)
        reliability = []
        for i, bucket in enumerate(bins):
            if not bucket:
                continue
            lo, hi = i / n_bins, (i + 1) / n_bins
            predicted_mean = sum(x["predicted_bullish_probability"] for x in bucket) / len(bucket)
            realized_rate = sum(1 for x in bucket if x["state"] in BULLISH_STATES) / len(bucket)
            reliability.append({"bin": f"{lo:.1f}-{hi:.1f}", "n": len(bucket),
                                 "predicted_mean": round(predicted_mean, 3),
                                 "realized_bullish_rate": round(realized_rate, 3)})

        state_counts = {}
        for x in readings:
            state_counts[x["state"]] = state_counts.get(x["state"], 0) + 1

        out[h] = {
            "n": n,
            "argmax_hit_rate": round(argmax_hits / n, 3),
            "mean_brier": round(sum(briers) / len(briers), 4) if briers else None,
            "mean_abs_move_error_pct": round(sum(move_errors) / len(move_errors), 4) if move_errors else None,
            "realized_state_distribution": state_counts,
            "reliability": reliability,
        }
    return out
