"""Research-only matched-control and simple classifier study for extreme BTC bars."""
from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import pandas as pd

REPO = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO))

from strategies.mlp_tail_event_research import (
    TailRule,
    moving_block_bootstrap_matched_lift,
    tail_classifier_metrics,
    volatility_matched_downside_profile,
)
from tools.run_mlp_extreme_bar_event_study import DISCOVERY_END, OOS_START, VALIDATION_START, _window, causal_event_frame


RULE_FEATURES = ("bb_pct_b_norm", "sopr_norm", "osc_norm")
MODEL_FEATURE_SETS = {
    "mlp_score": ["mlp_score"],
    "rvol": ["rvol_norm"],
    "downside_trio": list(RULE_FEATURES),
    "combined": ["mlp_score", "rvol_norm", *RULE_FEATURES],
}


def run(
    data_path: Path,
    *,
    tail_fractions: tuple[float, ...] = (0.05, 0.10, 0.15),
    bootstrap_repetitions: int = 200,
    bar_hours: float = 6.0,
) -> dict[str, object]:
    data = pd.read_csv(data_path)
    data.columns = data.columns.str.lower().str.strip()
    features = list(dict.fromkeys(["mlp_score", "rvol_norm", *RULE_FEATURES]))
    frame = causal_event_frame(data, 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)
    cuts = {feature: float(discovery[feature].quantile(0.10)) for feature in RULE_FEATURES}
    rules = [TailRule(feature, "low", cuts[feature]) for feature in RULE_FEATURES]
    profiles = {
        str(tail): {
            rule.feature: {
                window: volatility_matched_downside_profile(frame, rule, tail_fraction=tail)
                for window, frame in {"discovery": discovery, "validation": validation, "oos": oos}.items()
            }
            for rule in rules
        }
        for tail in tail_fractions
    }
    # A 10% tail supplies enough events for a modest uncertainty estimate in
    # the currently short OOS window. Rules remain frozen from discovery.
    bootstrap = {
        rule.feature: {
            window: moving_block_bootstrap_matched_lift(
                values, rule, tail_fraction=0.10, repetitions=bootstrap_repetitions,
            )
            for window, values in {"discovery": discovery, "validation": validation, "oos": oos}.items()
        }
        for rule in rules
    }
    chronological: dict[str, dict[str, object]] = {}
    for year in sorted(frame["time"].dt.year.dropna().unique()):
        yearly = _window(frame, pd.Timestamp(f"{year}-01-01"), pd.Timestamp(f"{year}-12-31 23:59:59"))
        if len(yearly) < 100:
            continue
        chronological[str(year)] = {
            rule.feature: volatility_matched_downside_profile(yearly, rule, tail_fraction=0.10)
            for rule in rules
        }
    # 10% tails give sufficient positives for a deliberately small baseline comparison.
    classifiers = {
        name: {
            "validation": tail_classifier_metrics(discovery, validation, features=features, tail_fraction=0.10),
            "oos": tail_classifier_metrics(discovery, oos, features=features, tail_fraction=0.10),
        }
        for name, features in MODEL_FEATURE_SETS.items()
    }
    return {
        "data_path": str(data_path),
        "definition": {
            "predictors": f"values at t; outcome is contiguous next {bar_hours:g}h close-to-close return",
            "matched_control": "calendar-year × realized-volatility decile standardization",
            "tail_fractions": list(tail_fractions),
            "classifier": "discovery-trained L2 logistic baseline; relative 10% tails in later reporting windows",
            "bootstrap": f"{bootstrap_repetitions} deterministic moving-block resamples (28 six-hour bars) for matched 10% downside lift",
            "chronological_slices": "calendar-year 10% relative tails using discovery-fixed rules; descriptive stability only",
        },
        "windows": {"discovery": len(discovery), "validation": len(validation), "oos": len(oos)},
        "downside_rules": [{"feature": rule.feature, "side": rule.side, "cut": rule.cut} for rule in rules],
        "volatility_matched_downside": profiles,
        "bootstrap_matched_downside_10pct": bootstrap,
        "calendar_year_matched_downside_10pct": chronological,
        "classifier": classifiers,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--data", type=Path, required=True)
    parser.add_argument("--output", type=Path)
    parser.add_argument("--bootstrap-repetitions", type=int, default=200)
    parser.add_argument("--bar-hours", type=float, default=6.0, help="Expected chart-bar duration; labels crossing >1.5 bars are dropped")
    args = parser.parse_args()
    result = run(args.data, bootstrap_repetitions=args.bootstrap_repetitions, bar_hours=args.bar_hours)
    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())
