"""Leakage-aware research utilities for relative extreme-bar classification."""
from __future__ import annotations

from dataclasses import dataclass

import numpy as np
import pandas as pd
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import average_precision_score
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler


@dataclass(frozen=True)
class TailRule:
    feature: str
    side: str
    cut: float


def relative_tail_labels(frame: pd.DataFrame, tail_fraction: float) -> pd.Series:
    """Label each complete outcome as down (-1), ordinary (0), or up (+1)."""
    if not 0 < tail_fraction < 0.5:
        raise ValueError("tail_fraction must be between 0 and 0.5")
    returns = frame["forward_return"]
    lower, upper = returns.quantile(tail_fraction), returns.quantile(1.0 - tail_fraction)
    labels = pd.Series(0, index=frame.index, dtype=np.int8)
    labels.loc[returns <= lower] = -1
    labels.loc[returns >= upper] = 1
    return labels


def _selection(frame: pd.DataFrame, rule: TailRule) -> pd.Series:
    if rule.side == "low":
        return frame[rule.feature] <= rule.cut
    if rule.side == "high":
        return frame[rule.feature] >= rule.cut
    raise ValueError("rule side must be 'low' or 'high'")


def volatility_matched_downside_profile(frame: pd.DataFrame, rule: TailRule, *, tail_fraction: float, volatility_feature: str = "rvol_norm", bins: int = 10) -> dict[str, float | int | None]:
    """Compare selected and control bars within calendar-year × volatility bins.

    This standardizes the control rate to the selected bars' stratum mix. It
    is deliberately a simple matched descriptive test, not a causal estimator
    of a deployable trading rule.
    """
    if bins < 2:
        raise ValueError("bins must be at least two")
    required = ["time", "forward_return", rule.feature, volatility_feature]
    missing = [column for column in required if column not in frame]
    if missing:
        raise ValueError(f"missing required columns: {missing}")
    data = frame.dropna(subset=required).copy()
    labels = relative_tail_labels(data, tail_fraction)
    data["selected"] = _selection(data, rule)
    data["down"] = labels.eq(-1)
    data["year"] = data["time"].dt.year
    data["vol_bin"] = data.groupby("year")[volatility_feature].rank(method="first", pct=True).mul(bins).clip(upper=bins - 1e-12).astype(int)

    selected_total = control_total = weighted_selected_rate = weighted_control_rate = 0.0
    strata = 0
    for _, group in data.groupby(["year", "vol_bin"], observed=True):
        selected, control = group[group["selected"]], group[~group["selected"]]
        if selected.empty or control.empty:
            continue
        weight = len(selected)
        selected_total += weight
        control_total += len(control)
        weighted_selected_rate += weight * selected["down"].mean()
        weighted_control_rate += weight * control["down"].mean()
        strata += 1
    if selected_total == 0:
        return {"selected": 0, "controls": 0, "strata": 0, "selected_rate": None, "matched_control_rate": None, "delta": None, "lift": None}
    selected_rate = weighted_selected_rate / selected_total
    control_rate = weighted_control_rate / selected_total
    return {
        "selected": int(selected_total),
        "controls": int(control_total),
        "strata": strata,
        "selected_rate": float(selected_rate),
        "matched_control_rate": float(control_rate),
        "delta": float(selected_rate - control_rate),
        "lift": float(selected_rate / control_rate) if control_rate > 0 else None,
    }


def moving_block_bootstrap_matched_lift(
    frame: pd.DataFrame,
    rule: TailRule,
    *,
    tail_fraction: float,
    block_bars: int = 28,
    repetitions: int = 200,
    seed: int = 7,
) -> dict[str, float | int | None]:
    """Estimate matched-lift uncertainty with deterministic contiguous resampling.

    Resampling individual bars would assume independent six-hour returns. This
    uses contiguous blocks (seven days by default), retaining short-run serial
    structure. It is an uncertainty description for this observational study,
    not an optimization input or a trading confidence score.
    """
    if block_bars < 2:
        raise ValueError("block_bars must be at least two")
    if repetitions < 20:
        raise ValueError("repetitions must be at least 20")
    data = frame.dropna(subset=["time", "forward_return", rule.feature, "rvol_norm"]).reset_index(drop=True)
    if len(data) < block_bars:
        raise ValueError("frame needs at least one complete block")
    rng = np.random.default_rng(seed)
    lifts: list[float] = []
    starts = np.arange(len(data) - block_bars + 1)
    for _ in range(repetitions):
        indexes: list[int] = []
        while len(indexes) < len(data):
            start = int(rng.choice(starts))
            indexes.extend(range(start, start + block_bars))
        sample = data.iloc[indexes[:len(data)]].copy()
        lift = volatility_matched_downside_profile(sample, rule, tail_fraction=tail_fraction)["lift"]
        if lift is not None and np.isfinite(lift):
            lifts.append(float(lift))
    if not lifts:
        return {"repetitions": repetitions, "valid_repetitions": 0, "p05": None, "p50": None, "p95": None, "share_above_one": None}
    values = np.asarray(lifts)
    return {
        "repetitions": repetitions,
        "valid_repetitions": len(values),
        "p05": float(np.quantile(values, 0.05)),
        "p50": float(np.quantile(values, 0.50)),
        "p95": float(np.quantile(values, 0.95)),
        "share_above_one": float((values > 1.0).mean()),
    }


def tail_classifier_metrics(train: pd.DataFrame, evaluation: pd.DataFrame, *, features: list[str], tail_fraction: float) -> dict[str, float | int]:
    """Fit on an earlier window and score relative tails in a later window."""
    missing = [column for column in [*features, "forward_return"] if column not in train or column not in evaluation]
    if missing:
        raise ValueError(f"missing required columns: {missing}")
    y_train = relative_tail_labels(train, tail_fraction).to_numpy() + 1
    y_eval = relative_tail_labels(evaluation, tail_fraction).to_numpy() + 1
    model = make_pipeline(
        SimpleImputer(strategy="median"),
        StandardScaler(),
        # ``lbfgs`` encounters a platform-dependent divide-by-zero warning in
        # its multiclass Hessian calculation for these correlated bounded
        # chart features.  ``saga`` is deterministic with random_state and
        # avoids that numerical path while retaining the same L2 objective.
        LogisticRegression(
            C=0.1,
            class_weight="balanced",
            max_iter=2_000,
            random_state=7,
            solver="saga",
        ),
    )
    # Chart exports may contain +/-inf normalizations. Treat them as missing;
    # the imputer is fit only on the earlier training window.
    train_features = train[features].replace([np.inf, -np.inf], np.nan)
    evaluation_features = evaluation[features].replace([np.inf, -np.inf], np.nan)
    model.fit(train_features, y_train)
    # Use an explicit stable multinomial softmax.  On the current macOS
    # NumPy/BLAS build, sklearn's otherwise equivalent dense matrix multiply
    # emits spurious floating-point warnings for finite, small matrices.
    # ``einsum`` avoids that platform path and makes the calculation auditable.
    classifier = model[-1]
    transformed = model[:-1].transform(evaluation_features)
    scores = np.einsum("ij,kj->ik", transformed, classifier.coef_) + classifier.intercept_
    scores -= scores.max(axis=1, keepdims=True)
    exp_scores = np.exp(scores)
    probabilities = exp_scores / exp_scores.sum(axis=1, keepdims=True)
    classes = classifier.classes_
    down_column, up_column = int(np.flatnonzero(classes == 0)[0]), int(np.flatnonzero(classes == 2)[0])
    return {
        "observations": len(evaluation),
        "down_pr_auc": float(average_precision_score(y_eval == 0, probabilities[:, down_column])),
        "up_pr_auc": float(average_precision_score(y_eval == 2, probabilities[:, up_column])),
        "tail_base_rate": float(tail_fraction),
    }
