"""Causal, non-promoting study of MLP feature states before extreme BTC bars.

The event at bar t is the close-to-close return from t-1 to t.  Every
candidate predictor is taken from t-1, so it was available before the event
bar began.  Tail thresholds and feature cuts are calibrated only in the
discovery window, then held fixed in validation and OOS windows.

This is an event-profile diagnostic, not a training command and not a source
of promoted MLP parameters.
"""
from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict, dataclass
from pathlib import Path

import numpy as np
import pandas as pd

REPO = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO))
from strategies.strategy_mlp_scores import FEATURE_COLS


DISCOVERY_END = pd.Timestamp("2023-12-31 23:59:59")
VALIDATION_START = pd.Timestamp("2024-01-01")
OOS_START = pd.Timestamp("2026-03-01")
DEFAULT_FEATURES = ("mlp_score", *dict.fromkeys(FEATURE_COLS))


@dataclass(frozen=True)
class TailDefinition:
    up_return: float
    down_return: float
    tail_fraction: float


def parse_chart_time(values: pd.Series) -> pd.Series:
    """Parse TradingView epoch-second exports and ISO timestamps deterministically."""
    numeric = pd.to_numeric(values, errors="coerce")
    if numeric.notna().all():
        return pd.to_datetime(numeric, unit="s", utc=True).dt.tz_localize(None)
    return pd.to_datetime(values, utc=True, errors="raise").dt.tz_localize(None)


def causal_event_frame(data: pd.DataFrame, features: list[str], *, bar_hours: float = 6.0) -> pd.DataFrame:
    """Return feature snapshots at t and their contiguous next-bar return."""
    if bar_hours <= 0:
        raise ValueError("bar_hours must be positive")
    missing = [column for column in ["time", "close", *features] if column not in data.columns]
    if missing:
        raise ValueError(f"missing required columns: {missing}")
    frame = data[["time", "close", *features]].copy()
    frame["time"] = parse_chart_time(frame["time"])
    frame = frame.sort_values("time").reset_index(drop=True)
    if frame["time"].duplicated().any() or not frame["time"].is_monotonic_increasing:
        raise ValueError("time must be unique and strictly increasing")
    next_time = frame["time"].shift(-1)
    max_gap = pd.Timedelta(hours=bar_hours * 1.5)
    contiguous = (next_time - frame["time"]) <= max_gap
    frame["forward_return"] = np.log(frame["close"].shift(-1) / frame["close"]).where(contiguous)
    frame["outcome_time"] = next_time.where(contiguous)
    return frame


def calibrate_tails(discovery: pd.DataFrame, tail_fraction: float) -> TailDefinition:
    """Set fixed extreme-move labels using discovery outcomes only."""
    if not 0 < tail_fraction < 0.5:
        raise ValueError("tail_fraction must be between 0 and 0.5")
    returns = discovery["forward_return"].dropna()
    if len(returns) < 20:
        raise ValueError("discovery window needs at least 20 complete outcomes")
    return TailDefinition(
        up_return=float(returns.quantile(1.0 - tail_fraction)),
        down_return=float(returns.quantile(tail_fraction)),
        tail_fraction=tail_fraction,
    )


def _window(frame: pd.DataFrame, start: pd.Timestamp | None, end: pd.Timestamp | None) -> pd.DataFrame:
    mask = pd.Series(True, index=frame.index)
    if start is not None:
        mask &= frame["time"] >= start
    if end is not None:
        mask &= frame["time"] <= end
        mask &= frame["outcome_time"] <= end
    return frame.loc[mask & frame["forward_return"].notna()].copy()


def _profile_window(frame: pd.DataFrame, feature: str, cut: float, side: str, event: str, tails: TailDefinition) -> dict[str, float | int | None]:
    event_mask = frame["forward_return"] >= tails.up_return if event == "up" else frame["forward_return"] <= tails.down_return
    selected = frame[feature] >= cut if side == "high" else frame[feature] <= cut
    usable = frame[feature].notna()
    selected &= usable
    base_rate = float(event_mask[usable].mean()) if usable.any() else None
    event_rate = float(event_mask[selected].mean()) if selected.any() else None
    return {
        "observations": int(usable.sum()),
        "selected": int(selected.sum()),
        "event_rate": event_rate,
        "base_rate": base_rate,
        "lift": event_rate / base_rate if event_rate is not None and base_rate and base_rate > 0 else None,
    }


def profile_feature(discovery: pd.DataFrame, validation: pd.DataFrame, oos: pd.DataFrame, *, feature: str, event: str, tails: TailDefinition, feature_tail: float) -> dict[str, object]:
    """Choose one discovery-tail direction, then hold that rule fixed later."""
    if event not in {"up", "down"}:
        raise ValueError("event must be 'up' or 'down'")
    values = discovery[feature].dropna()
    if len(values) < 20:
        raise ValueError(f"{feature} has too few discovery observations")
    low_cut, high_cut = float(values.quantile(feature_tail)), float(values.quantile(1.0 - feature_tail))
    low = _profile_window(discovery, feature, low_cut, "low", event, tails)
    high = _profile_window(discovery, feature, high_cut, "high", event, tails)
    low_lift, high_lift = low["lift"] or float("-inf"), high["lift"] or float("-inf")
    side, cut, discovery_profile = ("high", high_cut, high) if high_lift >= low_lift else ("low", low_cut, low)
    return {
        "feature": feature,
        "event": event,
        "side": side,
        "cut": cut,
        "discovery": discovery_profile,
        "validation": _profile_window(validation, feature, cut, side, event, tails),
        "oos": _profile_window(oos, feature, cut, side, event, tails),
    }


def _rank_key(profile: dict[str, object]) -> tuple[float, str]:
    discovery = profile["discovery"]
    assert isinstance(discovery, dict)
    return (float(discovery["lift"] or float("-inf")), str(profile["feature"]))


def relative_tail_profiles(discovery: pd.DataFrame, validation: pd.DataFrame, oos: pd.DataFrame, *, features: list[str], feature_tail: float, tail_fraction: float) -> dict[str, list[dict[str, object]]]:
    """Describe within-window tail behavior without reselecting feature rules.

    The event thresholds are recalculated within each reporting window, so this
    is descriptive of *relative* extremes rather than an executable fixed-size
    return rule.  Feature direction and cut remain selected in discovery.
    """
    discovery_tails = calibrate_tails(discovery, tail_fraction)
    windows = {"discovery": discovery, "validation": validation, "oos": oos}
    result: dict[str, list[dict[str, object]]] = {}
    for event in ("up", "down"):
        profiles = []
        for feature in features:
            selected = profile_feature(
                discovery, validation, oos, feature=feature, event=event,
                tails=discovery_tails, feature_tail=feature_tail,
            )
            relative = {
                name: _profile_window(window, feature, float(selected["cut"]), str(selected["side"]), event,
                                      calibrate_tails(window, tail_fraction))
                for name, window in windows.items()
            }
            profiles.append({
                "feature": feature,
                "event": event,
                "side": selected["side"],
                "cut": selected["cut"],
                "windows": relative,
            })
        result[event] = sorted(
            profiles,
            key=lambda profile: float(profile["windows"]["discovery"]["lift"] or float("-inf")),
            reverse=True,
        )
    return result


def run(data_path: Path, *, tail_fraction: float = 0.05, feature_tail: float = 0.10, bar_hours: float = 6.0, features: list[str] | None = None) -> dict[str, object]:
    data = pd.read_csv(data_path)
    data.columns = data.columns.str.lower().str.strip()
    selected_features = list(dict.fromkeys(features or [feature for feature in DEFAULT_FEATURES if feature in data.columns]))
    if not selected_features:
        raise ValueError("no requested feature columns are present")
    frame = causal_event_frame(data, selected_features, bar_hours=bar_hours)
    discovery = _window(frame, None, DISCOVERY_END)
    validation = _window(frame, VALIDATION_START, OOS_START - pd.Timedelta(nanoseconds=1))
    oos = _window(frame, OOS_START, None)
    tails = calibrate_tails(discovery, tail_fraction)
    studies: dict[str, list[dict[str, object]]] = {}
    for event in ("up", "down"):
        profiles = [profile_feature(discovery, validation, oos, feature=feature, event=event, tails=tails, feature_tail=feature_tail)
                    for feature in selected_features]
        studies[event] = sorted(profiles, key=_rank_key, reverse=True)
    return {
        "data_path": str(data_path),
        "definition": {
            "event": "next 6h close-to-close return; feature snapshot is the preceding close",
            "max_label_gap_hours": bar_hours * 1.5,
            "tail_definition": asdict(tails),
            "feature_tail": feature_tail,
            "discovery_end": str(DISCOVERY_END),
            "validation": [str(VALIDATION_START), str(OOS_START - pd.Timedelta(nanoseconds=1))],
            "oos_start": str(OOS_START),
        },
        "windows": {"discovery": len(discovery), "validation": len(validation), "oos": len(oos)},
        "studies": studies,
        "relative_tail_studies": relative_tail_profiles(
            discovery, validation, oos, features=selected_features,
            feature_tail=feature_tail, tail_fraction=tail_fraction,
        ),
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--data", type=Path, required=True, help="TradingView chart-data CSV to read without modifying")
    parser.add_argument("--tail-fraction", type=float, default=0.05)
    parser.add_argument("--feature-tail", type=float, default=0.10)
    parser.add_argument("--bar-hours", type=float, default=6.0, help="Expected chart-bar duration; labels crossing >1.5 bars are dropped")
    parser.add_argument("--features", nargs="*", help="Feature columns to profile; defaults to mlp_score when available")
    parser.add_argument("--output", type=Path, help="Optional JSON output path")
    args = parser.parse_args()
    result = run(args.data, tail_fraction=args.tail_fraction, feature_tail=args.feature_tail, bar_hours=args.bar_hours, features=args.features)
    text = json.dumps(result, indent=2)
    print(text)
    if args.output:
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(text + "\n")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
