import numpy as np
import pandas as pd

from strategies.adaptive_lag_correlation import LagCorrelationConfig, Transform
from strategies.adaptive_lag_walkforward import (
    WalkForwardConfig,
    evaluate_adaptive_lag,
    evaluate_fixed_lag,
    moving_block_bootstrap_correlation,
)


def _series(values):
    return pd.Series(values, index=pd.date_range("2020-01-01", periods=len(values), freq="D"))


def test_adaptive_walkforward_beats_wrong_fixed_lag_after_regime_change():
    rng = np.random.default_rng(8)
    n, change, first_lag, second_lag = 1_200, 600, 30, 60
    source_values = rng.normal(size=n)
    target_values = np.full(n, np.nan)
    for position in range(second_lag, n):
        lag = first_lag if position < change else second_lag
        target_values[position] = source_values[position - lag] + rng.normal(scale=0.05)
    source, target = _series(source_values), _series(target_values)
    estimator = LagCorrelationConfig(
        min_lag=20, max_lag=70, window=180, min_observations=140, transform=Transform.LEVELS
    )
    config = WalkForwardConfig(estimator=estimator, evaluation_stride=7, start_at=target.index[300])

    adaptive = evaluate_adaptive_lag(source, target, config)
    fixed = evaluate_fixed_lag(source, target, lag=30, transform=Transform.LEVELS, stride=7, start_at=target.index[300])

    assert adaptive.summary.correlation > 0.75
    assert adaptive.summary.correlation > fixed.summary.correlation + 0.15
    assert adaptive.records["matures_at"].ge(adaptive.records["evaluated_at"]).all()


def test_walkforward_uses_only_past_data_at_each_selection_time():
    rng = np.random.default_rng(9)
    source = _series(rng.normal(size=800))
    target = _series(np.r_[np.full(40, np.nan), source.to_numpy()[:-40]])
    config = WalkForwardConfig(
        estimator=LagCorrelationConfig(min_lag=30, max_lag=50, window=160, min_observations=120),
        evaluation_stride=10,
        start_at=target.index[300],
    )

    before = evaluate_adaptive_lag(source, target, config)
    cutoff = target.index[650]
    source.loc[source.index > cutoff] = 1_000_000
    target.loc[target.index > cutoff] = -1_000_000
    after = evaluate_adaptive_lag(source, target, config)

    pd.testing.assert_frame_equal(
        before.records[before.records["matures_at"] <= cutoff],
        after.records[after.records["matures_at"] <= cutoff],
    )


def test_fixed_lag_and_bootstrap_are_deterministic():
    rng = np.random.default_rng(3)
    source = _series(rng.normal(size=500))
    target = _series(np.r_[np.full(20, np.nan), source.to_numpy()[:-20]])

    result = evaluate_fixed_lag(source, target, lag=20, transform=Transform.LEVELS, stride=5)
    interval_one = moving_block_bootstrap_correlation(result.records["source_value"], result.records["target_value"], block_length=10, resamples=100, seed=11)
    interval_two = moving_block_bootstrap_correlation(result.records["source_value"], result.records["target_value"], block_length=10, resamples=100, seed=11)

    assert result.summary.correlation > 0.99
    assert interval_one == interval_two
    assert interval_one[0] > 0.99
