import numpy as np
import pandas as pd
import pytest

from tools.run_mlp_extreme_bar_event_study import _window, causal_event_frame, calibrate_tails, profile_feature, relative_tail_profiles


def _frame():
    return pd.DataFrame(
        {
            "time": [1_700_000_000 + 21_600 * index for index in range(40)],
            "close": np.exp(np.arange(40, dtype=float) / 100),
            "signal": np.arange(40, dtype=float),
        }
    )


def test_causal_frame_parses_epoch_seconds_and_aligns_feature_before_outcome():
    data = _frame()
    frame = causal_event_frame(data, ["signal"])

    assert frame.loc[0, "time"] == pd.Timestamp("2023-11-14 22:13:20")
    assert frame.loc[0, "forward_return"] == pytest.approx(0.01)
    assert pd.isna(frame.loc[len(frame) - 1, "forward_return"])


def test_causal_frame_drops_a_label_that_crosses_a_chart_gap():
    data = _frame().iloc[:3].copy()
    data.loc[2, "time"] += 21_600 * 10

    frame = causal_event_frame(data, ["signal"])

    assert pd.isna(frame.loc[1, "forward_return"])


def test_window_excludes_an_outcome_that_crosses_its_end_boundary():
    frame = causal_event_frame(_frame().iloc[:3], ["signal"])
    end = frame.loc[1, "time"]

    window = _window(frame, None, end)

    assert window.index.tolist() == [0]


def test_tail_thresholds_do_not_change_when_later_window_is_mutated():
    frame = causal_event_frame(_frame(), ["signal"])
    discovery = frame.iloc[:25].copy()
    before = calibrate_tails(discovery, 0.10)
    frame.loc[25:, "close"] *= 1_000
    after = calibrate_tails(discovery, 0.10)

    assert after == before


def test_profile_selects_discovery_direction_and_keeps_its_cut_in_later_windows():
    time = pd.date_range("2023-01-01", periods=60, freq="6h")
    values = np.tile(np.arange(10, dtype=float), 6)
    forward = np.where(values >= 8, 0.10, -0.01)
    frame = pd.DataFrame({"time": time, "close": 100.0, "signal": values, "forward_return": forward})
    tails = calibrate_tails(frame.iloc[:30], 0.10)

    profile = profile_feature(frame.iloc[:30], frame.iloc[30:45], frame.iloc[45:], feature="signal", event="up", tails=tails, feature_tail=0.20)

    assert profile["side"] == "high"
    assert profile["cut"] == pytest.approx(7.2)
    assert profile["validation"]["lift"] > 1.0


def test_relative_tail_profiles_do_not_reselect_the_feature_cut_out_of_sample():
    time = pd.date_range("2023-01-01", periods=90, freq="6h")
    values = np.tile(np.arange(10, dtype=float), 9)
    forward = np.where(values >= 8, 0.10, -0.01)
    frame = pd.DataFrame({"time": time, "close": 100.0, "signal": values, "forward_return": forward})

    profiles = relative_tail_profiles(frame.iloc[:30], frame.iloc[30:60], frame.iloc[60:], features=["signal"], feature_tail=0.20, tail_fraction=0.10)

    assert profiles["up"][0]["side"] == "high"
    assert profiles["up"][0]["cut"] == pytest.approx(7.2)
    assert profiles["up"][0]["windows"]["oos"]["lift"] > 1.0
