"""Out-of-sample validation for adaptive lag selections."""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np
import pandas as pd

from strategies.adaptive_lag_correlation import (
    AdaptiveLagCorrelationEngine,
    LagCorrelationConfig,
    Transform,
    transform_series,
)


@dataclass(frozen=True)
class WalkForwardConfig:
    estimator: LagCorrelationConfig
    evaluation_stride: int = 1
    start_at: pd.Timestamp | str | None = None
    source_availability_delay: int = 0

    def __post_init__(self) -> None:
        if self.evaluation_stride < 1:
            raise ValueError("evaluation_stride must be >= 1")
        if self.source_availability_delay < 0:
            raise ValueError("source_availability_delay must be >= 0")


@dataclass(frozen=True)
class ForwardSummary:
    observations: int
    correlation: float | None
    directional_hit_rate: float | None
    mean_lag: float | None


@dataclass(frozen=True)
class WalkForwardResult:
    records: pd.DataFrame
    summary: ForwardSummary


def evaluate_adaptive_lag(
    source: pd.Series, target: pd.Series, config: WalkForwardConfig
) -> WalkForwardResult:
    """Select lag at t, then score source[t] against target[t + selected_lag]."""
    _validate_aligned(source, target)
    available_source = source.shift(config.source_availability_delay)
    source_values = transform_series(available_source, config.estimator.transform)
    target_values = transform_series(target, config.estimator.transform)
    engine = AdaptiveLagCorrelationEngine(config.estimator)
    rows = []
    for position in _evaluation_positions(target.index, config.evaluation_stride, config.start_at):
        estimate = engine.estimate_at(available_source, target, target.index[position])
        if estimate.best_lag is None:
            continue
        maturity_position = position + estimate.best_lag
        if maturity_position >= len(target):
            continue
        source_value = source_values.iloc[position]
        target_value = target_values.iloc[maturity_position]
        if not np.isfinite(source_value) or not np.isfinite(target_value):
            continue
        rows.append(
            {
                "evaluated_at": target.index[position],
                "matures_at": target.index[maturity_position],
                "selected_lag": estimate.best_lag,
                "selected_correlation": estimate.best_correlation,
                "source_value": source_value,
                "target_value": target_value,
            }
        )
    return _result(rows)


def evaluate_fixed_lag(
    source: pd.Series,
    target: pd.Series,
    *,
    lag: int,
    transform: Transform,
    stride: int = 1,
    start_at: pd.Timestamp | str | None = None,
    source_availability_delay: int = 0,
) -> WalkForwardResult:
    """Score a predeclared lag with the exact same forward-pair protocol."""
    _validate_aligned(source, target)
    if lag < 0:
        raise ValueError("lag must be >= 0")
    if stride < 1:
        raise ValueError("stride must be >= 1")
    if source_availability_delay < 0:
        raise ValueError("source_availability_delay must be >= 0")
    source_values = transform_series(source.shift(source_availability_delay), transform)
    target_values = transform_series(target, transform)
    rows = []
    for position in _evaluation_positions(target.index, stride, start_at):
        maturity_position = position + lag
        if maturity_position >= len(target):
            continue
        source_value = source_values.iloc[position]
        target_value = target_values.iloc[maturity_position]
        if not np.isfinite(source_value) or not np.isfinite(target_value):
            continue
        rows.append(
            {
                "evaluated_at": target.index[position],
                "matures_at": target.index[maturity_position],
                "selected_lag": lag,
                "selected_correlation": np.nan,
                "source_value": source_value,
                "target_value": target_value,
            }
        )
    return _result(rows)


def moving_block_bootstrap_correlation(
    source: pd.Series, target: pd.Series, *, block_length: int, resamples: int, seed: int
) -> tuple[float, float]:
    """Return a deterministic 95% moving-block bootstrap interval for Pearson r."""
    x, y = np.asarray(source, dtype=float), np.asarray(target, dtype=float)
    valid = np.isfinite(x) & np.isfinite(y)
    x, y = x[valid], y[valid]
    if len(x) < 2 or block_length < 1 or resamples < 1:
        raise ValueError("need at least two observations, positive block_length, and positive resamples")
    block_length = min(block_length, len(x))
    starts = np.arange(len(x) - block_length + 1)
    rng = np.random.default_rng(seed)
    correlations = []
    blocks_needed = int(np.ceil(len(x) / block_length))
    for _ in range(resamples):
        sampled_starts = rng.choice(starts, size=blocks_needed, replace=True)
        indices = np.concatenate([np.arange(start, start + block_length) for start in sampled_starts])[: len(x)]
        correlations.append(float(np.corrcoef(x[indices], y[indices])[0, 1]))
    return tuple(float(value) for value in np.quantile(correlations, [0.025, 0.975]))


def _evaluation_positions(index: pd.DatetimeIndex, stride: int, start_at: pd.Timestamp | str | None):
    start = pd.Timestamp(start_at) if start_at is not None else index[0]
    if index.tz is not None and start.tz is None:
        start = start.tz_localize(index.tz)
    elif index.tz is None and start.tz is not None:
        start = start.tz_localize(None)
    return range(int(index.searchsorted(start)), len(index), stride)


def _result(rows: list[dict]) -> WalkForwardResult:
    records = pd.DataFrame(rows)
    if records.empty:
        return WalkForwardResult(records, ForwardSummary(0, None, None, None))
    source, target = records["source_value"], records["target_value"]
    correlation = float(np.corrcoef(source, target)[0, 1]) if len(records) > 1 else None
    nonzero = (source != 0) & (target != 0)
    hit_rate = float((np.sign(source[nonzero]) == np.sign(target[nonzero])).mean()) if nonzero.any() else None
    return WalkForwardResult(
        records,
        ForwardSummary(len(records), correlation, hit_rate, float(records["selected_lag"].mean())),
    )


def _validate_aligned(source: pd.Series, target: pd.Series) -> None:
    if not isinstance(source.index, pd.DatetimeIndex) or not source.index.equals(target.index):
        raise ValueError("source and target must share the same DatetimeIndex")
    if not source.index.is_monotonic_increasing or not source.index.is_unique:
        raise ValueError("series index must be unique and monotonic increasing")
